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

114 lines
4.3 KiB
Python

"""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,
}