89 lines
3.3 KiB
Python
89 lines
3.3 KiB
Python
"""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 ToolReturnPart, UserPromptPart
|
|
from pydantic_ai.models.function import DeltaToolCall, FunctionModel
|
|
|
|
|
|
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 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)}
|
|
|
|
|
|
@asynccontextmanager
|
|
async def fake_model(config, settings):
|
|
yield FunctionModel(stream_function=fake_stream, model_name="test-model")
|