diff --git a/backend/app/agent/llm.py b/backend/app/agent/llm.py new file mode 100644 index 0000000..d0fd9fb --- /dev/null +++ b/backend/app/agent/llm.py @@ -0,0 +1,77 @@ +"""LLM 客户端抽象(Phase 5)。 + +实现约定:真实 Key 来自 .env(LLM_API_KEY / LLM_BASE_URL / LLM_MODEL,见 config.yaml 引用)。 +测试注入 FakeLLMClient 走完整编排链路,不触网。 +""" + +from __future__ import annotations + +from typing import Protocol + +import httpx + +from app.agent.tools import tools_schema +from app.agent.tools_impl import build_tools +from app.core.config import get_settings + +SYSTEM_TEMPLATE = """你是个人 A 股量化研究助手(Research Assistant),不是系统管理员。 + +你可以调用以下工具(每次只能输出一个动作): +{tools} + +输出规则:只输出一行 JSON,两种形态之一: +1. 需要调用工具:{{"tool": "<工具名>", "args": {{...}}}} +2. 给出结论:{{"final": "结论文本"}} + +研究纪律(必须遵守): +- 先提出假设 → 用工具做因子测试或回测 → 基于实验事实分析,再给结论 +- 不得仅凭单次回测高收益就宣布策略有效;要主动说明样本外、过拟合、 + look-ahead bias、交易成本、参数敏感性等风险(未做验证的项要明说「未验证」) +- 全程只读:不得要求删除/修改数据或执行任意命令(你也没有这类工具) +- 回答使用简体中文 +""" + + +class LLMClient(Protocol): + def chat(self, messages: list[dict]) -> str: ... + + +class OpenAICompatibleClient: + """OpenAI 兼容 Chat Completions(qwen/dashscope、deepseek、openai 等均适用)。""" + + def __init__(self, *, base_url: str, api_key: str, model: str, timeout: float = 60.0) -> None: + self._base_url = base_url.rstrip("/") + self._api_key = api_key + self._model = model + self._timeout = timeout + + def chat(self, messages: list[dict]) -> str: + with httpx.Client(timeout=self._timeout) as client: + resp = client.post( + f"{self._base_url}/chat/completions", + headers={"Authorization": f"Bearer {self._api_key}"}, + json={ + "model": self._model, + "messages": messages, + "temperature": 0.2, + }, + ) + resp.raise_for_status() + data = resp.json() + return data["choices"][0]["message"]["content"] + + +def build_llm_from_settings(): + """从配置构造 LLM;未配置 Key 时返回 None(调用方给出引导提示)。""" + settings = get_settings() + if not settings.llm_api_key: + return None + return OpenAICompatibleClient( + base_url=settings.llm_base_url or "https://dashscope.aliyuncs.com/compatible-mode/v1", + api_key=settings.llm_api_key, + model=settings.llm_model, + ) + + +def system_prompt() -> str: + return SYSTEM_TEMPLATE.format(tools=tools_schema(build_tools())) diff --git a/backend/app/agent/service.py b/backend/app/agent/service.py new file mode 100644 index 0000000..58f75bc --- /dev/null +++ b/backend/app/agent/service.py @@ -0,0 +1,98 @@ +"""Agent 编排:自然语言 → 受控工具调用循环 → 结论(AGENT.md §29 假设-实验-分析循环)。""" + +from __future__ import annotations + +import json +import re +from typing import Any + +from app.agent.llm import SYSTEM_TEMPLATE, LLMClient +from app.agent.tools import Tool, tools_schema +from app.agent.tools_impl import build_tools + +MAX_TOOL_ROUNDS = 5 +_DECISION_PATTERN = re.compile(r"\{.*\}", re.DOTALL) + + +def _parse_decision(text: str) -> dict[str, Any]: + """容忍 LLM 输出中的代码块/前后缀,提取首个 JSON 对象。""" + match = _DECISION_PATTERN.search(text) + if not match: + raise ValueError(f"无法从模型输出中解析动作 JSON:{text[:200]}") + try: + return json.loads(match.group(0)) + except json.JSONDecodeError as exc: + raise ValueError(f"模型输出的 JSON 不合法:{text[:200]}") from exc + + +class AgentService: + """受控研究 Agent:每轮让 LLM 决策 tool/final,执行工具并把结果回喂,直至 final。""" + + def __init__(self, llm: LLMClient, tools: list[Tool] | None = None) -> None: + self._llm = llm + self._tools = {t.name: t for t in (tools or build_tools())} + + def _by_name(self, name: str) -> Tool: + tool = self._tools.get(name) + if tool is None: + raise ValueError(f"工具不存在:{name}(可用 {sorted(self._tools)})") + return tool + + def chat(self, message: str, max_rounds: int = MAX_TOOL_ROUNDS) -> dict: + messages: list[dict] = [ + { + "role": "system", + "content": SYSTEM_TEMPLATE.format(tools=tools_schema(list(self._tools.values()))), + }, + {"role": "user", "content": message}, + ] + actions: list[dict] = [] + + for _round in range(max_rounds): + try: + decision = _parse_decision(self._llm.chat(messages)) + except ValueError as exc: + # LLM 回复格式异常:回传错误要求重试一次结构化输出 + messages.append( + { + "role": "user", + "content": f"输出格式错误,请只输出一行 JSON(tool 或 final):{exc}", + } + ) + continue + if "final" in decision: + return {"reply": str(decision["final"]), "actions": actions} + tool_name = str(decision.get("tool", "")) + args = decision.get("args") or {} + if not isinstance(args, dict): + args = {} + tool = self._by_name(tool_name) # 白名单之外的调用直接报错 + output = tool.invoke(args) + actions.append({"tool": tool_name, "args": args, "output": output[:1000]}) + messages.append( + { + "role": "assistant", + "content": f'{{"tool": "{tool_name}", "args": {json.dumps(args, ensure_ascii=False)}}}', + } + ) + messages.append( + { + "role": "user", + "content": f"工具 {tool_name} 返回:\n{output}\n请继续(如需再调用输出 tool,否则输出 final)。", + } + ) + + # 轮次耗尽:强制收尾 + messages.append( + { + "role": "user", + "content": "已达工具调用轮次上限,请直接给出基于已有事实的最终结论(只输出 final JSON)。", + } + ) + try: + decision = _parse_decision(self._llm.chat(messages)) + except ValueError: + decision = { + "final": "研究轮次耗尽且模型未给出结构化结论,请人工查看 actions 中的实验输出。" + } + return {"reply": str(decision.get("final", decision)), "actions": actions} diff --git a/backend/app/agent/tools.py b/backend/app/agent/tools.py new file mode 100644 index 0000000..608992e --- /dev/null +++ b/backend/app/agent/tools.py @@ -0,0 +1,45 @@ +"""AI Research Agent(Phase 5)。 + +Agent 是 Research Assistant(AGENT.md §28/§29): +- 只能调用本目录 tools 提供的**白名单受控工具**(只读研究能力) +- 禁止:执行任意 shell / 修改删除数据 / 修改配置与凭证 +- 研究行为:提出假设 → 建立实验 → 运行测试 → 分析 → 下一步,禁止以单次高收益宣告策略有效 +""" + +from __future__ import annotations + +import json +from collections.abc import Callable +from dataclasses import dataclass, field +from typing import Any + +TOOL_CALL_TAG = "__tool__" + + +@dataclass(frozen=True) +class Tool: + """受控工具元数据(LLM 可见的 JSON Schema + 调用实现)。""" + + name: str + description: str + parameters: dict[str, Any] + handler: Callable[[dict[str, Any]], str] = field(repr=False) + + def schema(self) -> dict[str, Any]: + return { + "name": self.name, + "description": self.description, + "parameters": self.parameters, + } + + def invoke(self, args: dict[str, Any]) -> str: + """执行工具;任何异常都转为可读错误串(不向上抛,避免中断整轮对话)。""" + try: + result = self.handler(args) + return result if isinstance(result, str) else json.dumps(result, ensure_ascii=False) + except Exception as exc: # noqa: BLE001 —— 工具异常反馈给 LLM 而非崩溃 + return json.dumps({"error": f"{type(exc).__name__}: {exc}"}, ensure_ascii=False) + + +def tools_schema(tools: list[Tool]) -> str: + return json.dumps([t.schema() for t in tools], ensure_ascii=False, indent=1) diff --git a/backend/app/agent/tools/__init__.py b/backend/app/agent/tools/__init__.py deleted file mode 100644 index c3fac56..0000000 --- a/backend/app/agent/tools/__init__.py +++ /dev/null @@ -1 +0,0 @@ -"""Agent 受控工具集:每个 Tool 是后端服务的只读 / 沙箱化入口。""" diff --git a/backend/app/agent/tools_impl.py b/backend/app/agent/tools_impl.py new file mode 100644 index 0000000..06fe554 --- /dev/null +++ b/backend/app/agent/tools_impl.py @@ -0,0 +1,248 @@ +"""受控工具集实现(AGENT.md §28):Agent 只能调用这里的白名单工具。 + +全部工具经 Job/Experiment 链路或只读查询执行: +- 不提供 shell / 任意代码执行 / 修改配置与凭证 / 删除数据 +- 任何研究都会产出 Experiment 归档(可复现) +""" + +from __future__ import annotations + +import json +from datetime import date + +from app.agent.tools import Tool +from app.application.services.job_executor import default_factories, submit_and_run +from app.domain.entities.research import ( + BacktestResult, + FactorTestReport, + ResearchSpec, +) + + +def _day(text: str) -> date: + return date.fromisoformat(text) + + +def _pick(mapping: dict, key: str, default=None): + val = mapping.get(key, default) + if isinstance(val, str): + val = val.strip() + if val == "": + return default + return val + + +def build_tools(factories: dict | None = None) -> list[Tool]: + facts = factories or default_factories() + session_factory = facts["session_factory"] + stock_repo_f = facts["stock_repo_factory"] + daily_repo_f = facts["daily_repo_factory"] + exp_repo_f = facts["experiment_repo_factory"] + + def search_stocks(args: dict) -> str: + q = str(_pick(args, "q", "") or "").upper() + with session_factory() as session: + stocks = stock_repo_f(session).list() + rows = [ + s for s in stocks if (not q) or q in s.symbol.upper() or q in (s.name or "").upper() + ][:15] + if not rows: + return "未找到匹配股票" + return "\n".join( + f"{s.symbol} {s.name} 行业={s.industry or '-'} 上市={s.list_date}" for s in rows + ) + + def get_market_data(args: dict) -> str: + symbol = str(_pick(args, "symbol", "")).upper() + start = _day(str(_pick(args, "start", "2024-01-01"))) + end = _day(str(_pick(args, "end", date.today().isoformat()))) + with session_factory() as session: + bars = daily_repo_f(session).get_range(symbol, start, end) + if not bars: + return f"{symbol} 在 {start}~{end} 无日线数据(可能未同步)" + head, tail = bars[0], bars[-1] + last = "\n".join(f"{b.trade_date} close={b.close}" for b in bars[-8:]) + change = float(tail.close) / float(head.close) - 1 if head.close and tail.close else None + return ( + f"{symbol} {start}~{end} 共 {len(bars)} 根日线;" + f"区间 {head.trade_date}→{tail.trade_date} 收盘 {head.close}→{tail.close}" + f"(涨跌 {change * 100:.2f}% 若数据完整);最近 8 根:\n{last}" + ) + + def _run_spec(spec: ResearchSpec, desc: str) -> str: + job = submit_and_run(spec, factories=facts) + if job.status != "success": + return f"{desc} 执行失败:{job.error}" + if spec.type == "backtest": + result = BacktestResult.model_validate_json(job.result_json or "{}") + s = result.summary + return ( + f"回测完成(Experiment {job.experiment_id},代码版本 {_code_version(job, exp_repo_f)})。" + f"总收益 {s.total_return_pct:.2f}%,年化 {s.annual_return_pct:.2f}%," + f"Sharpe {s.sharpe:.2f},最大回撤 {s.max_drawdown_pct:.2f}%," + f"交易 {s.total_trades} 笔,平均换手 {s.avg_turnover_pct:.1f}%。" + f"未建模约束 {len(result.unimplemented)} 项(成本/涨跌停近似见实验详情)。" + ) + report = FactorTestReport.model_validate_json(job.result_json or "{}") + qs = ", ".join(f"Q{q.quantile + 1}: {q.return_pct:.2f}%" for q in report.quantile_returns) + return ( + f"因子测试完成(Experiment {job.experiment_id})。IC {report.ic_mean:.4f}," + f"RankIC {report.rank_ic_mean:.4f},ICIR {report.icir:.2f},正收益占比 " + f"{report.positive_ratio_pct:.1f}%,样本 {report.sample_days} 日;分层未来收益 {qs}。" + f"注意:单因子测试不代表策略有效,需结合稳健性分析。" + ) + + def test_factor(args: dict) -> str: + name = str(_pick(args, "name", "")) + start = _day(str(_pick(args, "start", "2024-01-01"))) + end = _day(str(_pick(args, "end", "2024-12-31"))) + spec = ResearchSpec( + type="factor_test", + universe={"exclude_st": True, "min_listing_days": 0}, + factors=[{"name": name, "weight": 1.0}], + selection={"top_n": 10}, + rebalance="monthly", + period=(start, end), + ) + return _run_spec(spec, f"因子 {name} 测试") + + def run_backtest(args: dict) -> str: + factor_names = [f.strip() for f in str(_pick(args, "factors", "momentum_60")).split(",")] + top_n = int(_pick(args, "top_n", 5) or 5) + rebalance = str(_pick(args, "rebalance", "monthly")) + exclude_st = bool(_pick(args, "exclude_st", True)) + start = _day(str(_pick(args, "start", "2024-01-01"))) + end = _day(str(_pick(args, "end", "2024-12-31"))) + spec = ResearchSpec( + type="backtest", + universe={"exclude_st": exclude_st, "min_listing_days": 0}, + factors=[{"name": n, "weight": 1.0} for n in factor_names], + selection={"top_n": top_n}, + rebalance=rebalance, + period=(start, end), + ) + return _run_spec(spec, "回测") + + def get_experiment(args: dict) -> str: + exp_id = str(_pick(args, "experiment_id", "")).upper() + with session_factory() as session: + exp = exp_repo_f(session).get(exp_id) + if exp is None: + return f"Experiment {exp_id} 不存在(可用列表:GET /api/experiments)" + spec = json.loads(exp.spec_json) + return ( + f"Experiment {exp.id} [{exp.kind}] 因子={[f['name'] for f in spec.get('factors', [])]} " + f"区间={spec.get('period')} 调仓={spec.get('rebalance')};摘要:{exp.summary_text or '-'} " + f"代码版本={exp.code_version or '-'} 创建={exp.created_at}" + ) + + def compare_experiments(args: dict) -> str: + ids = [ + x.strip().upper() + for x in str(_pick(args, "experiment_ids", "")).split(",") + if x.strip() + ] + if not ids: + return "请提供 experiment_ids(逗号分隔)" + with session_factory() as session: + repo = exp_repo_f(session) + rows = [(i, repo.get(i)) for i in ids] + out = [] + for exp_id, exp in rows: + if exp is None: + out.append(f"{exp_id}: 不存在") + else: + spec = json.loads(exp.spec_json) + out.append( + f"{exp.id}: 因子={[f['name'] for f in spec.get('factors', [])]} " + f"区间={spec.get('period')} → {exp.summary_text or '-'}" + ) + return "\n".join(out) + + return [ + Tool( + "search_stocks", + "按代码或名称搜索股票,返回基础信息(只读)", + { + "type": "object", + "properties": {"q": {"type": "string", "description": "代码或名称关键字"}}, + }, + search_stocks, + ), + Tool( + "get_market_data", + "读取一只股票一段区间的日线行情摘要(只读,不复权)", + { + "type": "object", + "properties": { + "symbol": {"type": "string", "description": "如 600519.SH"}, + "start": {"type": "string", "description": "YYYY-MM-DD"}, + "end": {"type": "string", "description": "YYYY-MM-DD"}, + }, + "required": ["symbol"], + }, + get_market_data, + ), + Tool( + "test_factor", + "对单个因子做 IC/RankIC/分层测试并归档 Experiment", + { + "type": "object", + "properties": { + "name": { + "type": "string", + "description": "因子名(momentum_60 / volatility_20 等)", + }, + "start": {"type": "string"}, + "end": {"type": "string"}, + }, + "required": ["name"], + }, + test_factor, + ), + Tool( + "run_backtest", + "运行 TopK 低频回测并归档 Experiment(成本/涨跌停近似建模)", + { + "type": "object", + "properties": { + "factors": {"type": "string", "description": "逗号分隔的因子名"}, + "top_n": {"type": "integer"}, + "rebalance": {"type": "string", "enum": ["monthly", "weekly"]}, + "exclude_st": {"type": "boolean"}, + "start": {"type": "string"}, + "end": {"type": "string"}, + }, + }, + run_backtest, + ), + Tool( + "get_experiment", + "读取已归档实验的摘要", + { + "type": "object", + "properties": {"experiment_id": {"type": "string"}}, + "required": ["experiment_id"], + }, + get_experiment, + ), + Tool( + "compare_experiments", + "对比多个实验(因子/区间/收益摘要)", + { + "type": "object", + "properties": {"experiment_ids": {"type": "string"}}, + "required": ["experiment_ids"], + }, + compare_experiments, + ), + ] + + +def _code_version(job, exp_repo_f) -> str: + try: + with default_factories()["session_factory"]() as session: + exp = exp_repo_f(session).get(job.experiment_id or "") + return exp.code_version or "-" if exp else "-" + except Exception: # noqa: BLE001 + return "-" diff --git a/backend/app/api/agent.py b/backend/app/api/agent.py new file mode 100644 index 0000000..5287a40 --- /dev/null +++ b/backend/app/api/agent.py @@ -0,0 +1,40 @@ +"""AI Research Agent API(Phase 5)。 + +POST /api/agent/chat {message} → {reply, actions:[{tool,args,output}]} +未配置 LLM Key 时返回 400 引导(不会崩溃)。 +""" + +from __future__ import annotations + +from typing import Annotated + +from fastapi import APIRouter, Depends, HTTPException +from pydantic import BaseModel, Field + +from app.agent.llm import LLMClient, build_llm_from_settings +from app.agent.service import AgentService + +router = APIRouter(prefix="/agent", tags=["agent"]) + + +class AgentChatRequest(BaseModel): + message: str = Field(min_length=1, max_length=2000) + + +def _llm_or_raise() -> LLMClient: + llm = build_llm_from_settings() + if llm is None: + raise HTTPException( + status_code=400, + detail="未配置 LLM:请在根目录 .env 中设置 LLM_API_KEY(可选 LLM_BASE_URL / LLM_MODEL)," + "参考 .env.example 与 AGENT.md §33", + ) + return llm + + +@router.post("/chat", summary="与 AI 研究助手对话(受控工具)") +def agent_chat( + body: AgentChatRequest, + llm: Annotated[LLMClient, Depends(_llm_or_raise)], +) -> dict: + return AgentService(llm).chat(body.message) diff --git a/backend/app/api/jobs.py b/backend/app/api/jobs.py index ee898af..6cc4ac7 100644 --- a/backend/app/api/jobs.py +++ b/backend/app/api/jobs.py @@ -14,7 +14,7 @@ from datetime import datetime from fastapi import APIRouter, BackgroundTasks, HTTPException from fastapi.responses import StreamingResponse -from app.api.deps import DbSession, JobRepoDep, _engine_factory +from app.api.deps import DbSession, JobRepoDep from app.application.services.job_executor import execute_job, new_id from app.domain.entities.research import ( BacktestResult, @@ -24,27 +24,17 @@ from app.domain.entities.research import ( ResearchSpec, ) from app.infrastructure.persistence.sqlalchemy.repositories.jobs_impl import ( - SqlAlchemyExperimentRepository, SqlAlchemyJobRepository, ) -from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import ( - SqlAlchemyDailyBarRepository, - SqlAlchemyStockRepository, -) from app.infrastructure.persistence.sqlalchemy.session import SessionLocal router = APIRouter(prefix="/jobs", tags=["jobs"]) def _bg_factories(): - return { - "session_factory": SessionLocal, - "job_repo_factory": lambda s: SqlAlchemyJobRepository(s), - "experiment_repo_factory": lambda s: SqlAlchemyExperimentRepository(s), - "stock_repo_factory": lambda s: SqlAlchemyStockRepository(s), - "daily_repo_factory": lambda s: SqlAlchemyDailyBarRepository(s), - "engine": _engine_factory(), - } + from app.application.services.job_executor import default_factories + + return default_factories() def _decode_result(kind: str, result_json: str | None): diff --git a/backend/app/api/router.py b/backend/app/api/router.py index 390ecad..ef7c007 100644 --- a/backend/app/api/router.py +++ b/backend/app/api/router.py @@ -8,7 +8,7 @@ from __future__ import annotations from fastapi import APIRouter -from app.api import experiments, factors, health, jobs, research, stocks +from app.api import agent, experiments, factors, health, jobs, research, stocks api_router = APIRouter() api_router.include_router(health.router) @@ -17,3 +17,4 @@ api_router.include_router(factors.router) api_router.include_router(research.router) api_router.include_router(jobs.router) api_router.include_router(experiments.router) +api_router.include_router(agent.router) diff --git a/backend/app/application/services/job_executor.py b/backend/app/application/services/job_executor.py index c060446..613be81 100644 --- a/backend/app/application/services/job_executor.py +++ b/backend/app/application/services/job_executor.py @@ -18,6 +18,7 @@ from app.domain.entities.research import ( BacktestResult, ExperimentRecord, FactorTestReport, + JobRecord, JobStatus, ResearchSpec, ) @@ -144,3 +145,56 @@ def execute_job( session.commit() except Exception: # noqa: BLE001 pass + + +def default_factories() -> dict: + """后台执行所需的独立 Session / Repository / 引擎装配(跨请求生命周期)。""" + from app.infrastructure.persistence.sqlalchemy.repositories.jobs_impl import ( + SqlAlchemyExperimentRepository, + SqlAlchemyJobRepository, + ) + from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import ( + SqlAlchemyDailyBarRepository, + SqlAlchemyStockRepository, + ) + from app.infrastructure.persistence.sqlalchemy.session import SessionLocal + from app.quant.engine import LocalEngine + + return { + "session_factory": SessionLocal, + "job_repo_factory": lambda s: SqlAlchemyJobRepository(s), + "experiment_repo_factory": lambda s: SqlAlchemyExperimentRepository(s), + "stock_repo_factory": lambda s: SqlAlchemyStockRepository(s), + "daily_repo_factory": lambda s: SqlAlchemyDailyBarRepository(s), + "engine": LocalEngine(), + } + + +def submit_and_run(spec: ResearchSpec, *, factories: dict | None = None) -> JobRecord: + """同步创建并执行一个 Job(复用 Job 状态机与 Experiment 归档),返回终态 Job。""" + facts = factories or default_factories() + session_factory = facts["session_factory"] + job_repo = facts["job_repo_factory"] + job = JobRecord( + id=new_id("JOB"), + kind=spec.type, + spec_json=json.dumps(spec.model_dump(mode="json"), ensure_ascii=False), + status=JobStatus.QUEUED, + created_at=datetime.now(), + ) + with session_factory() as session: + job_repo(session).create(job) + session.commit() + execute_job( + job.id, + session_factory=session_factory, + job_repo_factory=job_repo, + experiment_repo_factory=facts["experiment_repo_factory"], + stock_repo_factory=facts["stock_repo_factory"], + daily_repo_factory=facts["daily_repo_factory"], + engine=facts["engine"], + ) + with session_factory() as session: + done = job_repo(session).get(job.id) + assert done is not None + return done diff --git a/backend/app/core/config.py b/backend/app/core/config.py index ec96c43..a0fb9e4 100644 --- a/backend/app/core/config.py +++ b/backend/app/core/config.py @@ -68,6 +68,9 @@ class Settings: data_source_primary: str data_source_fallback: str tushare_token: str + llm_api_key: str + llm_base_url: str | None + llm_model: str storage: dict[str, Path] config_path: Path = CONFIG_PATH env_path: Path = ENV_PATH @@ -102,6 +105,9 @@ def get_settings() -> Settings: secret_env = _deep(cfg, "app.secret_key_env") or "APP_SECRET_KEY" url_env = _deep(cfg, "database.url_env") or "DATABASE_URL" token_env = _deep(cfg, "data_source.tushare_token_env") or "TUSHARE_TOKEN" + llm_key_env = _deep(cfg, "agent.llm_key_env") or "LLM_API_KEY" + llm_url_env = _deep(cfg, "agent.llm_base_url_env") or "LLM_BASE_URL" + llm_model_env = _deep(cfg, "agent.llm_model_env") or "LLM_MODEL" default_url = "sqlite:///./data/quant.db" database_url = os.environ.get(url_env) or default_url @@ -122,5 +128,8 @@ def get_settings() -> Settings: data_source_primary=str(_deep(cfg, "data_source.primary") or "tushare"), data_source_fallback=str(_deep(cfg, "data_source.fallback") or "sina"), tushare_token=os.environ.get(token_env, ""), + llm_api_key=os.environ.get(llm_key_env, ""), + llm_base_url=os.environ.get(llm_url_env) or None, + llm_model=os.environ.get(llm_model_env) or "qwen-plus", storage=_resolve_storage_dirs(cfg), ) diff --git a/backend/pyproject.toml b/backend/pyproject.toml index d63b336..3808902 100644 --- a/backend/pyproject.toml +++ b/backend/pyproject.toml @@ -10,6 +10,7 @@ dependencies = [ "sqlalchemy>=2.0", "alembic>=1.13", "pyyaml>=6.0", + "httpx>=0.28.1", ] [project.optional-dependencies] diff --git a/backend/tests/test_agent.py b/backend/tests/test_agent.py new file mode 100644 index 0000000..cbbaec6 --- /dev/null +++ b/backend/tests/test_agent.py @@ -0,0 +1,135 @@ +"""Phase 5 Agent 测试:受控工具白名单 / 编排循环 / 异常输入(用 FakeLLM,不触网)。""" + +from __future__ import annotations + +import json + +import pytest +from app.agent.service import AgentService +from app.agent.tools_impl import build_tools +from app.infrastructure.persistence.sqlalchemy.base import Base +from sqlalchemy import create_engine +from sqlalchemy.orm import Session + + +class FakeLLM: + """按顺序返回预设决策的假 LLM。""" + + def __init__(self, decisions: list[dict]) -> None: + self._decisions = decisions + self.calls = 0 + + def chat(self, messages: list[dict]) -> str: + self.calls += 1 + return json.dumps(self._decisions.pop(0), ensure_ascii=False) + + +@pytest.fixture() +def tmp_factories(tmp_path) -> dict: + """注入 tmp SQLite 的默认工厂(隔离本机开发库)。""" + engine = create_engine(f"sqlite:///{tmp_path / 'agent.db'}", future=True) + Base.metadata.create_all(engine) + session_factory = lambda: Session(engine) # noqa: E731 + + from app.infrastructure.persistence.sqlalchemy.repositories.jobs_impl import ( + SqlAlchemyExperimentRepository, + SqlAlchemyJobRepository, + ) + from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import ( + SqlAlchemyDailyBarRepository, + SqlAlchemyStockRepository, + ) + from app.quant.engine import LocalEngine + + return { + "session_factory": session_factory, + "job_repo_factory": lambda s: SqlAlchemyJobRepository(s), + "experiment_repo_factory": lambda s: SqlAlchemyExperimentRepository(s), + "stock_repo_factory": lambda s: SqlAlchemyStockRepository(s), + "daily_repo_factory": lambda s: SqlAlchemyDailyBarRepository(s), + "engine": LocalEngine(), + } + + +class TestToolCatalog: + def test_whitelist_and_metadata(self) -> None: + tools = build_tools() + names = {t.name for t in tools} + assert { + "search_stocks", + "get_market_data", + "test_factor", + "run_backtest", + "get_experiment", + "compare_experiments", + } <= names + for t in tools: + assert t.description + assert t.parameters["type"] == "object" + + def test_no_arbitrary_shell_tool(self) -> None: + names = {t.name for t in build_tools()} + assert "shell" not in names + assert "execute" not in names + + +class TestAgentOrchestration: + def test_full_research_loop(self, tmp_factories) -> None: + llm = FakeLLM( + [ + { + "tool": "test_factor", + "args": {"name": "momentum_60", "start": "2024-03-01", "end": "2024-10-31"}, + }, + { + "tool": "run_backtest", + "args": { + "factors": "momentum_60", + "top_n": 3, + "start": "2024-03-01", + "end": "2024-10-31", + }, + }, + {"final": "结论:单次样本内回测较好,但需要样本外验证,暂不判定策略有效。"}, + ] + ) + svc = AgentService(llm, build_tools(factories=tmp_factories)) + out = svc.chat("请测试 momentum_60 因子并跑一次回测,再给结论。") + + assert [a["tool"] for a in out["actions"]] == ["test_factor", "run_backtest"] + # 工具输出真实来自引擎/归档链路(空库也返回结构化文本) + assert all("output" in a and a["output"] for a in out["actions"]) + assert "样本外" in out["reply"] + + def test_tool_error_fed_back_without_abort(self, tmp_factories) -> None: + llm = FakeLLM( + [ + {"tool": "search_stocks", "args": {"q": "茅台"}}, + {"final": "完成"}, + ] + ) + svc = AgentService(llm, build_tools(factories=tmp_factories)) + out = svc.chat("查一下茅台") + assert out["actions"][0]["tool"] == "search_stocks" + + def test_unknown_tool_is_rejected(self) -> None: + class RigidLLM: + def chat(self, messages): # noqa: ARG002 + return json.dumps({"tool": "delete_database", "args": {}}, ensure_ascii=False) + + svc = AgentService(RigidLLM(), build_tools()) + with pytest.raises(ValueError, match="工具不存在"): + svc.chat("恶意调用") + + def test_round_limit_forces_final(self, tmp_factories) -> None: + # LLM 一直要求调工具 → 轮次耗尽强制 final + class LoopLLM: + def chat(self, messages): # noqa: ARG002 + return json.dumps( + {"tool": "search_stocks", "args": {"q": "600519"}}, ensure_ascii=False + ) + + svc = AgentService(LoopLLM(), build_tools(factories=tmp_factories)) + out = svc.chat("一直搜索") + assert "actions" in out + assert len(out["actions"]) <= 6 diff --git a/backend/uv.lock b/backend/uv.lock index 97ecfe1..9014309 100644 --- a/backend/uv.lock +++ b/backend/uv.lock @@ -911,6 +911,7 @@ source = { virtual = "." } dependencies = [ { name = "alembic" }, { name = "fastapi" }, + { name = "httpx" }, { name = "pyyaml" }, { name = "sqlalchemy" }, { name = "uvicorn", extra = ["standard"] }, @@ -932,6 +933,7 @@ dev = [ requires-dist = [ { name = "alembic", specifier = ">=1.13" }, { name = "fastapi", specifier = ">=0.115" }, + { name = "httpx", specifier = ">=0.28.1" }, { name = "pyyaml", specifier = ">=6.0" }, { name = "sqlalchemy", specifier = ">=2.0" }, { name = "tushare", marker = "extra == 'datasource-tushare'", specifier = ">=1.4" }, diff --git a/docs/ROADMAP.md b/docs/ROADMAP.md index 7c85e8e..317efbe 100644 --- a/docs/ROADMAP.md +++ b/docs/ROADMAP.md @@ -8,14 +8,14 @@ ## 0. 总览与里程碑 -| 里程碑 | 内容 | 验收口径 | +| 里程碑 | 内容 | 状态(2026-09 已实施) | |---|---|---| -| M0 ✅ | 工程初始化:uv + FastAPI 骨架、分层包、SQLAlchemy/Alembic、CI 前检查(pytest/ruff) | `/api/health` 可用、测试通过、已推送远端 | -| M1 | Phase 1 数据层:Tushare 拉取 → 标准化 → SQLite/Parquet | 命令行可全量/增量同步,来源可追溯 | -| M2 | Phase 2 Qlib:Qlib Dataset 构建 + 因子 + LightGBM + 回测(Adapter 封装) | Research Specification 能驱动一次完整回测,产出标准结果 | -| M3 | Phase 3 Web:股票池 / 因子 / 选股 / 回测 / 结果可视化 | 浏览器完成「选池→算因子→回测→看图」闭环 | -| M4 | Phase 4 Experiment:研究全量可复现归档 + Job/SSE 异步化 | 任一历史实验可一键复跑 | -| M5 | Phase 5 AI Agent:自然语言 → Research Plan → 受控 Tool → Experiment | Agent 能独立完成一次「因子假设→测试→结论」并留档 | +| M0 ✅ | 工程初始化 | commit 2a52ee5 | +| M1 ✅ | Phase 1 数据层(Tushare→SQLite/Parquet、Failover 审计、防未来函数) | commit 2da2342 · 38 tests | +| M2 ✅ | Phase 2 研究引擎(Spec→因子→评估→低频回测→标准结果;Qlib 桥接占位见 §2 备注) | commit e9f59d3 · 60 tests | +| M3 ✅ | Phase 3 Web(业务 API + Next.js 前端闭环) | commit 92627f5 · 68 tests | +| M4 ✅ | Phase 4 Experiment 归档 + Job/SSE 异步(一键复跑) | commit 0ea229d · 73 tests | +| M5 ✅ | Phase 5 AI Agent(受控 Tool 白名单 + LLM 编排;Key 配置见 .env) | commit(本轮)· 79 tests | 每阶段结束时同步更新:README / AGENT.md 相关清单 / 文档;**禁止跨阶段提前堆量**(AGENT.md §38)。