Files
worldquant-alpha-system/backend/app/research/assets.py
T

176 lines
7.6 KiB
Python
Raw Normal View History

"""Versioned research assets. Mutations use optimistic versions; revisions are immutable."""
from fastapi import HTTPException
from sqlalchemy import func, select, update
from ..backtests.contracts import fingerprint
from ..backtests.service import uid
from ..models import Account, ResearchAsset, ResearchRevision, now
from .serialization import encode_snapshot as jsonable_encoder
from .workspace_contracts import FeatureSpec, TemplateSpec, ViewSpec, WorkflowSpec
class Assets:
def __init__(self, db):
self.db = db
async def get(self, asset_id, version=None, expected_kind=None):
asset = await self.db.get(ResearchAsset, asset_id)
if not asset or (expected_kind and asset.kind != expected_kind):
raise HTTPException(404, "研究素材不存在或类型不匹配")
revision = await self.db.get(ResearchRevision, (asset_id, version or asset.version))
if not revision:
raise HTTPException(404, "素材版本不存在")
return jsonable_encoder(
{
"id": asset.id,
"kind": asset.kind,
"name": revision.content["name"],
"version": revision.version,
"latest_version": asset.version,
"archived": asset.archived,
"content": revision.content,
"provenance": revision.provenance,
"created_at": revision.created_at,
}
)
async def list(self, kind, q="", limit=25, offset=0):
query = select(ResearchAsset).where(ResearchAsset.kind == kind, ResearchAsset.archived.is_(False))
if q:
query = query.where(ResearchAsset.name.ilike(f"%{q}%"))
total = await self.db.scalar(select(func.count()).select_from(query.subquery()))
rows = await self.db.scalars(
query.order_by(ResearchAsset.updated_at.desc(), ResearchAsset.id).limit(limit).offset(offset)
)
return {
"items": [await self.get(row.id) for row in rows],
"total": total,
"limit": limit,
"offset": offset,
}
async def save(self, body, asset_id=None, provenance=None):
schema = {
"template": TemplateSpec,
"feature": FeatureSpec,
"view": ViewSpec,
"workflow": WorkflowSpec,
}[body.kind]
content = schema.model_validate(body.content).model_dump(mode="json")
if body.kind == "workflow":
from .workflows import validate_graph
validate_graph(WorkflowSpec.model_validate(content))
if body.kind == "feature":
from ..catalog.service import Catalog
for input_id in content["input_ids"]:
await Catalog(self.db).input(input_id)
if asset_id:
if body.version is None:
raise HTTPException(422, "更新需要素材版本")
changed = await self.db.execute(
update(ResearchAsset)
.where(
ResearchAsset.id == asset_id,
ResearchAsset.version == body.version,
ResearchAsset.kind == body.kind,
)
.values(version=body.version + 1, name=content["name"], updated_at=now())
)
if changed.rowcount != 1:
raise HTTPException(409, "素材已变化,保留草稿并读取最新版本")
version = body.version + 1
else:
asset_id, version = uid(), 1
self.db.add(ResearchAsset(id=asset_id, kind=body.kind, name=content["name"], version=version))
await self.db.flush()
self.db.add(
ResearchRevision(asset_id=asset_id, version=version, content=content, provenance=provenance or {})
)
await self.db.flush()
return await self.get(asset_id, version)
async def archive(self, asset_id, version):
result = await self.db.execute(
update(ResearchAsset)
.where(ResearchAsset.id == asset_id, ResearchAsset.version == version)
.values(archived=True, version=ResearchAsset.version + 1, updated_at=now())
)
if result.rowcount != 1:
raise HTTPException(409, "素材已变化或不存在")
# Archiving is itself a revision; old references remain resolvable.
previous = await self.db.get(ResearchRevision, (asset_id, version))
self.db.add(
ResearchRevision(
asset_id=asset_id,
version=version + 1,
content=previous.content,
provenance=previous.provenance,
)
)
return {"ok": True}
async def versions(self, asset_id):
await self.get(asset_id)
rows = await self.db.scalars(
select(ResearchRevision)
.where(ResearchRevision.asset_id == asset_id)
.order_by(ResearchRevision.version.desc())
)
return jsonable_encoder([{"version": row.version, "created_at": row.created_at} for row in rows])
async def import_preview(self, templates):
normalized, errors = [], []
for index, item in enumerate(templates):
try:
converted = dict(item)
if "templateConfigurations" in converted:
config = converted.pop("templateConfigurations")
if not isinstance(config, dict):
raise ValueError("旧变量配置需要对象格式,请转换后重试")
converted["variables"] = {
key: value
if isinstance(value, dict) and "kind" in value
else {
"kind": "fragment",
"values": value.get("variables", []) if isinstance(value, dict) else value,
}
for key, value in config.items()
}
for key in ("createdAt", "updatedAt", "id", "version"):
converted.pop(key, None)
normalized.append(TemplateSpec.model_validate(converted).model_dump(mode="json"))
except (ValueError, TypeError) as exc:
errors.append({"index": index, "message": str(exc)})
names = [item["name"] for item in normalized]
existing = list(
await self.db.scalars(
select(ResearchAsset).where(ResearchAsset.kind == "template", ResearchAsset.name.in_(names))
)
)
conflicts = [{"id": item.id, "name": item.name, "version": item.version} for item in existing]
if len(set(names)) != len(names):
errors.append({"index": -1, "message": "导入文件内模板名称重复"})
return {
"templates": normalized,
"conflicts": conflicts,
"errors": errors,
"digest": fingerprint({"templates": normalized, "conflicts": conflicts}),
"policy": "仅创建新模板;同名请修改名称,或在模板编辑器中查看差异后保存新版本",
}
async def import_commit(self, body):
from .workspace_contracts import AssetWrite
await self.db.scalar(select(Account).where(Account.id == 1).with_for_update())
preview = await self.import_preview([item.model_dump(mode="json") for item in body.templates])
if preview["digest"] != body.digest or preview["conflicts"] or preview["errors"]:
raise HTTPException(409, "导入预览已变化或存在冲突,请重新预览")
return {
"items": [
await self.save(AssetWrite(kind="template", content=item)) for item in preview["templates"]
]
}