Files
worldquant-alpha-system/backend/app/ai/tools.py
T

250 lines
10 KiB
Python
Raw Normal View History

"""Explicit business tool catalog. This module has no database or provider credentials."""
from datetime import datetime
from typing import Literal
from pydantic import Field
from ..backtests.contracts import ControlInput, PreviewInput, RerunInput, StartInput
from ..schemas import AlphaFilters, BulkInput, BulkUpdate, Contract, JobInput, ResearchInput, ResearchUpdate
class EmptyArgs(Contract):
pass
class SearchArgs(Contract):
filters: AlphaFilters = Field(default_factory=AlphaFilters)
class AlphaArgs(Contract):
alpha_id: str = Field(min_length=1, max_length=100, pattern=r"^[A-Za-z0-9_-]+$")
class JobArgs(Contract):
job_id: str = Field(min_length=1, max_length=100)
class ResearchArgs(AlphaArgs):
changes: ResearchInput
class ResultMetadata(Contract):
source: Literal["local_database"] = "local_database"
observed_at: datetime
nulls: str = "null 表示来源未提供,不等于零"
units: dict[str, str] = Field(
default_factory=lambda: {
"turnover": "比例,0.15 = 15%",
"returns": "比例",
"drawdown": "比例",
"margin": "比例",
"pnl": "供应商原始累计值,未提供货币/规模单位",
}
)
class BacktestRunArgs(Contract):
run_id: str = Field(min_length=1, max_length=36)
class BacktestListArgs(Contract):
limit: int = Field(default=20, ge=1, le=100)
offset: int = Field(default=0, ge=0)
source: str | None = Field(default=None, max_length=100)
class BacktestResultsArgs(BacktestRunArgs):
limit: int = Field(default=20, ge=1, le=100)
offset: int = Field(default=0, ge=0)
class BacktestPreviewArgs(Contract):
preview_id: str = Field(min_length=1, max_length=36)
limit: int = Field(default=20, ge=1, le=100)
offset: int = Field(default=0, ge=0)
class BacktestControlArgs(BacktestRunArgs):
action: Literal["pause", "resume", "stop", "recover"]
class BacktestRerunArgs(BacktestRunArgs):
item_ids: list[str] = Field(min_length=1, max_length=100)
CATALOG = {
"get_backtest_capabilities": (
EmptyArgs,
"读取回测输入 schema、支持类型和本地调度配置,不代表平台剩余额度。",
),
"prepare_backtest": (
PreviewInput,
"准备服务端固定回测预览,可用 inline 候选或草稿引用;只保存预览,不提交平台,不需要执行确认。",
),
"get_backtest_preview": (BacktestPreviewArgs, "分页读取完整固定预览,确认前核对表达式和最终参数。"),
"start_backtest": (
StartInput,
"对已保存预览请求一次用户确认,确认后后台运行全部固定候选,立即返回运行 ID;禁止循环等待。",
),
"list_backtests": (BacktestListArgs, "分页查询回测运行与统计,可按来源筛选。"),
"get_backtest": (BacktestRunArgs, "查询指定运行的真实进度,不循环等待完成。"),
"get_backtest_results": (
BacktestResultsArgs,
"分页读取逐项状态、历史指标和错误;未知结果不能推测为成功。",
),
"control_backtest": (
BacktestControlArgs,
"预览并确认暂停/继续/停止剩余项/找回原任务;不远端取消,不重新提交。",
),
"prepare_backtest_rerun": (
BacktestRerunArgs,
"从明确指定的已结束回测项准备新预览,保留来源;不会自动启动。",
),
"search_alphas": (
SearchArgs,
"按明确筛选条件查询本地 Alpha。Turnover 0.15 表示 15%;支持分页,禁止把当前页当作全部结果。",
),
"get_alpha_facets": (EmptyArgs, "获取可用地区、类型、状态、标签与本地 Alpha 总数。"),
"get_alpha": (AlphaArgs, "读取指定 Alpha 的表达式、指标、已有检查和本地研究记录。"),
"get_alpha_pnl": (AlphaArgs, "读取指定 Alpha 的 PnL 缓存摘要。缺失缓存时说明情况,不自动刷新。"),
"list_jobs": (EmptyArgs, "查询最近的同步任务,不要循环轮询等待。"),
"get_job_status": (JobArgs, "查询指定任务的状态、目标和错误,不要循环等待任务完成。"),
"update_research": (
ResearchArgs,
"提出指定 Alpha 的本地研究记录修改。仅传要修改的字段,等待用户在界面确认。",
),
"bulk_update_research": (BulkInput, "提出固定 1–100 个 Alpha 的批量标签或研究状态修改,等待用户确认。"),
"create_sync_job": (
JobInput,
"提出同步或本地自相关任务,等待确认。full_sync 仅同步已提交;待提交必须用 daily_sync 并指定 submission、date_from/date_to(UTC),待提交按创建日、已提交按提交日逐天同步。alpha_refresh/pnl_refresh/self_correlation 使用固定 alpha_ids;自相关缺失 PnL 时自动补取,不触发平台检查。创建后立即返回任务 ID。",
),
"cancel_job": (JobArgs, "提出取消指定同步任务,等待用户确认。"),
"retry_job": (JobArgs, "提出重试指定失败或暂停的同步任务,等待用户确认。"),
}
WRITES = {
"update_research",
"bulk_update_research",
"create_sync_job",
"cancel_job",
"retry_job",
"start_backtest",
"control_backtest",
}
def bounded(value):
if isinstance(value, str):
return value if len(value) <= 2000 else value[:2000] + "…(已截断)"
if isinstance(value, list):
return [bounded(v) for v in value[:100]]
if isinstance(value, dict):
return {k: bounded(v) for k, v in list(value.items())[:100]}
return value
async def read_tool(business, name, args):
from datetime import timezone
if name == "get_backtest_capabilities":
data = await business.backtests.capabilities()
elif name == "prepare_backtest":
data = await business.backtests.preview(args)
elif name == "get_backtest_preview":
data = await business.backtests.get_preview(**args.model_dump())
elif name == "list_backtests":
data = await business.backtests.runs(**args.model_dump())
elif name == "get_backtest":
data = await business.backtests.run(args.run_id)
elif name == "get_backtest_results":
data = await business.backtests.results(**args.model_dump())
# The complete historical response remains available through the business endpoint.
for item in data["items"]:
if item["result"]:
snapshot = item["result"].pop("snapshot")
item["result"].update({k: snapshot.get(k) for k in ("is", "os", "checks", "dateCreated")})
elif name == "prepare_backtest_rerun":
data = await business.backtests.rerun(args.run_id, RerunInput(item_ids=args.item_ids))
elif name == "search_alphas":
data = await business.search_alphas(args.filters)
data["filters"] = args.filters.model_dump(mode="json")
elif name == "get_alpha_pnl":
data = await business.get_alpha_pnl(args.alpha_id)
points = data.pop("points")
data.update(
alpha_id=args.alpha_id,
count=len(points),
first=points[0] if points else None,
last=points[-1] if points else None,
null_count=sum(p["value"] is None for p in points),
)
elif name in ("get_alpha", "get_job_status"):
data = await getattr(business, name)(*args.model_dump().values())
else:
data = await getattr(business, name)()
if isinstance(data, list):
data = {"items": data[:20]}
data["_meta"] = ResultMetadata(observed_at=datetime.now(timezone.utc)).model_dump(mode="json")
return bounded(data)
async def preview_tool(business, name, args):
if name == "start_backtest":
return {"backtest": await business.backtests.get_preview(args.preview_id)}
if name == "control_backtest":
return {"backtest_run": await business.backtests.run(args.run_id), "action": args.action}
if name in ("update_research", "bulk_update_research"):
ids = [args.alpha_id] if name == "update_research" else args.alpha_ids
targets, versions = [], {}
for alpha_id in ids:
detail = await business.get_alpha(alpha_id)
before = detail["research"]
versions[alpha_id] = before["version"]
if name == "update_research":
after = {**before, **args.changes.model_dump(exclude_unset=True)}
else:
after = {
**before,
"tags": sorted((set(before["tags"]) | set(args.add_tags)) - set(args.remove_tags)),
}
if args.state:
after["state"] = args.state
# Preview and execution use the same validation rules.
ResearchInput.model_validate({k: after[k] for k in ("note", "tags", "favorite", "state")})
targets.append({"alpha_id": alpha_id, "before": before, "after": after})
return {"targets": targets, "versions": versions}
if name in ("cancel_job", "retry_job"):
return {"job": await business.get_job_status(args.job_id)}
return {"operation": args.model_dump(mode="json")}
async def execute_tool(business, name, args, preview):
if name == "start_backtest":
current = await business.backtests.get_preview(args.preview_id)
if current["digest"] != preview["backtest"]["digest"] or current["version"] != args.version:
from fastapi import HTTPException
raise HTTPException(409, "回测预览不匹配,请重新确认")
return await business.backtests.start(args)
if name == "control_backtest":
return await business.backtests.control(
args.run_id, ControlInput(action=args.action, version=preview["backtest_run"]["version"])
)
if name == "update_research":
body = ResearchUpdate(
**args.changes.model_dump(exclude_unset=True), version=preview["versions"][args.alpha_id]
)
return await business.update_research(args.alpha_id, body)
if name == "bulk_update_research":
return await business.bulk_update_research(
BulkUpdate(**args.model_dump(), versions=preview["versions"])
)
if name == "create_sync_job":
return await business.create_sync_job(args)
current = await business.get_job_status(args.job_id)
if current["updated_at"] != preview["job"]["updated_at"] or current["status"] != preview["job"]["status"]:
from fastapi import HTTPException
raise HTTPException(409, "任务状态已变化,请重新确认操作")
return await getattr(business, name)(args.job_id)