czty's picture
Add files using upload-large-folder tool
b2c86fd verified
Raw
History Blame Contribute Delete
15.8 kB
"""
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()