feat: add MCP research access and browser key management
This commit is contained in:
+64
-3
@@ -5,7 +5,7 @@ import csv
|
||||
import io
|
||||
import time
|
||||
from collections import defaultdict
|
||||
from contextlib import asynccontextmanager
|
||||
from contextlib import AsyncExitStack, asynccontextmanager
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, Depends, FastAPI, HTTPException, Query, Request, Response
|
||||
@@ -24,6 +24,7 @@ from .catalog.routes import router as catalog_router
|
||||
from .config import Settings
|
||||
from .db import create_database
|
||||
from .jobs import AUTH_KINDS, Runner, create_job
|
||||
from .mcp_api.token_routes import router as mcp_token_router
|
||||
from .models import Account, Admin, BacktestConfig, Job, JobItem, LoginSession
|
||||
from .research.routes import router as research_router
|
||||
from .research.runtime import ResearchRuntime
|
||||
@@ -93,6 +94,12 @@ def create_app(settings=None, wq_client=None, ai_model_factory=None):
|
||||
ai_runtime = AIRuntime(sessions, settings, runner, ai_model_factory)
|
||||
research_runtime = ResearchRuntime(sessions, ai_runtime, runner)
|
||||
|
||||
mcp_runtime = None
|
||||
if settings.mcp_enabled:
|
||||
from .mcp_api.server import MCPResearchServer
|
||||
|
||||
mcp_runtime = MCPResearchServer(sessions, runner, settings)
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(app):
|
||||
async with sessions() as db:
|
||||
@@ -104,7 +111,10 @@ def create_app(settings=None, wq_client=None, ai_model_factory=None):
|
||||
if settings.enable_runner:
|
||||
await runner.start()
|
||||
await research_runtime.start()
|
||||
yield
|
||||
async with AsyncExitStack() as stack:
|
||||
if mcp_runtime:
|
||||
await stack.enter_async_context(mcp_runtime.server.session_manager.run())
|
||||
yield
|
||||
if settings.enable_runner:
|
||||
await research_runtime.stop()
|
||||
await ai_runtime.stop()
|
||||
@@ -124,6 +134,7 @@ def create_app(settings=None, wq_client=None, ai_model_factory=None):
|
||||
app.state.settings = settings
|
||||
app.state.ai = ai_runtime
|
||||
app.state.research = research_runtime
|
||||
app.state.mcp = mcp_runtime
|
||||
login_failures = defaultdict(list)
|
||||
|
||||
@app.exception_handler(RequestValidationError)
|
||||
@@ -139,7 +150,54 @@ def create_app(settings=None, wq_client=None, ai_model_factory=None):
|
||||
|
||||
@app.middleware("http")
|
||||
async def browser_security(request, call_next):
|
||||
if request.method not in ("GET", "HEAD", "OPTIONS"):
|
||||
is_mcp = request.url.path in ("/api/v1/mcp", "/api/v1/mcp/")
|
||||
if is_mcp:
|
||||
if not mcp_runtime:
|
||||
return JSONResponse({"detail": "MCP 未启用"}, status_code=404)
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
from .mcp_api.auth import authenticate
|
||||
from .mcp_api.server import TOOLS
|
||||
|
||||
if request.headers.get("host", "").lower() != urlsplit(settings.public_origin).netloc.lower():
|
||||
return JSONResponse({"detail": "MCP Host 不被允许"}, status_code=403)
|
||||
origin = request.headers.get("origin")
|
||||
if origin and origin.rstrip("/") != settings.public_origin.rstrip("/"):
|
||||
return JSONResponse({"detail": "MCP Origin 不被允许"}, status_code=403)
|
||||
scheme, _, secret = request.headers.get("authorization", "").partition(" ")
|
||||
if scheme.lower() != "bearer" or not secret or len(secret) > 256:
|
||||
return JSONResponse({"detail": "需要 MCP Bearer 令牌"}, status_code=401,
|
||||
headers={"WWW-Authenticate": "Bearer"})
|
||||
try:
|
||||
async with sessions() as db:
|
||||
principal = await authenticate(db, secret)
|
||||
except HTTPException as exc:
|
||||
return JSONResponse({"detail": exc.detail}, status_code=exc.status_code,
|
||||
headers={"WWW-Authenticate": "Bearer"})
|
||||
request.state.mcp_principal = principal
|
||||
if "research:read" not in principal.scopes:
|
||||
return JSONResponse({"detail": "缺少读取权限"}, status_code=403)
|
||||
if request.method == "POST":
|
||||
body = bytearray()
|
||||
async for chunk in request.stream():
|
||||
body.extend(chunk)
|
||||
if len(body) > 4 * 1024 * 1024:
|
||||
return JSONResponse({"detail": "MCP 请求过大"}, status_code=413)
|
||||
# BaseHTTPMiddleware replays cached bytes to the SDK; never log this payload.
|
||||
request._body = bytes(body)
|
||||
try:
|
||||
import json
|
||||
|
||||
message = json.loads(body)
|
||||
except (ValueError, UnicodeDecodeError):
|
||||
return JSONResponse({"detail": "无效 JSON"}, status_code=400)
|
||||
if isinstance(message, dict) and message.get("method") == "tools/call":
|
||||
params = message.get("params")
|
||||
tool = params.get("name") if isinstance(params, dict) else None
|
||||
definition = TOOLS.get(tool) if isinstance(tool, str) else None
|
||||
if definition and definition[2] not in principal.scopes:
|
||||
return JSONResponse({"detail": "MCP 令牌缺少所需权限"}, status_code=403)
|
||||
elif request.method not in ("GET", "HEAD", "OPTIONS"):
|
||||
if request.headers.get("X-WQ-Request") != "1":
|
||||
return JSONResponse({"detail": "缺少请求校验头"}, status_code=403)
|
||||
origin = request.headers.get("Origin")
|
||||
@@ -410,6 +468,9 @@ def create_app(settings=None, wq_client=None, ai_model_factory=None):
|
||||
await notify_job(runner, "retry_job", result)
|
||||
return result
|
||||
|
||||
if mcp_runtime:
|
||||
app.mount("/api/v1/mcp", mcp_runtime.app)
|
||||
app.include_router(mcp_token_router)
|
||||
app.include_router(backtest_router)
|
||||
app.include_router(api)
|
||||
app.include_router(catalog_router)
|
||||
|
||||
Reference in New Issue
Block a user