130 lines
5.4 KiB
Python
130 lines
5.4 KiB
Python
"""Resolve fixed data inputs and produce previews without submitting simulations.
|
||
|
||
Callers own authorization and transactions. Binding checks establish provenance,
|
||
not FASTEXPR operator semantics or the account's current platform permissions.
|
||
"""
|
||
|
||
from fastapi import HTTPException
|
||
from sqlalchemy import select
|
||
|
||
from ..backtests.contracts import Candidate, DraftInput, PreviewInput, Source
|
||
from ..catalog.contracts import EntryOutput, InputPreparation
|
||
from ..catalog.service import Catalog
|
||
from ..models import CatalogEntry
|
||
from .expressions import analyze, expand
|
||
|
||
|
||
class ResearchBuilder:
|
||
def __init__(self, db, backtests):
|
||
self.db = db
|
||
self.catalog = Catalog(db)
|
||
self.backtests = backtests
|
||
|
||
async def select_input(self, body):
|
||
"""Fix explicit fields in one published version; reject missing or stale members."""
|
||
collection = await self.catalog.collection(body.scope, body.dataset_id)
|
||
chosen = set(body.field_ids)
|
||
if len(chosen) != len(body.field_ids) or not chosen.issubset(collection["field_ids"]):
|
||
raise HTTPException(422, "字段选择含重复、未知或其他数据集字段")
|
||
saved = await self.catalog.prepare(
|
||
InputPreparation(
|
||
scope=body.scope,
|
||
dataset_id=body.dataset_id,
|
||
collection_version=body.collection_version,
|
||
selection="explicit",
|
||
excluded_ids=[field for field in collection["field_ids"] if field not in chosen],
|
||
)
|
||
)
|
||
return await self.input_page(saved["id"])
|
||
|
||
async def input_page(self, input_id, limit=25, offset=0, q="", field_type=None):
|
||
"""Read the saved version, including field descriptions, with explicit pagination."""
|
||
saved = await self.catalog.input(input_id)
|
||
ids = [
|
||
field
|
||
for field in saved["field_ids"]
|
||
if q.lower() in field.lower()
|
||
and (field_type is None or saved["field_types"].get(field) == field_type)
|
||
]
|
||
page = ids[offset : offset + limit]
|
||
entries = {
|
||
row.id: row
|
||
for row in await self.db.scalars(
|
||
select(CatalogEntry).where(
|
||
CatalogEntry.batch_id == saved["collection_version"], CatalogEntry.id.in_(page)
|
||
)
|
||
)
|
||
}
|
||
return {
|
||
**{
|
||
k: saved[k]
|
||
for k in ("id", "scope", "dataset_id", "collection_version", "selection", "created_at")
|
||
},
|
||
"field_count": len(saved["field_ids"]),
|
||
"items": [
|
||
EntryOutput.model_validate(entries[field], from_attributes=True).model_dump()
|
||
for field in page
|
||
],
|
||
"total": len(ids),
|
||
"limit": limit,
|
||
"offset": offset,
|
||
"has_more": offset + limit < len(ids),
|
||
}
|
||
|
||
async def prepare(self, body):
|
||
"""Bind templates against an immutable input, then reuse the fixed-preview interface.
|
||
|
||
Raises HTTPException(422) for wrong scope, membership or declared type.
|
||
No expression execution or implicit cleaning/aggregation takes place here.
|
||
"""
|
||
saved = await self.catalog.input(body.template_input_id)
|
||
scope = saved["scope"]
|
||
candidates = []
|
||
for item in body.candidates:
|
||
settings = item.settings
|
||
if (
|
||
settings.instrumentType != scope["instrument_type"]
|
||
or settings.region != scope["region"]
|
||
or settings.universe != scope["universe"]
|
||
or settings.delay != scope["delay"]
|
||
):
|
||
raise HTTPException(422, "候选模拟参数与输入快照的研究范围不一致")
|
||
for binding in item.bindings.values():
|
||
if binding.field_id not in saved["field_ids"]:
|
||
raise HTTPException(422, "绑定字段不属于该输入快照,不能使用被排除或其他数据集字段")
|
||
if saved["field_types"].get(binding.field_id) != binding.field_type:
|
||
raise HTTPException(422, "字段类型声明与输入快照不一致,未知类型不能自动构建")
|
||
expression = expand(
|
||
item.expression_template,
|
||
{name: [binding.field_id] for name, binding in item.bindings.items()},
|
||
limit=1,
|
||
)["items"][0]["expression"]
|
||
validation = analyze(expression, saved["field_types"])
|
||
if validation["syntax"] or validation["types"]:
|
||
raise HTTPException(422, ";".join(validation["syntax"] + validation["types"]))
|
||
if len(expression) > 20000:
|
||
raise HTTPException(422, "绑定后的表达式超过 20000 字符")
|
||
candidates.append(
|
||
Candidate(
|
||
client_item_id=item.client_item_id,
|
||
expression=expression,
|
||
settings=settings,
|
||
)
|
||
)
|
||
source = Source.model_validate(
|
||
{
|
||
**body.source.model_dump(),
|
||
"template_input_id": saved["id"],
|
||
"hypothesis": body.hypothesis,
|
||
}
|
||
)
|
||
return await self.backtests.preview(
|
||
PreviewInput(
|
||
inline=DraftInput(
|
||
name=body.name,
|
||
source=source,
|
||
candidates=candidates,
|
||
)
|
||
)
|
||
)
|