""" FastAPI backend for Biomni Web UI. 核心:把 Biomni agent 的执行过程包装成 SSE 流式响应。 """ import asyncio import json import os import queue import re import threading import traceback import uuid from datetime import datetime, timezone from contextlib import asynccontextmanager from pathlib import Path from typing import Any, AsyncGenerator from fastapi import FastAPI, File, Form, HTTPException, UploadFile from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import FileResponse from fastapi.responses import StreamingResponse from pydantic import BaseModel from mcp_registry import discover_mcp_servers, load_latest_graph_server_count from session import SessionManager from streaming import StreamingCallback from metabolomics_tools import generate_scmeta_mcp, list_metabolomics_tools from tool2mcp import ( build_docker_image, claude_addition, convert_mcptool, generate_environment_yaml, generate_requirements_with_pipreqs, ) DEFAULT_LLM_PROVIDER = os.getenv("BIOMNI_LLM_PROVIDER", "deepseek").strip().lower() # ── 全局 session 管理器(启动时初始化一次)──────────────────────── session_manager = SessionManager() def _artifact_root_dir(session_id: str, root: str) -> Path: if root == "session": return (Path("./data") / "biomni_sessions" / session_id).resolve() if root == "report": return (Path("./data") / "reports" / session_id).resolve() raise HTTPException(status_code=400, detail=f"Unsupported artifact root: {root}") def _resolve_artifact_file(session_id: str, root: str, relative_path: str) -> Path: root_dir = _artifact_root_dir(session_id, root) candidate = (root_dir / relative_path).resolve() if not str(candidate).startswith(str(root_dir)): raise HTTPException(status_code=400, detail="Invalid artifact path") if not candidate.exists() or not candidate.is_file(): raise HTTPException(status_code=404, detail="Artifact not found") return candidate @asynccontextmanager async def lifespan(_: FastAPI): print("Biomni backend starting...") yield session_manager.cleanup_all() app = FastAPI(lifespan=lifespan) app.add_middleware( CORSMiddleware, allow_origins=["*"], # 开发阶段先放开所有来源 allow_methods=["*"], allow_headers=["*"], ) # ── 数据模型 ───────────────────────────────────────────────────── class ChatRequest(BaseModel): message: str session_id: str | None = None # None = 新建 session class CreateSessionRequest(BaseModel): llm_provider: str = DEFAULT_LLM_PROVIDER class NewSessionResponse(BaseModel): session_id: str class UploadResponse(BaseModel): files: list[dict[str, str]] uploaded_count: int class MetabolomicsToolListResponse(BaseModel): tools: list[dict[str, Any]] class McpServerListResponse(BaseModel): servers: list[dict[str, Any]] catalog_server_count: int | None = None class ConverterGenerateResponse(BaseModel): ok: bool tool_name: str llm_provider: str run_help_command: bool server_path: str files: dict[str, str] warnings: list[str] = [] # ── 接口 ───────────────────────────────────────────────────────── @app.post("/session/new", response_model=NewSessionResponse) async def new_session(payload: CreateSessionRequest | None = None): """创建新的 session,返回 session_id""" sid = session_manager.create_session( llm_provider=(payload.llm_provider if payload else DEFAULT_LLM_PROVIDER) ) return NewSessionResponse(session_id=sid) @app.post("/session/{session_id}/upload", response_model=UploadResponse) async def upload_file_to_session( session_id: str, file: UploadFile | None = File(default=None), files: list[UploadFile] | None = File(default=None), ): """上传单个或多个文件并挂到 session,对话时将基于上传文件处理。""" if not session_manager.has_session(session_id): raise HTTPException(status_code=404, detail="Session not found") incoming_files: list[UploadFile] = [] if files: incoming_files.extend(files) if file is not None: incoming_files.append(file) if not incoming_files: raise HTTPException(status_code=400, detail="No file is provided") session_dir = Path("./data/uploads") / session_id session_dir.mkdir(parents=True, exist_ok=True) uploaded_items: list[dict[str, str]] = [] for upload in incoming_files: if not upload.filename: raise HTTPException(status_code=400, detail="Filename is empty") safe_name = Path(upload.filename).name stored_name = f"{uuid.uuid4()}_{safe_name}" save_path = session_dir / stored_name content = await upload.read() with open(save_path, "wb") as f: f.write(content) relative_path = os.path.relpath(save_path, start="./data") session_manager.add_uploaded_file(session_id, relative_path) uploaded_items.append({ "filename": safe_name, "relative_path": relative_path, }) return UploadResponse(files=uploaded_items, uploaded_count=len(uploaded_items)) @app.get("/metabolomics/tools", response_model=MetabolomicsToolListResponse) async def get_metabolomics_tools(): return MetabolomicsToolListResponse(tools=list_metabolomics_tools()) @app.get("/mcp/servers", response_model=McpServerListResponse) async def get_mcp_servers(): return McpServerListResponse( servers=discover_mcp_servers(), catalog_server_count=load_latest_graph_server_count(), ) @app.post("/metabolomics/scmeta/generate") async def generate_scmeta(): result = generate_scmeta_mcp() if not result.get("ok"): raise HTTPException( status_code=500, detail={ "message": "Failed to generate scmeta MCP server", "stderr": result.get("stderr", ""), "stdout": result.get("stdout", ""), }, ) return result @app.post("/mcp/converter/generate", response_model=ConverterGenerateResponse) async def converter_generate( tool_name: str = Form(...), llm_provider: str = Form("claude"), manual_pdf: UploadFile | None = File(default=None), ): safe_tool_name = re.sub(r"[^a-zA-Z0-9_\-]", "", tool_name).strip() if not safe_tool_name: raise HTTPException(status_code=400, detail="Invalid tool name") output_root = Path("./data/mcp_generated").resolve() output_root.mkdir(parents=True, exist_ok=True) server_path = output_root / f"mcp_{safe_tool_name}" app_path = server_path / "app" app_path.mkdir(parents=True, exist_ok=True) run_help_command = manual_pdf is None manual = "--help" if manual_pdf is not None: if not manual_pdf.filename: raise HTTPException(status_code=400, detail="manual_pdf filename is empty") allowed_exts = (".pdf", ".md", ".markdown", ".txt") lower_name = manual_pdf.filename.lower() if not lower_name.endswith(allowed_exts): raise HTTPException( status_code=400, detail="manual_pdf must be one of: .pdf, .md, .markdown, .txt", ) manual_dir = Path("./data/uploads/converter_manuals").resolve() manual_dir.mkdir(parents=True, exist_ok=True) manual_path = manual_dir / f"{uuid.uuid4()}_{Path(manual_pdf.filename).name}" content = await manual_pdf.read() manual_path.write_bytes(content) manual = str(manual_path) warnings: list[str] = [] try: convert_mcptool( safe_tool_name, manual=manual, run_help_command=run_help_command, server_path=server_path, llm_provider=llm_provider, ) except Exception as exc: raise HTTPException( status_code=500, detail=f"convert_mcptool failed: {exc}", ) from exc req_path = server_path / "requirements.txt" try: req_path = generate_requirements_with_pipreqs(safe_tool_name, server_path) except Exception as exc: warnings.append(f"generate requirements.txt failed: {exc}") env_path = generate_environment_yaml(safe_tool_name, server_path) try: build_status = build_docker_image(safe_tool_name, server_path, output_root, is_pipeline=False) if not build_status: warnings.append("build_docker_image returned failed status") except Exception as exc: warnings.append(f"build_docker_image failed: {exc}") return ConverterGenerateResponse( ok=True, tool_name=safe_tool_name, llm_provider=llm_provider, run_help_command=run_help_command, server_path=str(server_path), files={ "server": str(app_path / f"{safe_tool_name}_server.py"), "requirements": str(req_path), "environment": str(env_path), "dockerfile": str(server_path / "Dockerfile"), "docker_compose": str(server_path / "docker-compose.yml"), "claude_config_hint": claude_addition(safe_tool_name), }, warnings=warnings, ) @app.get("/chat/stream") async def chat_stream( message: str, session_id: str, metabolomics_enabled: bool = False, metabolomics_tools: str = "", paper_repro_enabled: bool = False, paper_profile: str = "foxj1", llm_provider: str = DEFAULT_LLM_PROVIDER, ): """ 主接口:接收用户消息,返回 SSE 流 SSE 事件格式: data: {"type": "thinking", "content": "..."} data: {"type": "code", "lang": "python", "content": "..."} data: {"type": "tool_use", "tool": "admet_prediction", "content": "..."} data: {"type": "result", "content": "..."} data: {"type": "error", "content": "..."} data: {"type": "done"} """ if not session_manager.has_session(session_id): raise HTTPException(status_code=404, detail="Session not found") session_manager.set_llm_provider(session_id, llm_provider) selected_tools = [x.strip() for x in metabolomics_tools.split(",") if x.strip()] session_manager.set_metabolomics_settings(session_id, metabolomics_enabled, selected_tools) session_manager.set_paper_reproduction_settings(session_id, paper_repro_enabled, paper_profile) return StreamingResponse( _stream_agent_response(message, session_id), media_type="text/event-stream", headers={ "Cache-Control": "no-cache", "X-Accel-Buffering": "no", # 关闭 nginx 缓冲 }, ) @app.delete("/session/{session_id}") async def delete_session(session_id: str): session_manager.delete_session(session_id) return {"ok": True} @app.post("/session/{session_id}/stop") async def stop_session_run(session_id: str): if not session_manager.has_session(session_id): raise HTTPException(status_code=404, detail="Session not found") session_manager.request_stop(session_id) return {"ok": True} @app.get("/session/{session_id}/artifacts/download") async def download_session_artifact(session_id: str, root: str, path: str): if not session_manager.has_session(session_id): raise HTTPException(status_code=404, detail="Session not found") artifact_path = _resolve_artifact_file(session_id, root, path) return FileResponse(path=artifact_path, filename=artifact_path.name) # ── 核心:同步 agent 转异步 SSE ─────────────────────────────────── def _create_run_log_file(session_id: str) -> tuple[Path, Any]: log_dir = Path("./data/logs") / session_id log_dir.mkdir(parents=True, exist_ok=True) ts = datetime.now(timezone.utc).strftime("%Y%m%dT%H%M%S_%fZ") log_path = log_dir / f"run_{ts}.log" return log_path, open(log_path, "a", encoding="utf-8") def _append_run_log(log_file: Any, record_type: str, payload: dict): row = { "timestamp": datetime.now(timezone.utc).isoformat(), "type": record_type, "payload": payload, } log_file.write(json.dumps(row, ensure_ascii=False) + "\n") log_file.flush() async def _stream_agent_response( message: str, session_id: str ) -> AsyncGenerator[str, None]: """ 关键流程: 1. 从 session 取出 agent 实例 2. 创建一个线程安全的 Queue 3. 在子线程中运行 agent(优先走 stream_events) 并把中间步骤放进 Queue 4. 主线程(asyncio event loop)从 Queue 取数据 yield 给客户端 """ agent = session_manager.get_agent(session_id) session_manager.clear_stop(session_id) event_queue: queue.Queue = queue.Queue() callback = StreamingCallback(event_queue) log_path, log_file = _create_run_log_file(session_id) _append_run_log(log_file, "run_start", {"session_id": session_id, "message": message}) def run_agent(): """在独立线程运行,防止阻塞 event loop""" try: if hasattr(agent, "clear_cancel_request"): agent.clear_cancel_request() if hasattr(agent, "stream_events"): for event in agent.stream_events(message): if session_manager.is_stop_requested(session_id): if hasattr(agent, "request_cancel"): agent.request_cancel() event_queue.put({"type": "error", "content": "Run cancelled by user."}) break event_queue.put(event) else: # 兼容旧实现:通过 stdout 拦截转事件 callback.attach(agent) agent.go(message) event_queue.put({"type": "done"}) except Exception: _append_run_log( log_file, "agent_exception", {"traceback": traceback.format_exc()}, ) event_queue.put({ "type": "error", "content": traceback.format_exc() }) event_queue.put({"type": "done"}) finally: if not hasattr(agent, "stream_events"): callback.detach(agent) # 在线程池运行 agent(不阻塞 asyncio) loop = asyncio.get_event_loop() thread = threading.Thread(target=run_agent, daemon=True) thread.start() # 持续从 Queue 取事件并 yield try: while True: if session_manager.is_stop_requested(session_id): yield f"data: {json.dumps({'type': 'error', 'content': 'Run cancelled by user.'}, ensure_ascii=False)}\n\n" yield "data: {\"type\": \"done\"}\n\n" break try: # 非阻塞轮询,让出控制权给 event loop event = await loop.run_in_executor( None, lambda: event_queue.get(timeout=0.1) ) except queue.Empty: # 心跳,防止连接超时 yield "data: {\"type\": \"heartbeat\"}\n\n" continue _append_run_log(log_file, "event", event) yield f"data: {json.dumps(event, ensure_ascii=False)}\n\n" if event.get("type") == "done": break finally: session_manager.clear_stop(session_id) _append_run_log(log_file, "run_end", {"log_path": str(log_path)}) log_file.close()