feat: integrate chatbox research with datasets and backtests

This commit is contained in:
yuxuanhui
2026-09-08 12:43:00 +08:00
parent 43336ad960
commit aef8e1d310
37 changed files with 1421 additions and 42 deletions
+1
View File
@@ -0,0 +1 @@
"""Research producers prepare candidates; the backtest module owns execution."""
+66
View File
@@ -0,0 +1,66 @@
"""Explicit snapshot and field-binding contracts for research producers."""
import re
from typing import Literal
from pydantic import Field, model_validator
from ..backtests.contracts import SimulationSettings, Source
from ..catalog.contracts import Scope
from ..schemas import Contract
PLACEHOLDER = re.compile(r"\{([A-Za-z_][A-Za-z0-9_]*)\}")
class ResearchInputSelection(Contract):
scope: Scope
dataset_id: str = Field(min_length=1, max_length=200)
collection_version: str = Field(min_length=1, max_length=36)
field_ids: list[str] = Field(min_length=1, max_length=100)
class InputPageArgs(Contract):
input_id: str = Field(min_length=1, max_length=36)
limit: int = Field(default=25, ge=1, le=100)
offset: int = Field(default=0, ge=0)
q: str = Field(default="", max_length=300)
field_type: str | None = Field(default=None, max_length=100)
class FieldBinding(Contract):
field_id: str = Field(min_length=1, max_length=200, pattern=r"^[A-Za-z_][A-Za-z0-9_]*$")
field_type: Literal["MATRIX", "VECTOR", "GROUP"]
class ResearchCandidate(Contract):
client_item_id: str = Field(min_length=1, max_length=100)
expression_template: str = Field(min_length=1, max_length=20000)
bindings: dict[str, FieldBinding] = Field(min_length=1, max_length=100)
settings: SimulationSettings
@model_validator(mode="after")
def complete_bindings(self):
placeholders = set(PLACEHOLDER.findall(self.expression_template))
remainder = PLACEHOLDER.sub("", self.expression_template)
if placeholders != set(self.bindings) or "{" in remainder or "}" in remainder:
raise ValueError("模板占位符必须与字段绑定逐一对应,例如 rank({price})")
return self
class ChatboxResearchInput(Contract):
"""Chatbox provenance is supplied by the server, never by model arguments."""
name: str = Field(min_length=1, max_length=200)
hypothesis: str = Field(min_length=1, max_length=2000)
template_input_id: str = Field(min_length=1, max_length=36)
candidates: list[ResearchCandidate] = Field(min_length=1, max_length=100)
@model_validator(mode="after")
def unique_candidates(self):
if len({item.client_item_id for item in self.candidates}) != len(self.candidates):
raise ValueError("client_item_id 在候选集合内必须唯一")
return self
class ResearchPreviewInput(ChatboxResearchInput):
source: Source = Field(default_factory=Source)
+70
View File
@@ -0,0 +1,70 @@
"""Read provenance from saved results, retaining every experiment for an Alpha."""
from collections import defaultdict
from fastapi import HTTPException
from fastapi.encoders import jsonable_encoder
from sqlalchemy import func, select
from ..models import Alpha, BacktestItem, BacktestResult, BacktestRun
def saved_sources():
"""Only persisted results establish provenance; pending platform IDs do not."""
return (
select(BacktestResult, BacktestItem, BacktestRun)
.select_from(BacktestResult)
.join(BacktestItem, BacktestResult.item_id == BacktestItem.id)
.join(BacktestRun, BacktestItem.run_id == BacktestRun.id)
)
def source_alpha_ids(source=None, source_reference=None, research_id=None, backtest_run_id=None):
"""An IN subquery keeps list counts and exports independent of source multiplicity."""
query = saved_sources().with_only_columns(BacktestResult.alpha_id)
for key, value in (("kind", source), ("reference", source_reference), ("research_id", research_id)):
if value:
query = query.where(BacktestRun.source[key].as_string() == value)
if backtest_run_id:
query = query.where(BacktestRun.id == backtest_run_id)
return query
async def source_kinds(db, alpha_ids=None):
query = saved_sources().with_only_columns(BacktestResult.alpha_id, BacktestRun.source["kind"].as_string())
if alpha_ids is not None:
query = query.where(BacktestResult.alpha_id.in_(alpha_ids))
values = defaultdict(list)
for alpha_id, kind in await db.execute(query.distinct()):
if kind:
values[alpha_id].append(kind)
return {alpha_id: sorted(kinds) for alpha_id, kinds in values.items()}
async def alpha_sources(db, alpha_id, limit=25, offset=0):
if not await db.get(Alpha, alpha_id):
raise HTTPException(404, "Alpha 尚未同步")
query = saved_sources().where(BacktestResult.alpha_id == alpha_id)
total = await db.scalar(select(func.count()).select_from(query.subquery()))
rows = await db.execute(
query.order_by(BacktestResult.observed_at.desc(), BacktestResult.item_id).limit(limit).offset(offset)
)
return jsonable_encoder(
{
"alpha_id": alpha_id,
"total": total,
"limit": limit,
"offset": offset,
"items": [
{
"backtest_run_id": run.id,
"name": run.name,
"source": run.source,
"item_id": item.id,
"client_item_id": item.client_item_id,
"observed_at": result.observed_at,
}
for result, item, run in rows
],
}
)
+124
View File
@@ -0,0 +1,124 @@
"""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 .contracts import PLACEHOLDER
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 = PLACEHOLDER.sub(
lambda match: item.bindings[match.group(1)].field_id, item.expression_template
)
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,
)
)
)