114 lines
4.3 KiB
Python
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,
|
||
|
|
}
|