diff --git a/README.md b/README.md index 460bb7a..3587d3e 100644 --- a/README.md +++ b/README.md @@ -33,6 +33,9 @@ SearchAgent 是一个本地优先的 AI 深度研究工作台。它用类似 Cod - OpenAI、Anthropic、DeepSeek、Xiaomi MiMo、Alibaba Qwen 和自定义服务预设 - 本地 stdio 与远程 HTTP MCP 服务配置、测试和编辑 - API Key 与 MCP 环境变量本地加密存储 +- 规划、检索、核验、写作四类 Agent 任务编排,检索任务最多三路协作 +- 本地 Trace、工具调用审计、Agent 级工具策略和运行中授权确认 +- 本地评测数据集、确定性指标及可选独立 LLM 裁判配置 - 中英文界面、浏览器语言检测和语言偏好持久化 ## 技术栈 @@ -123,6 +126,8 @@ SearchAgent 会识别名为 `bocha` 的配置,并使用保存的密钥调用 └── reports/ # 生成的报告导出文件 ``` +Trace、Agent 任务、工具策略、授权记录和评测结果也保存在 `searchagent.db`。权限策略按 Agent 角色和工具名匹配,未知工具默认暂停等待本次任务确认;批准不会写入持久化允许规则。首期是应用层隔离,不提供 MCP 进程的操作系统级沙箱。 + 可以用 `SEARCHAGENT_HOME` 指定其他目录: ```powershell diff --git a/backend/app/api/observability.py b/backend/app/api/observability.py new file mode 100644 index 0000000..18b6ac1 --- /dev/null +++ b/backend/app/api/observability.py @@ -0,0 +1,152 @@ +import asyncio +import time + +from fastapi import APIRouter, Depends, HTTPException +from sqlalchemy.orm import Session + +from app.api.deps import get_db +from app.db.models import Conversation, EvaluationCase, EvaluationDataset, EvaluationRun, EvaluationScore, LLMProfile, ResearchProject +from app.engine.runner import resume_research, start_research +from app.llm.factory import build_chat_model_from_profile +from app.schemas.observability import EvaluationCasePayload, EvaluationDatasetPayload, EvaluationRunPayload, ToolApprovalPayload, ToolPolicyPayload +from app.services.tracing_service import list_trace +from app.tools import policy + +router = APIRouter(tags=["observability"]) + + +@router.post("/research/{thread_id}/tool-approval") +async def approve_tool_call(thread_id: str, payload: ToolApprovalPayload, session: Session = Depends(get_db)): + from app.engine.runner import resume_research + + return await resume_research( + session, + thread_id=thread_id, + profile_id=payload.profile_id, + decision={"kind": "tool_approval", "approved": payload.approved, "agent_role": payload.agent_role, "tool_name": payload.tool_name, "args_fingerprint": payload.args_fingerprint}, + ) + + +@router.get("/tool-policies") +def list_tool_policies(session: Session = Depends(get_db)): + return policy.list_policies(session) + + +@router.post("/tool-policies") +def create_tool_policy(payload: ToolPolicyPayload, session: Session = Depends(get_db)): + return policy.create_policy(session, **payload.model_dump()) + + +@router.put("/tool-policies/{policy_id}") +def update_tool_policy(policy_id: int, payload: ToolPolicyPayload, session: Session = Depends(get_db)): + result = policy.update_policy(session, policy_id, **payload.model_dump()) + if result is None: + raise HTTPException(status_code=404, detail="Tool policy not found") + return result + + +@router.delete("/tool-policies/{policy_id}", status_code=204) +def delete_tool_policy(policy_id: int, session: Session = Depends(get_db)): + if not policy.delete_policy(session, policy_id): + raise HTTPException(status_code=404, detail="Tool policy not found") + + +@router.get("/research/projects/{project_id}/traces") +def get_project_traces(project_id: int, session: Session = Depends(get_db)): + if session.get(ResearchProject, project_id) is None: + raise HTTPException(status_code=404, detail="Research project not found") + return list_trace(session, project_id) + + +@router.post("/evaluation/datasets") +def create_dataset(payload: EvaluationDatasetPayload, session: Session = Depends(get_db)): + if payload.judge_profile_id and session.get(LLMProfile, payload.judge_profile_id) is None: + raise HTTPException(status_code=404, detail="Judge profile not found") + item = EvaluationDataset(**payload.model_dump()) + session.add(item) + session.commit() + return _dataset(item, session) + + +@router.get("/evaluation/datasets") +def list_datasets(session: Session = Depends(get_db)): + return [_dataset(item, session) for item in session.query(EvaluationDataset).order_by(EvaluationDataset.id.desc())] + + +@router.post("/evaluation/datasets/{dataset_id}/cases") +def create_case(dataset_id: int, payload: EvaluationCasePayload, session: Session = Depends(get_db)): + if session.get(EvaluationDataset, dataset_id) is None: + raise HTTPException(status_code=404, detail="Evaluation dataset not found") + item = EvaluationCase(dataset_id=dataset_id, input_text=payload.input_text, expected_json=payload.expected) + session.add(item) + session.commit() + return {"id": item.id, "dataset_id": item.dataset_id, "input_text": item.input_text, "expected": item.expected_json} + + +@router.post("/evaluation/datasets/{dataset_id}/runs") +def run_dataset(dataset_id: int, payload: EvaluationRunPayload, session: Session = Depends(get_db)): + dataset = session.get(EvaluationDataset, dataset_id) + if dataset is None: + raise HTTPException(status_code=404, detail="Evaluation dataset not found") + if session.get(LLMProfile, payload.profile_id) is None: + raise HTTPException(status_code=404, detail="Research profile not found") + judge_id = payload.judge_profile_id if payload.judge_profile_id is not None else dataset.judge_profile_id + if judge_id is not None and session.get(LLMProfile, judge_id) is None: + raise HTTPException(status_code=404, detail="Judge profile not found") + started = time.perf_counter() + cases = session.query(EvaluationCase).filter(EvaluationCase.dataset_id == dataset_id).all() + run = EvaluationRun(dataset_id=dataset_id, profile_id=payload.profile_id, judge_profile_id=judge_id, status="completed") + session.add(run) + session.flush() + for case in cases: + metrics = {"input_nonempty": bool(case.input_text.strip()), "expected_present": case.expected_json is not None, "source_count": 0, "citation_complete": False, "permission_violations": 0, "interrupted": False} + status = "failed" + if metrics["input_nonempty"]: + conversation = Conversation(title=f"Evaluation {dataset.name} #{case.id}") + session.add(conversation) + session.commit() + try: + result = asyncio.run(start_research(session, conversation_id=conversation.id, profile_id=payload.profile_id, user_message=case.input_text)) + metrics["interrupted"] = bool(result.get("interrupted")) + if result.get("interrupted") and (result.get("interrupt_payload") or {}).get("kind") in {"plan_ready", "plan_approval"}: + version = (result.get("state") or {}).get("plan_version") + resumed = asyncio.run(resume_research(session, thread_id=result["thread_id"], profile_id=payload.profile_id, decision={"kind": "execute_plan", "plan_version": version})) + result = resumed + state = result.get("state") or {} + findings = state.get("verified_findings") or state.get("findings") or [] + metrics["source_count"] = len(findings) + metrics["citation_complete"] = bool(state.get("report_md") and "[^" in str(state.get("report_md"))) + status = "passed" if state.get("report_id") else "blocked" + if judge_id is not None and state.get("report_md"): + try: + judge = build_chat_model_from_profile(session, judge_id) + verdict = judge.invoke( + "Score this research report from 0 to 1 for whether it addresses the expected criteria. " + "Return only a number.\n" + f"Expected: {case.expected_json or {}}\nReport: {str(state['report_md'])[:6000]}" + ) + metrics["judge_score"] = float(str(getattr(verdict, "content", verdict)).strip()) + except Exception as exc: + metrics["judge_error"] = str(exc)[:300] + except Exception as exc: + metrics["error"] = str(exc)[:500] + status = "failed" + session.add(EvaluationScore(run_id=run.id, case_id=case.id, status=status, metrics_json=metrics)) + run.metrics_json = {"case_count": len(cases), "completed": len(cases), "latency_ms": round((time.perf_counter() - started) * 1000, 2), "judge_profile_id": judge_id} + session.commit() + return _run(run, session) + + +@router.get("/evaluation/runs") +def list_runs(session: Session = Depends(get_db)): + return [_run(item, session) for item in session.query(EvaluationRun).order_by(EvaluationRun.id.desc())] + + +def _dataset(item: EvaluationDataset, session: Session) -> dict: + cases = session.query(EvaluationCase).filter(EvaluationCase.dataset_id == item.id).all() + return {"id": item.id, "name": item.name, "version": item.version, "description": item.description, "judge_profile_id": item.judge_profile_id, "case_count": len(cases), "created_at": item.created_at.isoformat()} + + +def _run(item: EvaluationRun, session: Session) -> dict: + scores = session.query(EvaluationScore).filter(EvaluationScore.run_id == item.id).all() + return {"id": item.id, "dataset_id": item.dataset_id, "profile_id": item.profile_id, "judge_profile_id": item.judge_profile_id, "status": item.status, "metrics": item.metrics_json, "scores": [{"case_id": score.case_id, "status": score.status, "metrics": score.metrics_json} for score in scores], "created_at": item.created_at.isoformat()} diff --git a/backend/app/api/research.py b/backend/app/api/research.py index a033418..1ad4ca6 100644 --- a/backend/app/api/research.py +++ b/backend/app/api/research.py @@ -2,7 +2,7 @@ from sqlalchemy.orm import Session from app.api.deps import get_db -from app.db.models import Conversation, LLMProfile, ResearchProject, Source, Step +from app.db.models import AgentTask, Conversation, LLMProfile, ResearchProject, Source, Step, ToolApproval, ToolCall from app.engine import runner from app.schemas.research import ( ResearchResumeRequest, @@ -116,6 +116,22 @@ def project_execution_detail( } for report in sorted(project.reports, key=lambda item: (item.version, item.id)) ], + "agent_tasks": [ + {"id": task.id, "role": task.role, "title": task.title, "status": task.status, + "input": task.input_json, "output": task.output_json} + for task in session.query(AgentTask).filter(AgentTask.project_id == project.id).order_by(AgentTask.id) + ], + "tool_calls": [ + {"id": call.id, "task_id": call.task_id, "agent_role": call.agent_role, + "tool_name": call.tool_name, "status": call.status, "result_summary": call.result_summary, + "error": call.error} + for call in session.query(ToolCall).filter(ToolCall.project_id == project.id).order_by(ToolCall.id) + ], + "approvals": [ + {"id": approval.id, "agent_role": approval.agent_role, "tool_name": approval.tool_name, + "decision": approval.decision, "args_fingerprint": approval.args_fingerprint} + for approval in session.query(ToolApproval).filter(ToolApproval.project_id == project.id).order_by(ToolApproval.id) + ], } diff --git a/backend/app/api/router.py b/backend/app/api/router.py index 066a885..b662bb2 100644 --- a/backend/app/api/router.py +++ b/backend/app/api/router.py @@ -1,6 +1,6 @@ from fastapi import APIRouter -from app.api import conversations, events, health, mcp, reports, research, settings +from app.api import conversations, events, health, mcp, observability, reports, research, settings api_router = APIRouter() api_router.include_router(health.router) @@ -10,3 +10,4 @@ api_router.include_router(mcp.router) api_router.include_router(reports.router) api_router.include_router(events.router) +api_router.include_router(observability.router) diff --git a/backend/app/db/models.py b/backend/app/db/models.py index 7139807..97ce273 100644 --- a/backend/app/db/models.py +++ b/backend/app/db/models.py @@ -189,3 +189,134 @@ class AppPreference(Base): key: Mapped[str] = mapped_column(primary_key=True) value_json: Mapped[Optional[dict]] = mapped_column(JSON, default=None) + + +class AgentTask(Base): + __tablename__ = "agent_tasks" + + id: Mapped[int] = mapped_column(primary_key=True) + project_id: Mapped[int] = mapped_column(ForeignKey("research_projects.id", ondelete="CASCADE")) + trace_id: Mapped[str] = mapped_column(index=True) + role: Mapped[str] = mapped_column() + title: Mapped[str] = mapped_column() + status: Mapped[str] = mapped_column(default="pending") + input_json: Mapped[Optional[dict]] = mapped_column(JSON, default=None) + output_json: Mapped[Optional[dict]] = mapped_column(JSON, default=None) + created_at: Mapped[dt.datetime] = mapped_column(default=_now) + started_at: Mapped[Optional[dt.datetime]] = mapped_column(default=None) + completed_at: Mapped[Optional[dt.datetime]] = mapped_column(default=None) + + +class ToolPolicy(Base): + __tablename__ = "tool_policies" + + id: Mapped[int] = mapped_column(primary_key=True) + version: Mapped[int] = mapped_column(default=1) + agent_role: Mapped[str] = mapped_column(index=True) + tool_name: Mapped[str] = mapped_column(index=True) + allowed_domains_json: Mapped[Optional[list]] = mapped_column(JSON, default=None) + require_approval: Mapped[bool] = mapped_column(Boolean, default=False) + enabled: Mapped[bool] = mapped_column(Boolean, default=True) + created_at: Mapped[dt.datetime] = mapped_column(default=_now) + + +class ToolApproval(Base): + __tablename__ = "tool_approvals" + + id: Mapped[int] = mapped_column(primary_key=True) + project_id: Mapped[int] = mapped_column(ForeignKey("research_projects.id", ondelete="CASCADE")) + trace_id: Mapped[str] = mapped_column(index=True) + agent_role: Mapped[str] = mapped_column() + tool_name: Mapped[str] = mapped_column() + args_fingerprint: Mapped[str] = mapped_column(index=True) + decision: Mapped[str] = mapped_column(default="approved") + created_at: Mapped[dt.datetime] = mapped_column(default=_now) + + +class ToolCall(Base): + __tablename__ = "tool_calls" + + id: Mapped[int] = mapped_column(primary_key=True) + project_id: Mapped[int] = mapped_column(ForeignKey("research_projects.id", ondelete="CASCADE")) + trace_id: Mapped[str] = mapped_column(index=True) + task_id: Mapped[Optional[int]] = mapped_column(ForeignKey("agent_tasks.id", ondelete="SET NULL"), default=None) + agent_role: Mapped[str] = mapped_column() + tool_name: Mapped[str] = mapped_column() + args_json: Mapped[Optional[dict]] = mapped_column(JSON, default=None) + status: Mapped[str] = mapped_column(default="started") + result_summary: Mapped[Optional[str]] = mapped_column(Text, default=None) + error: Mapped[Optional[str]] = mapped_column(Text, default=None) + created_at: Mapped[dt.datetime] = mapped_column(default=_now) + completed_at: Mapped[Optional[dt.datetime]] = mapped_column(default=None) + + +class TraceRun(Base): + __tablename__ = "trace_runs" + + id: Mapped[str] = mapped_column(primary_key=True) + project_id: Mapped[int] = mapped_column(ForeignKey("research_projects.id", ondelete="CASCADE"), index=True) + root_name: Mapped[str] = mapped_column(default="research") + status: Mapped[str] = mapped_column(default="running") + created_at: Mapped[dt.datetime] = mapped_column(default=_now) + completed_at: Mapped[Optional[dt.datetime]] = mapped_column(default=None) + + +class TraceSpan(Base): + __tablename__ = "trace_spans" + + id: Mapped[str] = mapped_column(primary_key=True) + trace_id: Mapped[str] = mapped_column(ForeignKey("trace_runs.id", ondelete="CASCADE"), index=True) + parent_id: Mapped[Optional[str]] = mapped_column(ForeignKey("trace_spans.id", ondelete="SET NULL"), default=None) + name: Mapped[str] = mapped_column() + kind: Mapped[str] = mapped_column(default="internal") + status: Mapped[str] = mapped_column(default="running") + attributes_json: Mapped[Optional[dict]] = mapped_column(JSON, default=None) + input_summary: Mapped[Optional[str]] = mapped_column(Text, default=None) + output_summary: Mapped[Optional[str]] = mapped_column(Text, default=None) + error: Mapped[Optional[str]] = mapped_column(Text, default=None) + started_at: Mapped[dt.datetime] = mapped_column(default=_now) + ended_at: Mapped[Optional[dt.datetime]] = mapped_column(default=None) + + +class EvaluationDataset(Base): + __tablename__ = "evaluation_datasets" + + id: Mapped[int] = mapped_column(primary_key=True) + name: Mapped[str] = mapped_column(unique=True) + version: Mapped[int] = mapped_column(default=1) + description: Mapped[str] = mapped_column(Text, default="") + judge_profile_id: Mapped[Optional[int]] = mapped_column(ForeignKey("llm_profiles.id", ondelete="SET NULL"), default=None) + created_at: Mapped[dt.datetime] = mapped_column(default=_now) + + +class EvaluationCase(Base): + __tablename__ = "evaluation_cases" + + id: Mapped[int] = mapped_column(primary_key=True) + dataset_id: Mapped[int] = mapped_column(ForeignKey("evaluation_datasets.id", ondelete="CASCADE")) + input_text: Mapped[str] = mapped_column(Text) + expected_json: Mapped[Optional[dict]] = mapped_column(JSON, default=None) + created_at: Mapped[dt.datetime] = mapped_column(default=_now) + + +class EvaluationRun(Base): + __tablename__ = "evaluation_runs" + + id: Mapped[int] = mapped_column(primary_key=True) + dataset_id: Mapped[int] = mapped_column(ForeignKey("evaluation_datasets.id", ondelete="CASCADE")) + profile_id: Mapped[int] = mapped_column(ForeignKey("llm_profiles.id", ondelete="RESTRICT")) + judge_profile_id: Mapped[Optional[int]] = mapped_column(ForeignKey("llm_profiles.id", ondelete="SET NULL"), default=None) + status: Mapped[str] = mapped_column(default="completed") + metrics_json: Mapped[Optional[dict]] = mapped_column(JSON, default=None) + created_at: Mapped[dt.datetime] = mapped_column(default=_now) + + +class EvaluationScore(Base): + __tablename__ = "evaluation_scores" + + id: Mapped[int] = mapped_column(primary_key=True) + run_id: Mapped[int] = mapped_column(ForeignKey("evaluation_runs.id", ondelete="CASCADE")) + case_id: Mapped[int] = mapped_column(ForeignKey("evaluation_cases.id", ondelete="CASCADE")) + status: Mapped[str] = mapped_column() + metrics_json: Mapped[dict] = mapped_column(JSON, default=dict) + created_at: Mapped[dt.datetime] = mapped_column(default=_now) diff --git a/backend/app/engine/nodes.py b/backend/app/engine/nodes.py index e585233..c1cda57 100644 --- a/backend/app/engine/nodes.py +++ b/backend/app/engine/nodes.py @@ -3,16 +3,20 @@ import asyncio import inspect import threading +import datetime as dt +from concurrent.futures import ThreadPoolExecutor from typing import Any from langgraph.types import interrupt from app.core.events import get_event_bus -from app.db.models import Source, Step +from app.db.models import AgentTask, ResearchProject, Source, Step, TraceRun from app.engine.context import EngineContext from app.engine.state import ResearchState from app.services.settings_service import get_preference from app.services.report_markdown import normalize_report_markdown +from app.services.tracing_service import finish_span, start_span +from app.tools.gateway import execute_tool def _get_llm(ctx: EngineContext) -> Any: @@ -315,7 +319,7 @@ def _report_references(findings: list[dict], *, chinese: bool = True) -> list[st def _build_report_markdown(state: ResearchState, analysis_md: str | None = None) -> str: - findings = state.get("findings", []) + findings = state.get("verified_findings") or state.get("findings", []) chinese = _is_chinese(state) lines = [ "# 研究报告" if chinese else "# Research Report", @@ -365,7 +369,7 @@ def _normalize_citation(match: re.Match[str]) -> str: def _synthesize_report(ctx: EngineContext, state: ResearchState) -> str | None: - findings = state.get("findings", []) + findings = state.get("verified_findings") or state.get("findings", []) evidence = [ { "source_id": idx, @@ -427,6 +431,13 @@ def publish_progress(event_type: str, state: ResearchState, **details: Any) -> N def plan_conversation(state: ResearchState) -> dict: from app.engine.persistence import persist_plan_version + planner_task = None + trace_id = state.get("trace_id") + if isinstance(trace_id, str) and isinstance(state.get("project_id"), int) and ctx.session.get(ResearchProject, state["project_id"]) is not None and ctx.session.get(TraceRun, trace_id) is not None: + planner_task = AgentTask(project_id=state["project_id"], trace_id=trace_id, role="planner", title="Clarify objective and build plan", status="running", input_json={"message": _latest_user_message(state.get("messages", []))}) + ctx.session.add(planner_task) + ctx.session.commit() + latest_user = _latest_user_message(state.get("messages", [])) current_objective = str(state.get("objective") or "").strip() previous_plan = state.get("plan") or {} @@ -474,6 +485,11 @@ def plan_conversation(state: ResearchState) -> dict: objective = str(parsed.get("objective") or current_objective).strip() questions = _normalize_planning_questions(parsed.get("questions")) publish_progress("research.planning_message", state, message=message) + if planner_task is not None: + planner_task.status = "completed" + planner_task.completed_at = dt.datetime.now(dt.timezone.utc).replace(tzinfo=None) + planner_task.output_json = {"ready": False, "question_count": len(questions)} + ctx.session.commit() return { "objective": objective, "planner_message": message, @@ -495,6 +511,11 @@ def plan_conversation(state: ResearchState) -> dict: else "Please clarify the objective, scope, time horizon, or expected output." ) publish_progress("research.planning_message", state, message=message) + if planner_task is not None: + planner_task.status = "completed" + planner_task.completed_at = dt.datetime.now(dt.timezone.utc).replace(tzinfo=None) + planner_task.output_json = {"ready": False, "question_count": 0} + ctx.session.commit() return { "planner_message": message, "planner_questions": [], @@ -539,6 +560,11 @@ def plan_conversation(state: ResearchState) -> dict: plan_version=version, step_count=len(plan["steps"]), ) + if planner_task is not None: + planner_task.status = "completed" + planner_task.completed_at = dt.datetime.now(dt.timezone.utc).replace(tzinfo=None) + planner_task.output_json = {"ready": True, "step_count": len(plan["steps"]), "plan_version": version} + ctx.session.commit() return { "objective": objective, "plan": plan, @@ -680,6 +706,18 @@ def execute_research(state: ResearchState) -> dict: step_title=steps[0].get("title"), ) max_sources = _get_max_sources(ctx) + trace_id = state.get("trace_id") + agent_tasks: list[dict] = [] + if not isinstance(trace_id, str) or ctx.session.get(TraceRun, trace_id) is None: + trace_id = None + if trace_id: + for index, step in enumerate(steps[:3], start=1): + task = AgentTask(project_id=state["project_id"], trace_id=trace_id, role="retriever", title=str(step.get("title") or f"Retrieval task {index}"), status="running", input_json={"step": step}) + ctx.session.add(task) + ctx.session.flush() + agent_tasks.append({"id": task.id, "role": task.role, "title": task.title, "status": task.status}) + ctx.session.commit() + coordinator_span = start_span(ctx.session, trace_id, "agent:coordinator", kind="agent", attributes={"role": "coordinator", "parallel_limit": 3}, input_value=state.get("objective")) if trace_id else None tool_result = _run_async(registry.get_research_tools(ctx.session)) tools = tool_result.get("tools", []) tools_by_name = {_tool_name(tool): tool for tool in tools} @@ -697,22 +735,38 @@ def execute_research(state: ResearchState) -> dict: response = bound_llm.invoke(prompt) findings = list(state.get("findings", [])) - for call in _response_tool_calls(response): - if len(findings) >= max_sources: - break - + tool_calls = [] + for call in _response_tool_calls(response)[:3]: name = _tool_call_name(call) tool = tools_by_name.get(name) - if tool is None: - continue - - args = _tool_call_args(call) - publish_progress( - "research.tool_started", - state, - tool_name=name, + if tool is not None: + tool_calls.append((name, tool, _tool_call_args(call))) + + parallel_results = None + if trace_id and len(tool_calls) > 1: + # Interrupts stay on the graph thread; once every call is allowed, actual I/O runs in up to three workers. + from app.db import session as db_session + from app.tools.policy import decide + for name, tool, args in tool_calls: + decision = decide(ctx.session, project_id=state["project_id"], trace_id=trace_id, agent_role="retriever", tool_name=name, args=args) + if decision.action != "allow": + execute_tool(ctx.session, project_id=state["project_id"], trace_id=trace_id, agent_role="retriever", tool_name=name, args=args, tool=tool, invoke=_invoke_tool, publish=lambda event_type, **details: publish_progress(event_type, state, **details)) + def _parallel_call(item): + name, tool, args = item + with db_session.SessionLocal() as worker_session: + return execute_tool(worker_session, project_id=state["project_id"], trace_id=trace_id, agent_role="retriever", tool_name=name, args=args, tool=tool, invoke=_invoke_tool, publish=lambda event_type, **details: publish_progress(event_type, state, **details)) + with ThreadPoolExecutor(max_workers=min(3, len(tool_calls))) as executor: + parallel_results = list(executor.map(_parallel_call, tool_calls)) + + for index, (name, tool, args) in enumerate(tool_calls): + if len(findings) >= max_sources: + break + task_id = agent_tasks[0]["id"] if agent_tasks else None + publish_progress("research.tool_started", state, tool_name=name, agent_role="retriever") + result = parallel_results[index] if parallel_results is not None else ( + execute_tool(ctx.session, project_id=state["project_id"], trace_id=trace_id, agent_role="retriever", tool_name=name, args=args, tool=tool, invoke=_invoke_tool, publish=lambda event_type, **details: publish_progress(event_type, state, **details)) + if trace_id else _invoke_tool(tool, args) ) - result = _invoke_tool(tool, args) collected_before = len(findings) for source in _iter_source_items(result, tool_name=name, args=args): if len(findings) >= max_sources: @@ -740,11 +794,14 @@ def execute_research(state: ResearchState) -> dict: "research.tool_completed", state, tool_name=name, + agent_role="retriever", source_count=len(findings) - collected_before, ) ctx.session.commit() if not findings: + if coordinator_span: + finish_span(ctx.session, coordinator_span, status="error", error="No usable sources") raise RuntimeError( "No usable sources were collected. Configure a search-capable MCP server " "or verify that the requested public sources are reachable." @@ -773,16 +830,38 @@ def execute_research(state: ResearchState) -> dict: step_seq=step.get("seq"), step_title=step.get("title"), ) + for task in ctx.session.query(AgentTask).filter(AgentTask.project_id == state["project_id"], AgentTask.trace_id == trace_id).all(): + task.status = "completed" + task.completed_at = dt.datetime.now(dt.timezone.utc).replace(tzinfo=None) + task.output_json = {"source_count": len(findings)} + ctx.session.commit() + if coordinator_span: + finish_span(ctx.session, coordinator_span, output_value={"source_count": len(findings)}) return {"findings": findings, "steps": steps} def aggregate_evidence(state: ResearchState) -> dict: findings = _dedupe_findings(state.get("findings", [])) - steps = _mark_acceptance_criteria(state.get("steps", []), findings) - return {"findings": findings, "steps": steps} + trace_id = state.get("trace_id") + verifier_span = start_span(ctx.session, trace_id, "agent:verifier", kind="agent", attributes={"role": "verifier"}, input_value={"finding_count": len(findings)}) if isinstance(trace_id, str) and ctx.session.get(TraceRun, trace_id) is not None else None + verified = [] + for finding in findings: + if finding.get("url"): + verified.append({**finding, "verified": True}) + steps = _mark_acceptance_criteria(state.get("steps", []), verified) + if verifier_span: + finish_span(ctx.session, verifier_span, output_value={"verified_count": len(verified)}) + return {"findings": findings, "verified_findings": verified, "steps": steps} def write_report(state: ResearchState) -> dict: from app.engine.persistence import persist_report, sync_plan_and_steps + writer_task = None + trace_id = state.get("trace_id") + if isinstance(trace_id, str) and isinstance(state.get("project_id"), int) and ctx.session.get(ResearchProject, state["project_id"]) is not None and ctx.session.get(TraceRun, trace_id) is not None: + writer_task = AgentTask(project_id=state["project_id"], trace_id=trace_id, role="writer", title="Synthesize verified evidence", status="running", input_json={"verified_count": len(state.get("verified_findings") or state.get("findings", []))}) + ctx.session.add(writer_task) + ctx.session.commit() + try: analysis_md = _synthesize_report(ctx, state) except Exception: @@ -790,6 +869,11 @@ def write_report(state: ResearchState) -> dict: report_md = _build_report_markdown(state, analysis_md=analysis_md) sync_plan_and_steps(ctx.session, state) report = persist_report(ctx.session, project_id=state["project_id"], content_md=report_md) + if writer_task is not None: + writer_task.status = "completed" + writer_task.completed_at = dt.datetime.now(dt.timezone.utc).replace(tzinfo=None) + writer_task.output_json = {"report_id": report.id, "citation_count": report_md.count("[^")} + ctx.session.commit() publish_progress("research.report_ready", state, report_id=report.id) return {"report_md": report_md, "report_id": report.id} diff --git a/backend/app/engine/runner.py b/backend/app/engine/runner.py index 340add3..06c28a9 100644 --- a/backend/app/engine/runner.py +++ b/backend/app/engine/runner.py @@ -12,6 +12,8 @@ from app.engine.context import EngineContext from app.engine.graph import compile_research_graph from app.services.report_markdown import normalize_report_markdown +from app.services.tracing_service import finish_trace, start_trace +from app.tools.policy import ensure_default_policies def _config(thread_id: str) -> dict: @@ -58,6 +60,7 @@ def _update_project_status(session: Session, state: dict, status: str) -> None: project.conversation.status = { "planning": "planning", "awaiting_execution": "awaiting_execution", + "awaiting_approval": "awaiting_approval", "executing": "executing", "revising_report": "revising_report", "done": "completed", @@ -157,12 +160,16 @@ def _record_message_once( def _waiting_status(payload: Optional[dict]) -> str: + if payload and payload.get("kind") == "tool_approval": + return "awaiting_approval" if payload and payload.get("kind") in {"planning_input", "clarification"}: return "planning" return "awaiting_execution" def _waiting_event(payload: dict) -> str: + if payload.get("kind") == "tool_approval": + return "research.tool_approval_required" if payload.get("kind") in {"planning_input", "clarification"}: return "research.planning_message" return "research.plan_ready" @@ -171,11 +178,12 @@ def _waiting_event(payload: dict) -> str: def _persist_waiting_message( session: Session, *, thread_id: str, state: dict, profile_id: int, payload: Optional[dict] ) -> None: - if not payload or payload.get("kind") not in { + if not payload or payload.get("kind") == "tool_approval" or payload.get("kind") not in { "planning_input", "plan_ready", "clarification", "plan_approval", + "tool_approval", }: return question = str(payload.get("message") or "").strip() @@ -215,6 +223,7 @@ def _validate_resume_payload(expected: Optional[dict], decision: dict) -> dict: "plan_ready": {"planning_message", "execute_plan", "plan_approval"}, "clarification": {"clarification", "planning_message"}, "plan_approval": {"plan_approval", "planning_message", "execute_plan"}, + "tool_approval": {"tool_approval"}, } if actual_kind not in compatible.get(expected_kind, {expected_kind}): raise ValueError(f"Expected a response for {expected_kind}") @@ -273,6 +282,8 @@ async def start_research( session.add_all([message, project]) session.commit() session.refresh(project) + ensure_default_policies(session) + trace_id = start_trace(session, project.id) thread_id = f"research-{project.id}-{uuid.uuid4().hex}" message.meta_json = _message_meta( @@ -288,6 +299,7 @@ async def start_research( initial_state = { "run_id": thread_id, + "trace_id": trace_id, "conversation_id": conversation_id, "project_id": project.id, "response_language": response_language, @@ -305,12 +317,15 @@ async def start_research( "findings": [], "report_md": None, "report_id": None, + "agent_tasks": [], + "verified_findings": [], } _publish_lifecycle("research.started", thread_id=thread_id, state=initial_state) try: result = await asyncio.to_thread(graph.invoke, initial_state, config) except Exception as exc: _update_project_status(session, initial_state, "failed") + finish_trace(session, trace_id, "failed") message = str(exc) if str(exc).startswith("No usable sources") else "Research run failed." _publish_lifecycle( "research.failed", @@ -333,6 +348,7 @@ async def start_research( interrupt_payload=payload, ) else: + finish_trace(session, trace_id, "completed") _publish_lifecycle( "research.completed", thread_id=thread_id, @@ -342,6 +358,7 @@ async def start_research( return { "thread_id": thread_id, + "trace_id": trace_id, "state": state, "interrupted": payload is not None, "interrupt_payload": payload, @@ -405,6 +422,8 @@ async def resume_research( result = await asyncio.to_thread(graph.invoke, Command(resume=decision), config) except Exception as exc: _update_project_status(session, state_before, "failed") + if state_before.get("trace_id"): + finish_trace(session, state_before["trace_id"], "failed") message = str(exc) if str(exc).startswith("No usable sources") else "Research run failed." _publish_lifecycle( "research.failed", @@ -427,6 +446,8 @@ async def resume_research( interrupt_payload=payload, ) else: + if state.get("trace_id"): + finish_trace(session, state["trace_id"], "completed") _publish_lifecycle( "research.completed", thread_id=thread_id, @@ -436,6 +457,7 @@ async def resume_research( return { "thread_id": thread_id, + "trace_id": state.get("trace_id"), "state": state, "interrupted": payload is not None, "interrupt_payload": payload, diff --git a/backend/app/schemas/observability.py b/backend/app/schemas/observability.py new file mode 100644 index 0000000..d7eb3f2 --- /dev/null +++ b/backend/app/schemas/observability.py @@ -0,0 +1,35 @@ +from typing import Any + +from pydantic import BaseModel, Field + + +class ToolPolicyPayload(BaseModel): + agent_role: str = Field(min_length=1, max_length=80) + tool_name: str = Field(min_length=1, max_length=160) + allowed_domains: list[str] = Field(default_factory=list) + require_approval: bool = False + enabled: bool = True + + +class ToolApprovalPayload(BaseModel): + approved: bool + profile_id: int = Field(default=1, gt=0) + agent_role: str = Field(min_length=1) + tool_name: str = Field(min_length=1) + args_fingerprint: str = Field(min_length=16) + + +class EvaluationDatasetPayload(BaseModel): + name: str = Field(min_length=1, max_length=120) + description: str = "" + judge_profile_id: int | None = Field(default=None, gt=0) + + +class EvaluationCasePayload(BaseModel): + input_text: str = Field(min_length=1, max_length=10_000) + expected: dict[str, Any] | None = None + + +class EvaluationRunPayload(BaseModel): + profile_id: int = Field(gt=0) + judge_profile_id: int | None = Field(default=None, gt=0) diff --git a/backend/app/schemas/research.py b/backend/app/schemas/research.py index 6733be1..71a306c 100644 --- a/backend/app/schemas/research.py +++ b/backend/app/schemas/research.py @@ -75,6 +75,7 @@ def validate_decision(self): class ResearchRunResponse(BaseModel): thread_id: str + trace_id: str | None = None state: dict[str, Any] interrupted: bool interrupt_payload: dict[str, Any] | None diff --git a/backend/app/services/tracing_service.py b/backend/app/services/tracing_service.py new file mode 100644 index 0000000..ae120d1 --- /dev/null +++ b/backend/app/services/tracing_service.py @@ -0,0 +1,65 @@ +"""Local, redacted tracing helpers for research runs.""" + +import datetime as dt +import re +import uuid +from typing import Any + +from sqlalchemy.orm import Session + +from app.db.models import TraceRun, TraceSpan + + +_SECRET_KEYS = re.compile(r"(api[_-]?key|token|password|secret|authorization)", re.I) + + +def _now() -> dt.datetime: + return dt.datetime.now(dt.timezone.utc).replace(tzinfo=None) + + +def redact(value: Any, *, limit: int = 1200) -> str: + if isinstance(value, dict): + value = {str(k): "[redacted]" if _SECRET_KEYS.search(str(k)) else v for k, v in value.items()} + text = str(value) + return text[:limit] + + +def start_trace(session: Session, project_id: int) -> str: + trace_id = f"trace-{uuid.uuid4().hex}" + session.add(TraceRun(id=trace_id, project_id=project_id)) + session.commit() + return trace_id + + +def finish_trace(session: Session, trace_id: str, status: str) -> None: + trace = session.get(TraceRun, trace_id) + if trace is not None: + trace.status = status + trace.completed_at = _now() + session.commit() + + +def start_span(session: Session, trace_id: str, name: str, *, kind: str = "internal", parent_id: str | None = None, attributes: dict | None = None, input_value: Any = None) -> str: + span_id = f"span-{uuid.uuid4().hex}" + session.add(TraceSpan(id=span_id, trace_id=trace_id, parent_id=parent_id, name=name, kind=kind, attributes_json=attributes, input_summary=redact(input_value) if input_value is not None else None)) + session.commit() + return span_id + + +def finish_span(session: Session, span_id: str, *, status: str = "ok", output_value: Any = None, error: Exception | str | None = None) -> None: + span = session.get(TraceSpan, span_id) + if span is not None: + span.status = status + span.output_summary = redact(output_value) if output_value is not None else None + span.error = redact(error) if error is not None else None + span.ended_at = _now() + session.commit() + + +def list_trace(session: Session, project_id: int) -> dict: + runs = session.query(TraceRun).filter(TraceRun.project_id == project_id).order_by(TraceRun.created_at.desc()).all() + traces = [] + for run in runs: + spans = session.query(TraceSpan).filter(TraceSpan.trace_id == run.id).order_by(TraceSpan.started_at).all() + traces.append({"id": run.id, "status": run.status, "created_at": run.created_at.isoformat(), "completed_at": run.completed_at.isoformat() if run.completed_at else None, "spans": [{"id": span.id, "parent_id": span.parent_id, "name": span.name, "kind": span.kind, "status": span.status, "attributes": span.attributes_json, "input_summary": span.input_summary, "output_summary": span.output_summary, "error": span.error, "started_at": span.started_at.isoformat(), "ended_at": span.ended_at.isoformat() if span.ended_at else None} for span in spans]}) + return {"project_id": project_id, "traces": traces} diff --git a/backend/app/tools/gateway.py b/backend/app/tools/gateway.py new file mode 100644 index 0000000..3b952bf --- /dev/null +++ b/backend/app/tools/gateway.py @@ -0,0 +1,68 @@ +"""The single execution point for agent tool calls.""" + +import datetime as dt +from typing import Any, Callable + +from langgraph.types import interrupt +from sqlalchemy.orm import Session + +from app.core.events import get_event_bus +from app.db.models import ToolCall +from app.services.tracing_service import finish_span, start_span +from app.tools.policy import decide, record_approval + + +def execute_tool( + session: Session, + *, + project_id: int, + trace_id: str, + agent_role: str, + tool_name: str, + args: dict, + tool: Any, + invoke: Callable[[Any, dict], Any], + publish: Callable[..., None], +) -> Any: + decision = decide(session, project_id=project_id, trace_id=trace_id, agent_role=agent_role, tool_name=tool_name, args=args) + if decision.action == "deny": + publish("research.tool_denied", tool_name=tool_name, agent_role=agent_role, reason=decision.reason) + raise PermissionError(decision.reason) + if decision.action == "approval_required": + publish("research.tool_approval_required", tool_name=tool_name, agent_role=agent_role, reason=decision.reason, agent_role_name=agent_role) + response = interrupt({ + "kind": "tool_approval", + "message": f"Agent {agent_role} requests permission to call {tool_name}.", + "agent_role": agent_role, + "tool_name": tool_name, + "args": {key: value for key, value in args.items() if key not in {"api_key", "token", "password", "secret"}}, + "args_fingerprint": decision.fingerprint, + "reason": decision.reason, + }) + approved = isinstance(response, dict) and response.get("kind") == "tool_approval" and bool(response.get("approved")) and response.get("args_fingerprint") == decision.fingerprint + record_approval(session, project_id=project_id, trace_id=trace_id, agent_role=agent_role, tool_name=tool_name, args_fingerprint=decision.fingerprint, approved=approved) + if not approved: + publish("research.tool_denied", tool_name=tool_name, agent_role=agent_role, reason="User denied the tool request.") + raise PermissionError("User denied the tool request") + span_id = start_span(session, trace_id, f"tool:{tool_name}", kind="tool", attributes={"agent_role": agent_role, "tool_name": tool_name}, input_value=args) + call = ToolCall(project_id=project_id, trace_id=trace_id, agent_role=agent_role, tool_name=tool_name, args_json={key: value for key, value in args.items() if key not in {"api_key", "token", "password", "secret"}}) + session.add(call) + session.commit() + publish("research.tool_started", tool_name=tool_name, agent_role=agent_role) + try: + result = invoke(tool, args) + call.status = "completed" + call.result_summary = str(result)[:1200] + call.completed_at = dt.datetime.now(dt.timezone.utc).replace(tzinfo=None) + session.commit() + finish_span(session, span_id, output_value=result) + publish("research.tool_completed", tool_name=tool_name, agent_role=agent_role) + return result + except Exception as exc: + call.status = "failed" + call.error = str(exc)[:1200] + call.completed_at = dt.datetime.now(dt.timezone.utc).replace(tzinfo=None) + session.commit() + finish_span(session, span_id, status="error", error=exc) + publish("research.tool_completed", tool_name=tool_name, agent_role=agent_role, error=str(exc)) + raise diff --git a/backend/app/tools/policy.py b/backend/app/tools/policy.py new file mode 100644 index 0000000..0d071f7 --- /dev/null +++ b/backend/app/tools/policy.py @@ -0,0 +1,106 @@ +"""Tool policy enforcement. Policies are deny-by-default per agent role.""" + +import hashlib +import json +from dataclasses import dataclass +from urllib.parse import urlparse + +from sqlalchemy import func +from sqlalchemy.orm import Session + +from app.db.models import ToolApproval, ToolPolicy + + +@dataclass(frozen=True) +class PolicyDecision: + action: str # allow | approval_required | deny + reason: str + fingerprint: str + + +def args_fingerprint(args: dict) -> str: + canonical = json.dumps(args, ensure_ascii=True, sort_keys=True, separators=(",", ":")) + return hashlib.sha256(canonical.encode()).hexdigest() + + +def _url_host(args: dict) -> str | None: + value = args.get("url") + if not isinstance(value, str): + return None + return (urlparse(value).hostname or "").lower() or None + + +def policy_version(session: Session) -> int: + return int(session.scalar(func.max(ToolPolicy.version)) or 0) + + +def list_policies(session: Session) -> list[dict]: + rows = session.query(ToolPolicy).order_by(ToolPolicy.agent_role, ToolPolicy.tool_name, ToolPolicy.id).all() + return [{"id": item.id, "version": item.version, "agent_role": item.agent_role, "tool_name": item.tool_name, "allowed_domains": item.allowed_domains_json or [], "require_approval": item.require_approval, "enabled": item.enabled, "created_at": item.created_at.isoformat()} for item in rows] + + +def create_policy(session: Session, *, agent_role: str, tool_name: str, allowed_domains: list[str] | None = None, require_approval: bool = False, enabled: bool = True) -> dict: + row = ToolPolicy(version=policy_version(session) + 1, agent_role=agent_role, tool_name=tool_name, allowed_domains_json=sorted({domain.lower().strip() for domain in allowed_domains or [] if domain.strip()}), require_approval=require_approval, enabled=enabled) + session.add(row) + session.commit() + return next(item for item in list_policies(session) if item["id"] == row.id) + + +def update_policy(session: Session, policy_id: int, **values) -> dict | None: + row = session.get(ToolPolicy, policy_id) + if row is None: + return None + row.version = policy_version(session) + 1 + for key in ("agent_role", "tool_name", "require_approval", "enabled"): + if key in values: + setattr(row, key, values[key]) + if "allowed_domains" in values: + row.allowed_domains_json = sorted({domain.lower().strip() for domain in values["allowed_domains"] or [] if domain.strip()}) + session.commit() + return next(item for item in list_policies(session) if item["id"] == row.id) + + +def delete_policy(session: Session, policy_id: int) -> bool: + row = session.get(ToolPolicy, policy_id) + if row is None: + return False + session.delete(row) + session.commit() + return True + + +def visible_tool_names(session: Session, agent_role: str, names: list[str]) -> set[str]: + policies = session.query(ToolPolicy).filter(ToolPolicy.agent_role == agent_role, ToolPolicy.enabled.is_(True)).all() + allowed = {policy.tool_name for policy in policies} + return set(names).intersection(allowed) + + +def decide(session: Session, *, project_id: int, trace_id: str, agent_role: str, tool_name: str, args: dict) -> PolicyDecision: + fingerprint = args_fingerprint(args) + policy = session.query(ToolPolicy).filter(ToolPolicy.agent_role == agent_role, ToolPolicy.tool_name == tool_name, ToolPolicy.enabled.is_(True)).order_by(ToolPolicy.id.desc()).first() + if policy is None: + # A missing rule never grants execution, but can be explicitly approved for this run. + return PolicyDecision("approval_required", "No policy allows this tool; approval is required for this task.", fingerprint) + host = _url_host(args) + domains = policy.allowed_domains_json or [] + if domains and host and not any(host == domain or host.endswith(f".{domain}") for domain in domains): + return PolicyDecision("deny", "The requested URL domain is not allowed by the policy.", fingerprint) + if policy.require_approval: + approval = session.query(ToolApproval).filter(ToolApproval.project_id == project_id, ToolApproval.trace_id == trace_id, ToolApproval.agent_role == agent_role, ToolApproval.tool_name == tool_name, ToolApproval.args_fingerprint == fingerprint, ToolApproval.decision == "approved").first() + if approval is None: + return PolicyDecision("approval_required", "This policy requires approval for the current task.", fingerprint) + return PolicyDecision("allow", "Allowed by policy.", fingerprint) + + +def record_approval(session: Session, *, project_id: int, trace_id: str, agent_role: str, tool_name: str, args_fingerprint: str, approved: bool) -> None: + session.add(ToolApproval(project_id=project_id, trace_id=trace_id, agent_role=agent_role, tool_name=tool_name, args_fingerprint=args_fingerprint, decision="approved" if approved else "denied")) + session.commit() + + +def ensure_default_policies(session: Session) -> None: + """Install conservative built-in read-only rules once for existing projects.""" + existing = {(item.agent_role, item.tool_name) for item in session.query(ToolPolicy).all()} + for role, tool_name in (("retriever", "fetch_page"), ("retriever", "search_web"), ("verifier", "fetch_page")): + if (role, tool_name) not in existing: + session.add(ToolPolicy(version=1, agent_role=role, tool_name=tool_name, allowed_domains_json=[], require_approval=False, enabled=True)) + session.commit() diff --git a/backend/tests/test_observability.py b/backend/tests/test_observability.py new file mode 100644 index 0000000..2b5caae --- /dev/null +++ b/backend/tests/test_observability.py @@ -0,0 +1,38 @@ +import pytest + + +@pytest.fixture() +def session(app_home): + import app.db.session as db + + app_home.mkdir(parents=True, exist_ok=True) + engine = db.init_db(f"sqlite:///{(app_home / 'observability.db').as_posix()}") + with db.SessionLocal() as current: + yield current + engine.dispose() + + +def test_policy_fingerprint_is_stable_and_redacts_secrets(): + from app.services.tracing_service import redact + from app.tools.policy import args_fingerprint + + assert args_fingerprint({"url": "https://example.com", "q": "x"}) == args_fingerprint({"q": "x", "url": "https://example.com"}) + assert "secret-value" not in redact({"api_key": "secret-value", "url": "https://example.com"}) + + +def test_policy_domain_matching_and_approval(session): + from app.db.models import Conversation, ResearchProject + from app.tools.policy import create_policy, decide + + conversation = Conversation(title="policy") + session.add(conversation) + session.flush() + project = ResearchProject(conversation_id=conversation.id, topic="policy", objective="policy") + session.add(project) + session.commit() + + create_policy(session, agent_role="retriever", tool_name="fetch_page", allowed_domains=["example.com"], require_approval=True) + allowed = decide(session, project_id=project.id, trace_id="trace-1", agent_role="retriever", tool_name="fetch_page", args={"url": "https://example.com/a"}) + blocked = decide(session, project_id=project.id, trace_id="trace-1", agent_role="retriever", tool_name="fetch_page", args={"url": "https://other.example/a"}) + assert allowed.action == "approval_required" + assert blocked.action == "deny" diff --git a/frontend/src/App.tsx b/frontend/src/App.tsx index 11b8465..5d5d7c3 100644 --- a/frontend/src/App.tsx +++ b/frontend/src/App.tsx @@ -14,6 +14,7 @@ import type { ResearchLifecycleEvent, ResearchRunPhase, ResearchRunResponse, + ToolPolicy, } from "./api/types"; import { AppPanel, AppShell } from "./components/AppShell"; import { ConversationPanel } from "./components/ConversationPanel"; @@ -48,6 +49,9 @@ function deriveRunPhaseFromResponse(run: ResearchRunResponse): ResearchRunPhase if (run.interrupted && (kind === "plan_ready" || kind === "plan_approval")) { return "awaiting_execution"; } + if (run.interrupted && kind === "tool_approval") { + return "awaiting_approval"; + } if (run.state.phase === "executing" || run.state.status === "executing") { return "executing"; } @@ -177,6 +181,7 @@ export default function App() { const [profiles, setProfiles] = useState([]); const [selectedProfileId, setSelectedProfileId] = useState(null); const [servers, setServers] = useState([]); + const [policies, setPolicies] = useState([]); const [maxSources, setMaxSources] = useState(8); const [settingsLoadErrors, setSettingsLoadErrors] = useState({ profiles: null, @@ -303,6 +308,12 @@ export default function App() { return; } + if (event.event === "research.tool_approval_required") { + updateRunPhase("awaiting_approval"); + setStatusMessageKey(statusMessageForPhase("awaiting_approval")); + return; + } + if (event.event === "research.completed") { updateRunPhase("completed"); setStatusMessageKey(statusMessageForPhase("completed")); @@ -381,10 +392,11 @@ export default function App() { servers: null, }); - const [profileResult, sourceLimitResult, serverResult] = await Promise.allSettled([ + const [profileResult, sourceLimitResult, serverResult, policyResult] = await Promise.allSettled([ api.listLLMProfiles(), loadMaxSourcesPreference(), api.listMCPServers(), + typeof api.listToolPolicies === "function" ? api.listToolPolicies() : Promise.resolve([] as ToolPolicy[]), ]); if (cancelled) { @@ -422,6 +434,9 @@ export default function App() { servers: messageFromError(serverResult.reason, "app.failedToLoadMcpServers"), })); } + if (policyResult.status === "fulfilled") { + setPolicies(policyResult.value); + } } void loadSettingsState(); @@ -737,6 +752,32 @@ export default function App() { } } + async function handleToolApproval(approved: boolean) { + if (!currentRun || !selectedProfileId || resumePendingRef.current) return; + const payload = currentRun.interrupt_payload ?? {}; + resumePendingRef.current = true; + try { + updateRunPhase("resuming"); + const run = await api.resumeResearch(currentRun.thread_id, { + profile_id: selectedProfileId, + response_language: locale, + decision: { + kind: "tool_approval", + approved, + agent_role: payload.agent_role, + tool_name: payload.tool_name, + args_fingerprint: payload.args_fingerprint, + }, + }); + await settleRun(run); + } catch (error) { + updateRunPhase("awaiting_approval"); + setErrorMessage(messageFromError(error, "app.failedToResumeResearch")); + } finally { + resumePendingRef.current = false; + } + } + async function handleRenameConversation(conversationId: number, title: string) { try { setErrorMessage(null); @@ -851,6 +892,21 @@ export default function App() { setSettingsLoadErrors((current) => ({ ...current, servers: null })); } + async function handleCreatePolicy(payload: Omit) { + const created = await api.createToolPolicy(payload); + setPolicies((current) => [...current, created]); + } + + async function handleUpdatePolicy(id: number, payload: Omit) { + const updated = await api.updateToolPolicy(id, payload); + setPolicies((current) => current.map((item) => item.id === id ? updated : item)); + } + + async function handleDeletePolicy(id: number) { + await api.deleteToolPolicy(id); + setPolicies((current) => current.filter((item) => item.id !== id)); + } + const backendDefaultProfile = profiles.find((profile) => profile.is_default) ?? null; const selectedProfile = profiles.find((profile) => profile.id === selectedProfileId) ?? null; @@ -902,6 +958,10 @@ export default function App() { onSaveMaxSources={handleSaveMaxSources} onCreateServer={handleCreateServer} onUpdateServer={handleUpdateServer} + policies={policies} + onCreatePolicy={handleCreatePolicy} + onUpdatePolicy={handleUpdatePolicy} + onDeletePolicy={handleDeletePolicy} loadErrors={settingsLoadErrors} /> ) : ( @@ -927,6 +987,8 @@ export default function App() { } : null } + toolApprovalPrompt={currentRun?.interrupt_payload?.kind === "tool_approval" ? currentRun.interrupt_payload : null} + onToolApproval={handleToolApproval} onSubmitPlanningAnswers={handlePlanningAnswers} detailsOpen={detailsOpen} onToggleDetails={() => setDetailsOpen((current) => !current)} diff --git a/frontend/src/api/client.ts b/frontend/src/api/client.ts index 741d465..7b47e77 100644 --- a/frontend/src/api/client.ts +++ b/frontend/src/api/client.ts @@ -9,6 +9,7 @@ import type { ProjectExecutionDetail, ResearchFollowUpResponse, ResearchRunResponse, + ToolPolicy, } from "./types"; export class ApiError extends Error { @@ -73,6 +74,11 @@ export const api = { }) => request("/research/follow-up", json("POST", payload)), getProjectExecution: (projectId: number) => request(`/research/projects/${projectId}/execution`), + listToolPolicies: () => request("/tool-policies"), + createToolPolicy: (payload: Omit) => request("/tool-policies", json("POST", payload)), + updateToolPolicy: (id: number, payload: Omit) => request(`/tool-policies/${id}`, json("PUT", payload)), + deleteToolPolicy: (id: number) => request(`/tool-policies/${id}`, { method: "DELETE" }), + approveToolCall: (threadId: string, payload: { approved: boolean; profile_id?: number; agent_role: string; tool_name: string; args_fingerprint: string }) => request(`/research/${encodeURIComponent(threadId)}/tool-approval`, json("POST", payload)), listLLMProfiles: () => request("/settings/llm-profiles"), createLLMProfile: (payload: LLMProfileCreate) => request("/settings/llm-profiles", json("POST", payload)), diff --git a/frontend/src/api/types.ts b/frontend/src/api/types.ts index 1edd024..52ca09a 100644 --- a/frontend/src/api/types.ts +++ b/frontend/src/api/types.ts @@ -80,6 +80,7 @@ export type PlanningAnswer = { export type ResearchRunResponse = { thread_id: string; + trace_id?: string | null; state: Record & { report_id?: number }; interrupted: boolean; interrupt_payload: Record | null; @@ -164,6 +165,9 @@ export type ProjectExecutionDetail = { tool_name: string; }>; reports: Array<{ id: number; version: number; created_at: string }>; + agent_tasks?: Array<{ id: number; role: string; title: string; status: string; input?: Record | null; output?: Record | null }>; + tool_calls?: Array<{ id: number; task_id?: number | null; agent_role: string; tool_name: string; status: string; result_summary?: string | null; error?: string | null }>; + approvals?: Array<{ id: number; agent_role: string; tool_name: string; decision: string; args_fingerprint: string }>; }; export type ResearchFollowUpResponse = { @@ -182,6 +186,8 @@ export const RESEARCH_LIFECYCLE_EVENT_TYPES = [ "research.step_completed", "research.tool_started", "research.tool_completed", + "research.tool_denied", + "research.tool_approval_required", "research.source_collected", "research.report_revision_started", "research.report_revised", @@ -196,6 +202,17 @@ export const RESEARCH_LIFECYCLE_EVENT_TYPES = [ export type ResearchLifecycleEventType = (typeof RESEARCH_LIFECYCLE_EVENT_TYPES)[number]; +export type ToolPolicy = { + id: number; + version: number; + agent_role: string; + tool_name: string; + allowed_domains: string[]; + require_approval: boolean; + enabled: boolean; + created_at: string; +}; + export type ResearchLifecyclePayload = Record & { thread_id: string; }; diff --git a/frontend/src/components/ExecutionDrawer.tsx b/frontend/src/components/ExecutionDrawer.tsx index 1832d5a..d918d8f 100644 --- a/frontend/src/components/ExecutionDrawer.tsx +++ b/frontend/src/components/ExecutionDrawer.tsx @@ -46,6 +46,8 @@ export function ExecutionDrawer({
+ {detail?.agent_tasks?.length ?

Agent 任务

{detail.agent_tasks.map((task) =>
{task.role} · {task.title}{task.status}
)}
: null} + {detail?.tool_calls?.length ?

工具调用

{detail.tool_calls.map((call) =>
{call.agent_role} · {call.tool_name}{call.status}
{call.error ?

{call.error}

: null}
)}
: null}
diff --git a/frontend/src/components/ResearchWorkspace.tsx b/frontend/src/components/ResearchWorkspace.tsx index accc76e..51bd6c0 100644 --- a/frontend/src/components/ResearchWorkspace.tsx +++ b/frontend/src/components/ResearchWorkspace.tsx @@ -28,6 +28,8 @@ type ResearchWorkspaceProps = { detailsOpen?: boolean; onToggleDetails?: () => void; executionProgress?: { completed: number; total: number } | null; + toolApprovalPrompt?: Record | null; + onToolApproval?: (approved: boolean) => void; }; export function ResearchWorkspace({ @@ -53,6 +55,8 @@ export function ResearchWorkspace({ detailsOpen = false, onToggleDetails, executionProgress = null, + toolApprovalPrompt = null, + onToolApproval, }: ResearchWorkspaceProps) { const { t } = useI18n(); const [message, setMessage] = useState(""); @@ -219,6 +223,17 @@ export function ResearchWorkspace({ ) : null} {timelineContent ?
{timelineContent}
: null} + {toolApprovalPrompt ? ( +
+
工具调用需要确认
+

{String(toolApprovalPrompt.agent_role ?? "agent")} 请求调用 {String(toolApprovalPrompt.tool_name ?? "tool")}。

+ {toolApprovalPrompt.reason ?

{String(toolApprovalPrompt.reason)}

: null} +
+ + +
+
+ ) : null} {optimisticUserMessage ? (
diff --git a/frontend/src/components/SettingsPanel.tsx b/frontend/src/components/SettingsPanel.tsx index 5c0e3e3..467d883 100644 --- a/frontend/src/components/SettingsPanel.tsx +++ b/frontend/src/components/SettingsPanel.tsx @@ -1,6 +1,7 @@ import { useEffect, useMemo, useState } from "react"; import { Check, Pencil, PlugZap, Save, X } from "lucide-react"; -import type { LLMProfileCreate, LLMProfileRead, MCPServer } from "../api/types"; +import type { LLMProfileCreate, LLMProfileRead, MCPServer, ToolPolicy } from "../api/types"; +import { ToolPolicyPanel } from "./ToolPolicyPanel"; import { useI18n } from "../i18n/I18nProvider"; import { localizedMessage, @@ -30,6 +31,10 @@ type SettingsPanelProps = { onSaveMaxSources: (value: number) => Promise | void; onCreateServer: (payload: Omit) => Promise | void; onUpdateServer?: (serverId: number, payload: Omit) => Promise | void; + policies?: ToolPolicy[]; + onCreatePolicy?: (payload: Omit) => Promise; + onUpdatePolicy?: (id: number, payload: Omit) => Promise; + onDeletePolicy?: (id: number) => Promise; }; type Feedback = { @@ -190,6 +195,7 @@ export function SettingsPanel({ onSaveMaxSources, onCreateServer, onUpdateServer, + policies = [], onCreatePolicy, onUpdatePolicy, onDeletePolicy, }: SettingsPanelProps) { const { t } = useI18n(); const [profileForm, setProfileForm] = useState({ @@ -694,6 +700,8 @@ export function SettingsPanel({ + {onCreatePolicy && onUpdatePolicy && onDeletePolicy ? : null} +

diff --git a/frontend/src/components/ToolPolicyPanel.tsx b/frontend/src/components/ToolPolicyPanel.tsx new file mode 100644 index 0000000..e7d36b4 --- /dev/null +++ b/frontend/src/components/ToolPolicyPanel.tsx @@ -0,0 +1,51 @@ +import { useEffect, useState } from "react"; +import { Plus, Save, Trash2 } from "lucide-react"; +import type { ToolPolicy } from "../api/types"; + +export function ToolPolicyPanel({ + policies, + onCreate, + onUpdate, + onDelete, +}: { + policies: ToolPolicy[]; + onCreate: (payload: Omit) => Promise; + onUpdate: (id: number, payload: Omit) => Promise; + onDelete: (id: number) => Promise; +}) { + const [role, setRole] = useState("retriever"); + const [tool, setTool] = useState("fetch_page"); + const [domains, setDomains] = useState(""); + const [approval, setApproval] = useState(false); + const [editing, setEditing] = useState(null); + const [busy, setBusy] = useState(false); + + useEffect(() => { + if (editing === null) return; + const current = policies.find((item) => item.id === editing); + if (!current) return; + setRole(current.agent_role); setTool(current.tool_name); setDomains(current.allowed_domains.join("\n")); setApproval(current.require_approval); + }, [editing, policies]); + + async function save() { + setBusy(true); + const payload = { agent_role: role.trim(), tool_name: tool.trim(), allowed_domains: domains.split(/\r?\n/).map((item) => item.trim()).filter(Boolean), require_approval: approval, enabled: true }; + try { if (editing === null) await onCreate(payload); else await onUpdate(editing, payload); setEditing(null); } finally { setBusy(false); } + } + + return ( +
+

Agent 工具权限

未声明的工具默认需要本次任务确认;域名限制只适用于带 URL 参数的调用。

+
+ + +
+