"""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 snapshots = [await Catalog(self.db).input(input_id) for input_id in content["input_ids"]] provenance = {**(provenance or {}), "inputs": snapshots} 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=jsonable_encoder(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"] ] }