Files
worldquant-alpha-system/backend/tests/ai_fake.py
T

134 lines
5.2 KiB
Python
Raw Normal View History

"""Deterministic model for isolated acceptance; never calls a provider or platform."""
import asyncio
import json
from contextlib import asynccontextmanager
from uuid import uuid4
from pydantic_ai.messages import ModelResponse, TextPart, ToolCallPart, ToolReturnPart, UserPromptPart
from pydantic_ai.models.function import DeltaToolCall, FunctionModel
from tests.research_fake import research_step
async def fake_stream(messages, info):
latest = max(
(i for i, m in enumerate(messages) if any(isinstance(p, UserPromptPart) for p in m.parts)), default=0
)
text = " ".join(
str(p.content) for m in messages[latest:] for p in m.parts if isinstance(p, UserPromptPart)
)
returns = [p for m in messages[latest:] for p in m.parts if isinstance(p, ToolReturnPart)]
if any(marker in text for marker in ("研究此输入", "自行选字段研究", "解读研究结果")):
step = research_step(
text, returns, [p for m in messages for p in m.parts if isinstance(p, ToolReturnPart)]
)
if isinstance(step, str):
yield step
else:
name, args = step
yield {0: DeltaToolCall(name=name, json_args=json.dumps(args), tool_call_id=uuid4().hex)}
return
if returns and returns[-1].tool_name == "prepare_backtest":
content = returns[-1].content
content = json.loads(content) if isinstance(content, str) else content
yield {
0: DeltaToolCall(
name="start_backtest",
json_args=json.dumps(
{
"preview_id": content["preview_id"],
"version": 1,
"idempotency_key": content["preview_id"],
}
),
tool_call_id=uuid4().hex,
)
}
return
if returns and "LOOP" not in text:
if returns[-1].tool_name == "capability_probe":
yield str(returns[-1].content)
else:
yield "操作结果已返回。"
yield "请查看下方业务记录与数据来源。"
return
if any(t.name == "capability_probe" for t in info.function_tools):
name, args = "capability_probe", {}
elif "READY" in text:
yield "REA"
yield "DY"
return
elif "SLOW" in text:
yield "正在查询"
await asyncio.sleep(2)
yield ",查询完成。"
return
elif "回测" in text:
name, args = (
"prepare_backtest",
{
"inline": {
"name": "AI 固定回测",
"source": {"kind": "ai"},
"candidates": [
{
"client_item_id": "ai-1",
"expression": "rank(close)",
"settings": {"region": "USA", "universe": "TOP3000", "delay": 1},
}
],
}
},
)
elif "批量" in text:
name, args = "bulk_update_research", {"alpha_ids": ["a0000", "a0001"], "add_tags": ["AI"]}
elif "修改" in text or "update" in text:
alpha_id = "TEST0001" if "TEST0001" in text else "a0000"
name, args = "update_research", {"alpha_id": alpha_id, "changes": {"note": "AI 测试研究记录"}}
elif "同步" in text:
name, args = "create_sync_job", {"kind": "full_sync"}
elif "PnL" in text:
name, args = "get_alpha_pnl", {"alpha_id": "TEST0001" if "TEST" in text else "a0000"}
else:
name, args = "search_alphas", {"filters": {"turnover_max": 0.15, "limit": 5}}
yield {0: DeltaToolCall(name=name, json_args=json.dumps(args), tool_call_id=uuid4().hex)}
def fake_structured(messages, info):
if not info.output_tools:
return ModelResponse(parts=[TextPart("READY")])
tool = info.output_tools[0]
context = next(
(json.loads(p.content) for m in reversed(messages) for p in m.parts if isinstance(p, UserPromptPart)),
{},
)
fields = [name for name, kind in context.get("fields", {}).items() if kind == "MATRIX"][:2]
template = {
"name": "合成流水线模板",
"description": "合成模型研究假设",
"expression": "rank({field})",
"variables": {
"field": {"kind": "field", "field_type": "MATRIX", "values": fields or ["TEST_FIN_001"]}
},
}
properties = tool.parameters_json_schema.get("properties", {})
if "summary" in properties:
data = {"summary": "合成评估建议", "risks": ["仅供验收"], "suggestions": ["继续核实缺失证据"]}
elif "input_ids" in properties:
data = {
"name": "合成特征方案",
"hypothesis": context.get("hypothesis", "合成假设"),
"input_ids": [i["id"] for i in context.get("inputs", [])],
"steps": [],
"template": template,
}
else:
data = template
return ModelResponse(parts=[ToolCallPart(tool.name, data)])
@asynccontextmanager
async def fake_model(config, settings):
yield FunctionModel(function=fake_structured, stream_function=fake_stream, model_name="test-model")