"""Bounded graph traversal with explicit continuation, retaining every result source.""" from fastapi import HTTPException from sqlalchemy import or_, select from ..models import Alpha, BacktestResult, BacktestRun, ResearchExperiment, ResearchParent from .experiments import Experiments from .provenance import alpha_sources, saved_sources async def lineage(db, alpha_id=None, experiment_id=None, limit=25, offset=0): if bool(alpha_id) == bool(experiment_id): raise HTTPException(422, "指定 Alpha 或实验之一") sources = None frontier = set() truncated = False if alpha_id: if not await db.get(Alpha, alpha_id): raise HTTPException(404, "Alpha 尚未同步") sources = await alpha_sources(db, alpha_id, limit, offset) source_ids = list( await db.scalars( saved_sources() .with_only_columns(BacktestRun.source["research_id"].as_string()) .where(BacktestResult.alpha_id == alpha_id) .distinct() .limit(101) ) ) children = list( await db.scalars( select(ResearchParent.child_id) .where(ResearchParent.parent_kind == "alpha", ResearchParent.parent_id == alpha_id) .order_by(ResearchParent.child_id) .limit(101) ) ) frontier = set(source_ids + children) - {None} truncated = len(frontier) > 100 else: await Experiments(db).get(experiment_id) frontier.add(experiment_id) found, edges = {}, {} for _ in range(8): wanted = sorted(frontier - set(found)) if not wanted: break remaining = 100 - len(found) if len(wanted) > remaining: truncated = True wanted = wanted[:remaining] if not wanted: break rows = list(await db.scalars(select(ResearchExperiment).where(ResearchExperiment.id.in_(wanted)))) for row in rows: found[row.id] = await Experiments(db).get(row.id) produced = list( await db.scalars( saved_sources() .with_only_columns(BacktestResult.alpha_id) .where(BacktestRun.source["research_id"].as_string().in_(wanted)) .distinct() .limit(1001) ) ) truncated |= len(produced) > 1000 parent_alphas = {p["id"] for row in rows for p in row.parents if p["kind"] == "alpha"} related_sources = list( await db.scalars( saved_sources() .with_only_columns(BacktestRun.source["research_id"].as_string()) .where(BacktestResult.alpha_id.in_(parent_alphas | set(produced[:1000]))) .distinct() .limit(101) ) ) truncated |= len(related_sources) > 100 relations = list( await db.scalars( select(ResearchParent) .where( or_( ResearchParent.child_id.in_(wanted), (ResearchParent.parent_kind == "alpha") & ResearchParent.parent_id.in_(produced[:1000]), (ResearchParent.parent_kind == "experiment") & ResearchParent.parent_id.in_(wanted), ) ) .order_by(ResearchParent.child_id, ResearchParent.parent_kind, ResearchParent.parent_id) .limit(1001) ) ) truncated |= len(relations) > 1000 frontier = set(related_sources[:100]) - {None} for edge in relations[:1000]: edges[(edge.child_id, edge.parent_kind, edge.parent_id)] = { "child_id": edge.child_id, "parent_kind": edge.parent_kind, "parent_id": edge.parent_id, } frontier.add(edge.child_id) if edge.parent_kind == "experiment": frontier.add(edge.parent_id) unresolved = sorted(frontier - set(found)) return { "items": list(found.values()), "edges": list(edges.values()), "sources": sources, "truncated": truncated or bool(unresolved), "unresolved_experiment_ids": unresolved, "limit": 100, "max_depth": 8, }