"""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")