feat: integrate chatbox research with datasets and backtests
This commit is contained in:
@@ -0,0 +1 @@
|
||||
"""Research producers prepare candidates; the backtest module owns execution."""
|
||||
@@ -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)
|
||||
@@ -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
|
||||
],
|
||||
}
|
||||
)
|
||||
@@ -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,
|
||||
)
|
||||
)
|
||||
)
|
||||
Reference in New Issue
Block a user