feat: add AI research chatbot with confirmed business tools

This commit is contained in:
yuxuanhui
2026-09-07 23:02:55 +08:00
parent 3cd280d068
commit 79ab20b4ea
44 changed files with 4647 additions and 197 deletions
+150
View File
@@ -0,0 +1,150 @@
"""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 ..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": "供应商原始累计值,未提供货币/规模单位",
}
)
CATALOG = {
"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,
"提出全量同步、指定 Alpha 刷新或 PnL 刷新任务,等待确认;创建后立即返回任务 ID。",
),
"cancel_job": (JobArgs, "提出取消指定同步任务,等待用户确认。"),
"retry_job": (JobArgs, "提出重试指定失败或暂停的同步任务,等待用户确认。"),
}
WRITES = {"update_research", "bulk_update_research", "create_sync_job", "cancel_job", "retry_job"}
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 == "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 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 == "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)