ai-agent-book 精选快照(<2MB 代码与文档,来自 github.com/bojieli/ai-agent-book)
Build latest book artifacts / build (push) Canceled after 0s
dependency resolution / resolve (3.11) (push) Canceled after 0s
dependency resolution / resolve (3.13) (push) Canceled after 0s
deploy-pages / build (push) Canceled after 0s
deploy-pages / deploy (push) Canceled after 0s
i18n consistency check / check (push) Canceled after 0s
provider adoption tests / test (chapter2/context-compression) (push) Canceled after 0s
provider adoption tests / test (chapter2/prompt-injection) (push) Canceled after 0s
provider adoption tests / test (chapter2/system-hint) (push) Canceled after 0s
provider adoption tests / test (chapter3/log-sanitization) (push) Canceled after 0s
web-search-agent tests / test (push) Canceled after 0s
web-search-agent tests / agentbook (push) Canceled after 0s

This commit is contained in:
2026-08-20 13:12:50 +00:00
commit b119135836
10275 changed files with 3284984 additions and 0 deletions
@@ -0,0 +1,39 @@
FROM python:3.13-slim-bookworm
# Avoid prompts from apt
ENV DEBIAN_FRONTEND=noninteractive
# Configure mirror
ARG ENABLE_MIRROR=false
ENV PIP_INDEX_URL=${ENABLE_MIRROR:+https://mirrors.aliyun.com/pypi/simple}
RUN if [ "$ENABLE_MIRROR" = "true" ]; then \
rm -rfv /etc/apt/sources.list.d/* && \
echo "deb http://ftp.cn.debian.org/debian bookworm main contrib non-free non-free-firmware" > /etc/apt/sources.list && \
echo "deb http://ftp.cn.debian.org/debian bookworm-updates main contrib non-free non-free-firmware" >> /etc/apt/sources.list && \
echo "deb http://ftp.cn.debian.org/debian bookworm-backports main contrib non-free non-free-firmware" >> /etc/apt/sources.list && \
echo "deb http://ftp.cn.debian.org/debian-security bookworm-security main contrib non-free non-free-firmware" >> /etc/apt/sources.list && \
apt-get clean && \
rm -rf /var/lib/apt/lists/*; \
fi
# Install dependencies
RUN apt-get update && apt-get install -y \
curl \
--no-install-recommends \
&& rm -rf /var/lib/apt/lists/*
RUN pip install -U --break-system-packages uv
COPY pyproject.toml /app/mcp_gateway/pyproject.toml
COPY src /app/mcp_gateway/src
WORKDIR /app/mcp_gateway
RUN uv sync --python-preference=only-system
EXPOSE 8000
HEALTHCHECK --interval=30s --timeout=10s --start-period=10s --retries=5 \
CMD curl -f http://localhost:8000/health || exit 1
CMD ["uv", "run", "--no-sync", "-m", "mcp_gateway.main"]
@@ -0,0 +1,28 @@
[project]
name = "mcp-gateway-server"
version = "0.1.0"
description = "MCP Gateway Server"
requires-python = ">=3.12"
dependencies = [
"docker",
"aiohttp",
"pydantic",
"pydantic-settings",
"fastapi[standard]",
"typer",
"aiofiles",
"python-dotenv",
"websockets",
"pyyaml",
"mcp>=1.12.2",
"httpx[http2]",
"pyjwt>=2.10.1",
"redis>=6.4.0",
]
[build-system]
requires = ["setuptools"]
build-backend = "setuptools.build_meta"
[project.scripts]
mcp-gateway = "mcp_gateway.main:app"
@@ -0,0 +1,43 @@
import logging
import traceback
from fastapi import Request
import jwt
from ..utils.common_utils import get_remote_addr
from ..configs import token_secret
logger = logging.getLogger(__name__)
def check_auth(request: Request) -> bool:
payload = get_auth_payload(dict(request.headers))
logger.info(
f"Gateway auth: remote.addr={get_remote_addr(request)}, payload={payload}"
)
return payload is not None
def get_auth_payload(headers: dict) -> str | None:
"""Check if the request is authorized"""
try:
token = headers.get("Authorization")
if not token:
return None
token = token[len("Bearer ") :]
if not token:
return None
payload = decode_token(token)
return payload
except Exception as e:
logger.error(
f"Failed to check auth, remote.addr={get_remote_addr(request)}, request.headers={request.headers} \n{traceback.format_exc()}"
)
return None
def decode_token(token: str) -> str:
"""Decode the token"""
payload = jwt.decode(token, token_secret, algorithms=["HS256"])
return payload
@@ -0,0 +1,12 @@
import os
cluster_name = os.getenv("CLUSTER_NAME", "mcp_gateway")
debug_mode = os.getenv("DEBUG_MODE", "false").lower() == "true"
vnc_auth = os.getenv("VNC_AUTH", "true").lower() == "true"
redis_url = os.getenv("MCP_GATEWAY_REDIS_URL", "redis://localhost:6379/0")
token_secret = os.getenv("MCP_GATEWAY_TOKEN_SECRET")
assert token_secret, "MCP_GATEWAY_TOKEN_SECRET is not set"
@@ -0,0 +1 @@
from .container_server_manager import ContainerServer, ContainerServerManager, container_server_manager_builder
@@ -0,0 +1,300 @@
from typing import List, Tuple, Optional
from pydantic import BaseModel, Field
from ..configs import cluster_name
from ..sessions import SessionId
import logging
import httpx
import traceback
import random
import asyncio
from typing import Optional
import json
from redis.asyncio import Redis
from ..utils.common_utils import check_container_server_health
logger = logging.getLogger(__name__)
class ContainerServer(BaseModel):
"""Container server configuration"""
ip_addr: str = Field(description="IP address of the container server")
port: int = Field(description="Manager port of the container server")
token: Optional[str] = Field(
default=None, description="Token of the container server"
)
cpu_load: Optional[List[float]] = Field(
default=None, description="CPU load of the container server"
)
memory_usage: Optional[List[float]] = Field(
default=None, description="Memory usage of the container server"
)
@property
def server_id(self) -> str:
return f"{self.ip_addr}:{self.port}"
class ContainerServerRepo:
def __init__(self):
self._container_servers: List[ContainerServer] = []
async def initialize(self):
pass
async def get_server(self, container_server_id: str) -> ContainerServer | None:
servers = await self.get_servers()
return next(
(s for s in servers if s.server_id == container_server_id),
None,
)
async def get_servers(self) -> List[ContainerServer]:
return self._container_servers
async def add_server(self, server: ContainerServer):
self._container_servers.append(server)
async def remove_server(self, server: ContainerServer):
self._container_servers.remove(server)
class ContainerServerRedisRepo(ContainerServerRepo):
def __init__(self, redis_url: str):
self._redis_client = Redis.from_url(redis_url)
self._redis_key = f"{cluster_name}.container_servers"
async def initialize(self):
"""Initialize the container server redis repo"""
await self._redis_client.ping()
async def get_servers(self) -> List[ContainerServer]:
"""Get all container servers from Redis"""
try:
# Get all server data from Redis hash
server_data = await self._redis_client.hgetall(self._redis_key)
servers = []
for server_id, server_json in server_data.items():
try:
server_dict = json.loads(server_json)
server = ContainerServer(**server_dict)
servers.append(server)
except (json.JSONDecodeError, ValueError) as e:
logger.warning(f"Failed to deserialize server {server_id}: {e}")
continue
return servers
except Exception as e:
logger.error(f"Failed to get servers from Redis: {e}")
return []
async def add_server(self, server: ContainerServer):
"""Add a container server to Redis"""
try:
server_json = json.dumps(server.model_dump())
await self._redis_client.hset(
self._redis_key, server.server_id, server_json
)
logger.info(
f"Added server: server_id={server.server_id}, server_json={server_json}"
)
except Exception as e:
logger.error(f"Failed to add server {server.server_id} to Redis: {e}")
async def remove_server(self, server: ContainerServer):
"""Remove a container server from Redis"""
try:
await self._redis_client.hdel(self._redis_key, server.server_id)
logger.info(f"Removed server: server_id={server.server_id}")
except Exception as e:
logger.error(f"Failed to remove server {server.server_id} from Redis: {e}")
async def clear_all(self):
"""Clear all container servers from Redis"""
try:
await self._redis_client.delete(self._redis_key)
logger.debug("Cleared all container servers from Redis")
except Exception as e:
logger.error(f"Failed to clear all servers from Redis: {e}")
class ContainerServerManager:
def __init__(
self, container_server_repo: ContainerServerRepo = ContainerServerRepo()
):
self._container_server_repo: ContainerServerRepo = container_server_repo
async def create_container(
self, session_id: SessionId
) -> Tuple[str, str, int, int, str]:
"""
Create a new container for the session id
Return: (container_id, container_ip, container_mcp_port, container_novnc_port, container_server_id)
"""
container_server = await self._select_container_server()
if not container_server:
raise Exception("No container servers available")
async with httpx.AsyncClient() as client:
try:
response = await client.post(
f"http://{container_server.ip_addr}:{container_server.port}/api/container/create",
json={
"token": container_server.token,
"session_id": session_id.session_id,
},
timeout=httpx.Timeout(300.0),
)
response.raise_for_status()
ret = response.json()
logger.info(f"Create container response: {ret}")
if ret.pop("status") == "success":
return (
ret.get("data", {}).get("container_id"),
ret.get("data", {}).get("ip_addr"),
ret.get("data", {}).get("mcp_port"),
ret.get("data", {}).get("novnc_port"),
container_server.server_id,
)
else:
raise Exception(ret.get("message", "Failed to create container"))
except Exception as e:
logger.error(
f"Failed to create container, remote.addr={container_server.ip_addr}: {traceback.format_exc()}"
)
raise
async def shutdown_container(
self, container_server_id: str, container_id: str
) -> bool:
async with httpx.AsyncClient() as client:
try:
container_server = await self._container_server_repo.get_server(
container_server_id
)
assert (
container_server
), f"Container server not found! container_server_id={container_server_id}"
container_server_url = (
f"http://{container_server.ip_addr}:{container_server.port}"
)
response = await client.post(
f"{container_server_url}/api/container/shutdown",
json={
"token": container_server.token,
"container_id": container_id,
},
timeout=httpx.Timeout(60.0),
)
response.raise_for_status()
ret = response.json()
logger.info(f"Shutdown container response: {ret}")
return ret.pop("status") == "success"
except Exception as e:
logger.error(f"Failed to shutdown container! {traceback.format_exc()}")
raise
async def _select_container_server(self) -> Optional[ContainerServer]:
"""Select a container server for the request"""
servers = await self._container_server_repo.get_servers()
if not servers:
logger.error("No container servers available")
return None
# Use the first available backend (can be enhanced with load balancing logic)
server = None
for i in range(10):
server = random.choice(servers)
if await check_container_server_health(ip=server.ip_addr, port=server.port):
return server
else:
logger.warning(
f"Container server {server.ip_addr}:{server.port} is not healthy, retry {i+1}/10"
)
await asyncio.sleep(1)
return None
async def shutdown(self):
"""Shutdown the container server manager"""
pass
async def register_container_server(self, new_server: ContainerServer) -> bool:
"""Register a container server"""
assert new_server, "server is required"
assert (
new_server.ip_addr and new_server.port
), "ip_addr and manager_port are required"
# Check for duplicate using list comprehension
servers = await self._container_server_repo.get_servers()
if any(s.server_id == new_server.server_id for s in servers):
logger.debug(
f"Container server already registered: {new_server.ip_addr}:{new_server.port}"
)
return False
if not await check_container_server_health(
ip=new_server.ip_addr, port=new_server.port
):
logger.info(
f"Container server not healthy: {new_server.ip_addr}:{new_server.port}"
)
return False
# Add new server
await self._container_server_repo.add_server(new_server)
current_servers = await self._container_server_repo.get_servers()
logger.info(
f"Container server register success! current_servers={current_servers}"
)
return True
async def initialize(self):
"""Initialize the container server manager"""
await self._container_server_repo.initialize()
async def health_check_task():
while True:
invalid_servers = []
servers = await self._container_server_repo.get_servers()
for s in servers:
if not await check_container_server_health(
ip=s.ip_addr, port=s.port
):
invalid_servers.append(s)
if invalid_servers:
for s in invalid_servers:
await self._container_server_repo.remove_server(s)
current_servers = await self._container_server_repo.get_servers()
logger.warning(
f"Container server health check failed: invalid_servers={invalid_servers}, current_container_servers={current_servers}"
)
await asyncio.sleep(10)
asyncio.create_task(health_check_task())
async def get_container_servers(self) -> List[ContainerServer]:
"""Get all container servers"""
return await self._container_server_repo.get_servers()
async def container_server_manager_builder():
from ..configs import redis_url
if redis_url:
container_server_manager = ContainerServerManager(
ContainerServerRedisRepo(redis_url)
)
else:
container_server_manager = ContainerServerManager(ContainerServerRepo())
await container_server_manager.initialize()
return container_server_manager
@@ -0,0 +1,96 @@
"""
Streamable HTTP MCP Service Proxy
A gateway service that forwards MCP requests to bound backend MCP services
with persistent HTTP connections. Implements the streamable HTTP protocol
for Model Context Protocol (MCP) communication.
"""
import logging
import traceback
from contextlib import asynccontextmanager
from fastapi import FastAPI, Request, Response, HTTPException
from .utils.common_utils import get_remote_addr, get_mcp_operation
from .auth import check_auth
from .mcp_gateway import MCPGateway
from .routers import gateway_api, novnc_proxy_pass, dashboard
# Configure logging
logging.basicConfig(
level=logging.INFO, format="%(asctime)s - %(name)s - %(levelname)s - %(message)s"
)
logger = logging.getLogger(__name__)
# Suppress httpx INFO logs
logging.getLogger("httpx").setLevel(logging.WARNING)
@asynccontextmanager
async def lifespan(app: FastAPI):
"""FastAPI lifespan context manager for startup and shutdown events"""
# Initialize gateway
gateway = MCPGateway()
await gateway.startup()
app.state.gateway = gateway
try:
yield
finally:
# Shutdown
await gateway.shutdown()
# FastAPI application setup with lifespan
app = FastAPI(
title="MCP",
description="MCP Service",
version="1.0.0",
lifespan=lifespan,
docs_url=None,
redoc_url=None,
)
@app.api_route("/mcp", methods=["GET", "POST", "PUT", "DELETE", "PATCH"])
async def mcp_proxy(request: Request):
"""
Main MCP proxy endpoint.
Accepts all HTTP methods and forwards them to backend MCP servers.
"""
if not check_auth.check_auth(request):
logger.warning(
f"MCP Request unauthorized! remote.addr={get_remote_addr(request)}, request.headers={request.headers}"
)
return Response(status_code=401, content="Unauthorized")
try:
gateway: MCPGateway = request.app.state.gateway
return await gateway.handle_mcp_request(request)
except Exception as e:
logger.error(
f"Request error: remote.addr={get_remote_addr(request)}, request.headers={request.headers}\n{traceback.format_exc()}"
)
raise HTTPException(status_code=500, detail="Internal server error, {e}")
# NoVNC proxy pass
app.include_router(novnc_proxy_pass.router, prefix="/novnc")
# Gateway rest api
app.include_router(gateway_api.router, prefix="/api")
# Gateway dashboard
app.include_router(dashboard.router, prefix="/dashboard")
@app.get("/health")
async def health(request: Request):
return {"status": "success", "message": "MCP Gateway is healthy"}
if __name__ == "__main__":
import uvicorn
# Run the server
uvicorn.run(app, host="0.0.0.0", port=8000)
@@ -0,0 +1,88 @@
from abc import ABC
import logging
from fastapi import Request, Response
from .sessions import SessionId
from .utils.common_utils import (
get_mcp_operation,
get_remote_addr,
)
from .sessions import (
SessionConnectionManager,
session_connection_manager_builder,
)
from .containers import (
ContainerServerManager,
container_server_manager_builder,
)
logger = logging.getLogger(__name__)
class MCPGateway(ABC):
"""
MCP Gateway service that proxies requests to backend MCP servers
with persistent connections based on client connection caching.
"""
def __init__(self):
self.session_connection_manager: SessionConnectionManager
self.container_server_manager: ContainerServerManager
async def startup(self):
"""Initialize the gateway service"""
logger.info("Starting MCP Gateway service")
self.container_server_manager = await container_server_manager_builder()
self.session_connection_manager = await session_connection_manager_builder(
self.container_server_manager
)
async def shutdown(self):
"""Cleanup resources"""
logger.info("Shutting down MCP Gateway service")
await self.session_connection_manager.shutdown()
await self.container_server_manager.shutdown()
async def handle_mcp_request(self, request: Request) -> Response:
"""
Handle incoming MCP requests and forward them to appropriate backend servers.
Uses connection-based caching for persistent connections.
"""
remote_addr = get_remote_addr(request)
session_id = SessionId.from_request(request)
http_method, mcp_client_method, mcp_tool_method = await get_mcp_operation(
request
)
# Handle session initialize request
if (
http_method == "POST"
and not session_id.mcp_session_id
and mcp_client_method == "initialize"
):
return await self.session_connection_manager.handle_initialize_request(
request
)
# Handle seesion finalize request
if http_method == "DELETE" and session_id.mcp_session_id:
return await self.session_connection_manager.handle_delete_request(
request, session_id
)
# Handle other requests
assert session_id.mcp_session_id, "Mcp-Session-Id is required"
response = None
try:
response = await self.session_connection_manager.forward_request(
request, session_id
)
return response
finally:
logger.info(
f"MCP client request: request.addr={remote_addr}, session_id={session_id}, mcp_operation={http_method, mcp_client_method, mcp_tool_method}, response.status_code={response.status_code if response else 'None'}, response.headers={response.headers if response else 'None'}"
)
@@ -0,0 +1,44 @@
import json
import logging
import traceback
from fastapi import APIRouter, Request
from fastapi.responses import JSONResponse, Response
from ..utils.common_utils import get_remote_addr
router = APIRouter()
logger = logging.getLogger(__name__)
@router.get("/status")
async def dashboard(request: Request):
"""Show dashboard"""
try:
gateway = request.app.state.gateway
sessions = await gateway.session_connection_manager.get_sessions()
container_servers = (
await gateway.container_server_manager.get_container_servers()
)
status = {
"status": "success",
"data": {
"sessions": [session.model_dump() for session in sessions],
"container_servers": [
server.model_dump(exclude={"token"})
for server in container_servers
],
},
}
return Response(
content=json.dumps(status, ensure_ascii=False, indent=2).encode("utf-8"),
status_code=200,
headers={"Content-Type": "application/json; charset=utf-8"},
)
except Exception as e:
logger.error(
f"Failed to get dashboard, remote.addr={get_remote_addr(request)}\n{traceback.format_exc()}"
)
return JSONResponse(
content={"status": "error", "message": "Failed to get dashboard"},
status_code=500,
)
@@ -0,0 +1,25 @@
import logging
import traceback
from fastapi import APIRouter, Request
from ..utils.common_utils import get_remote_addr
router = APIRouter()
logger = logging.getLogger(__name__)
from ..containers import ContainerServer
from ..mcp_gateway import MCPGateway
@router.post("/container_server/register")
async def container_server_register(request: Request, body: dict):
"""Register a container server with deduplication"""
try:
gateway: MCPGateway = request.app.state.gateway
new_server = ContainerServer(**body)
await gateway.container_server_manager.register_container_server(new_server)
except Exception as e:
logger.error(
f"Failed to register container server, remote.addr={get_remote_addr(request)}\n{traceback.format_exc()}"
)
return {"status": "error", "message": "Failed to register container server"}
@@ -0,0 +1,105 @@
import asyncio
import logging
import websockets
from fastapi import APIRouter, Request, HTTPException, Response, WebSocket
from ..configs import vnc_auth
from ..sessions.session_connection import SessionId
from ..auth import check_auth
from ..utils.common_utils import get_remote_addr
logger = logging.getLogger(__name__)
router = APIRouter()
@router.api_route(
"/{mcp_session_id}/{full_path:path}",
methods=["GET", "POST", "PUT", "DELETE", "PATCH"],
)
async def novnc_proxy(request: Request, mcp_session_id: str):
"""
Proxy for the VNC server.
"""
# Auth check
if vnc_auth and not check_auth.check_auth(request):
logger.warning(
f"Request unauthorized! remote.addr={get_remote_addr(request)}, request.headers={request.headers}"
)
return Response(status_code=401, content="Unauthorized")
gateway = request.app.state.gateway
session_connection = await gateway.session_connection_manager.get_session(
SessionId(mcp_session_id=mcp_session_id)
)
if not session_connection:
logger.warning(
f"Session is invalid or expired: remote.addr={get_remote_addr(request)}, request.headers={request.headers}, session_id={mcp_session_id}"
)
raise HTTPException(status_code=400, detail="Session is invalid or expired")
return await session_connection.novnc_proxy(request)
@router.websocket("/{mcp_session_id}/websockify")
async def websocket_novnc_proxy(websocket: WebSocket, mcp_session_id: str):
"""WebSocket proxy for noVNC websockify connections"""
await websocket.accept()
# Auth check
if vnc_auth and not check_auth.get_auth_payload(dict(websocket.headers)):
logger.warning(
f"Request unauthorized! request.url={websocket.url}, request.headers={websocket.headers}"
)
await websocket.close(code=4000, reason="Unauthorized")
return
# Get session connection
gateway = websocket.app.state.gateway
session_connection = await gateway.session_connection_manager.get_session(
SessionId(mcp_session_id=mcp_session_id)
)
if not session_connection:
logger.warning(
f"WebSocket: Session invalid or expired: session_id={mcp_session_id}"
)
await websocket.close(code=4000, reason="Session invalid or expired")
return
# Proxy WebSocket to backend container
backend_ws_url = f"ws://{session_connection.container_ip_addr}:{session_connection.novnc_port}/websockify"
logger.info(f"WebSocket: Proxying to {backend_ws_url}")
async def _proxy_websocket(client_ws, backend_ws):
"""Simple bidirectional WebSocket proxy"""
async def client_to_backend():
while True:
try:
message = await client_ws.receive_bytes()
except:
try:
message = await client_ws.receive_text()
except:
break
await backend_ws.send(message)
async def backend_to_client():
async for message in backend_ws:
if isinstance(message, bytes):
await client_ws.send_bytes(message)
else:
await client_ws.send_text(message)
await asyncio.gather(
client_to_backend(), backend_to_client(), return_exceptions=True
)
try:
async with websockets.connect(backend_ws_url) as backend_ws:
await _proxy_websocket(websocket, backend_ws)
except Exception as e:
logger.error(f"WebSocket proxy error: {e}")
@@ -0,0 +1,5 @@
from .session_connection import SessionId, VpcSession
from .session_connection_manager import (
SessionConnectionManager,
session_connection_manager_builder,
)
@@ -0,0 +1,105 @@
import asyncio
import logging
import re
import traceback
from typing import List, Optional
from fastapi import Request
from fastapi.responses import StreamingResponse
import httpx
from pydantic import BaseModel, Field
from ..utils.proxy_utils import proxy_pass_bytes, proxy_pass_lines
logger = logging.getLogger(__name__)
class SessionId(BaseModel):
session_id: Optional[str] = Field(
default=None,
description="Session ID from mcp client, used for multi step session affinity",
)
mcp_session_id: Optional[str] = Field(
default=None,
description="MCP Session ID for mcp client session",
)
@classmethod
def from_request(cls, request: Request) -> "SessionId":
"""Create SessionId from request headers"""
return cls(
session_id=request.headers.get("SESSION_ID"),
mcp_session_id=request.headers.get("Mcp-Session-Id"),
)
class VpcSession(BaseModel):
"""Represents a persistent HTTP connection to a backend server"""
container_id: str = Field(description="Container ID")
container_ip_addr: str = Field(description="Container IP")
mcp_port: int = Field(description="MCP Port")
novnc_port: int = Field(description="NoVNC Port")
container_server_id: str = Field(description="Container Server ID")
mcp_session_ids: List[str] = Field(default=[], description="MCP Session IDs")
session_ids: List[str] = Field(default=[], description="Session IDs")
def is_bind(self, session_id: SessionId) -> bool:
"""Check if the session id matches the session connection"""
return (
session_id.mcp_session_id in self.mcp_session_ids
or session_id.session_id in self.session_ids
)
def _get_client(self):
"""Establish connection to backend server"""
return httpx.AsyncClient(
timeout=httpx.Timeout(360, connect=30, pool=30),
limits=httpx.Limits(
max_keepalive_connections=30,
max_connections=30,
keepalive_expiry=600.0, # Keep alive for 1 hour
),
http2=True, # Enable HTTP/2 for better multiplexing
)
async def forward_request(
self, method: str, headers: dict, content: bytes
) -> StreamingResponse:
"""Send request through the persistent connection with true streaming proxy"""
try:
client = self._get_client()
await client.__aenter__()
return await proxy_pass_lines(
client=client,
method=method,
url=f"http://{self.container_ip_addr}:{self.mcp_port}/mcp",
headers=headers,
content=content,
)
except Exception as e:
logger.error(
f"Connection error to {self.container_ip_addr}:{self.mcp_port}: {e}"
)
raise
async def novnc_proxy(self, request: Request) -> StreamingResponse:
"""
Proxy the request to the VNC server.
"""
try:
client = self._get_client()
await client.__aenter__()
# Remove the prefix "/novnc/{session_id}" from the URL path
target_url = re.sub(r"^/novnc/[^/]+", "", request.url.path) or "/"
return await proxy_pass_bytes(
client=client,
method=request.method,
url=f"http://{self.container_ip_addr}:{self.novnc_port}{target_url}",
headers=dict(request.headers),
content=await request.body(),
)
except:
logger.error(
f"Connection error to {self.container_ip_addr}:{self.mcp_port}: {traceback.format_exc()}"
)
raise
@@ -0,0 +1,295 @@
import asyncio
import json
import logging
import time
import traceback
from typing import List
from fastapi import Request, Response
import redis.asyncio as redis
from ..containers import ContainerServerManager
from ..utils.common_utils import get_remote_addr
from ..sessions import VpcSession, SessionId
from ..configs import cluster_name, debug_mode, redis_url
logger = logging.getLogger(__name__)
class SessionRepo:
def __init__(self):
self.sessions: List[VpcSession] = []
async def initialize(self):
pass
async def get_bind_session(self, session_id: SessionId) -> VpcSession | None:
return next(
(v for v in await self.get_sessions() if v.is_bind(session_id)),
None,
)
async def get_sessions(self) -> List[VpcSession]:
return self.sessions
async def update_vpc_session(self, vpc_session: VpcSession):
pass
async def remove_vpc_session(self, vpc_session: VpcSession):
self.sessions.remove(vpc_session)
class SessionRedisRepo(SessionRepo):
def __init__(self, redis_url: str):
self._redis_client = redis.Redis.from_url(redis_url)
self._sessions_key = f"{cluster_name}.vpc_sessions"
async def initialize(self):
await self._redis_client.ping()
async def get_sessions(self) -> List[VpcSession]:
"""Get all VPC sessions from Redis"""
sessions = []
try:
session_data = await self._redis_client.hgetall(self._sessions_key)
for _, session_json in session_data.items():
try:
if session_json:
session_data_str = session_json.decode("utf-8")
sessions.append(self._deserialize_vpc_session(session_data_str))
except Exception as e:
logger.error(
f"Error deserializing VPC session: {traceback.format_exc()}"
)
except Exception as e:
logger.error(f"Error getting sessions: {traceback.format_exc()}")
return sessions
async def update_vpc_session(self, vpc_session: VpcSession):
"""Update VPC session in Redis"""
try:
session_json = self._serialize_vpc_session(vpc_session)
await self._redis_client.hset(
self._sessions_key, vpc_session.container_id, session_json
)
except Exception as e:
logger.error(f"Error updating VPC session: {traceback.format_exc()}")
async def remove_vpc_session(self, vpc_session: VpcSession):
"""Remove VPC session from Redis"""
try:
await self._redis_client.hdel(self._sessions_key, vpc_session.container_id)
except Exception as e:
logger.error(f"Error removing VPC session: {e}")
def _serialize_vpc_session(self, vpc_session: VpcSession) -> str:
"""Serialize VpcSession to JSON string"""
session_data = {
"container_id": vpc_session.container_id,
"container_ip_addr": vpc_session.container_ip_addr,
"mcp_port": vpc_session.mcp_port,
"novnc_port": vpc_session.novnc_port,
"container_server_id": vpc_session.container_server_id,
"mcp_session_ids": vpc_session.mcp_session_ids,
"session_ids": vpc_session.session_ids,
}
return json.dumps(session_data)
def _deserialize_vpc_session(self, session_json: str) -> VpcSession:
"""Deserialize JSON string to VpcSession"""
session_data = json.loads(session_json)
vpc_session = VpcSession(
container_id=session_data.get("container_id", ""),
container_ip_addr=session_data.get("container_ip_addr", ""),
mcp_port=session_data.get("mcp_port", -1),
novnc_port=session_data.get("novnc_port", -1),
container_server_id=session_data.get("container_server_id", ""),
)
vpc_session.mcp_session_ids = session_data.get("mcp_session_ids", [])
vpc_session.session_ids = session_data.get("session_ids", [])
return vpc_session
class SessionConnectionManager:
"""Manages session connections to container servers"""
def __init__(
self,
container_server_manager: ContainerServerManager,
session_repo: SessionRepo = SessionRepo(),
):
self._session_repo: SessionRepo = session_repo
self.container_server_manager: ContainerServerManager = container_server_manager
async def initialize(self):
"""Initialize the session connection manager"""
await self._session_repo.initialize()
async def get_session(self, session_id: SessionId) -> VpcSession | None:
return await self._session_repo.get_bind_session(session_id)
async def create_vpc_session(self, session_id: SessionId) -> VpcSession:
"""Get the session connection"""
vpc_session = await self.get_session(session_id)
if not vpc_session:
logger.info(f"Create new vpc session: session_id={session_id}")
(
container_id,
container_ip_addr,
container_mcp_port,
container_novnc_port,
container_server_id,
) = await self.container_server_manager.create_container(session_id)
vpc_session = VpcSession(
container_id=container_id,
container_ip_addr=container_ip_addr,
mcp_port=container_mcp_port,
novnc_port=container_novnc_port,
container_server_id=container_server_id,
)
return vpc_session
async def forward_request(
self,
request: Request,
session_id: SessionId,
) -> Response:
vpc_session = await self.get_session(session_id)
assert vpc_session, f"Vpc session not found! session_id={session_id}"
return await self._forward_request(request, vpc_session)
async def _forward_request(
self, request: Request, vpc_session: VpcSession
) -> Response:
"""Forward the request through the persistent connection"""
try:
headers = dict(request.headers)
# Remove headers that should not be forwarded or can cause conflicts
for header in [
"host",
"connection",
"upgrade",
]:
headers.pop(header, None)
# Forward request through connection
backend_response = await vpc_session.forward_request(
method=request.method, headers=headers, content=await request.body()
)
return backend_response
except Exception as e:
logger.error(
f"Unexpected error forwarding request: {traceback.format_exc()}"
)
raise e
async def handle_initialize_request(self, request: Request) -> Response:
# Handle session init
logger.info(
f"MCP client session initialize request: remote.addr={get_remote_addr(request)}, request.headers={request.headers}"
)
time_start = time.time()
session_id = SessionId.from_request(request)
vpc_session = None
if debug_mode:
vpc_session = VpcSession(
container_server_id="http://mcp-server-debug:4242",
container_ip_addr="mcp-server-debug",
mcp_port=4242,
novnc_port=5901,
container_id="mcp-server-debug",
)
logger.info(
f"Gateway debug mode, skip container creation, default mcp_server connection={vpc_session}"
)
else:
vpc_session = await self.create_vpc_session(session_id)
# Forward request through the new persistent connection
response = await self._forward_request(request, vpc_session)
mcp_session_id = response.headers.get("Mcp-Session-Id")
assert (
mcp_session_id
), "Mcp session id is required in initialize response header!"
session_id.mcp_session_id = mcp_session_id
# Bind session to VPC
if (
session_id.mcp_session_id
and session_id.mcp_session_id not in vpc_session.mcp_session_ids
):
vpc_session.mcp_session_ids.append(session_id.mcp_session_id)
if (
session_id.session_id
and session_id.session_id not in vpc_session.session_ids
):
vpc_session.session_ids.append(session_id.session_id)
await self._session_repo.update_vpc_session(vpc_session)
logger.info(
f"Client session initialize success: session_id={session_id}, vpc_session.container_id={vpc_session.container_id}, time_cost_sec={time.time() - time_start}"
)
return response
async def handle_delete_request(
self, request: Request, session_id: SessionId
) -> Response:
"""Handle session delete request"""
logger.info(
f"McpClient session release request: {get_remote_addr(request)}, session_id={session_id}"
)
mcp_response = await self.forward_request(request, session_id)
await self._release_mcp_session(session_id)
return mcp_response
async def _release_mcp_session(self, session_id: SessionId):
"""Release the mcp session"""
vpc_session = await self._session_repo.get_bind_session(session_id)
assert vpc_session, "Mcp session not found"
assert session_id.mcp_session_id, "Mcp session id is required"
assert vpc_session.is_bind(session_id), "Mcp session id not bind to vpc_session"
vpc_session.mcp_session_ids.remove(session_id.mcp_session_id)
if session_id.session_id:
# Release MCP Session only
vpc_session.mcp_session_ids.remove(session_id.mcp_session_id)
await self._session_repo.update_vpc_session(vpc_session)
else:
await self._session_repo.remove_vpc_session(vpc_session)
asyncio.create_task(
self.container_server_manager.shutdown_container(
container_server_id=vpc_session.container_server_id,
container_id=vpc_session.container_id,
)
)
async def shutdown(self):
"""Shutdown the session connection manager"""
pass
async def get_sessions(self) -> List[VpcSession]:
"""Get all sessions"""
return await self._session_repo.get_sessions()
async def session_connection_manager_builder(
container_server_manager: ContainerServerManager,
):
if redis_url:
session_connection_manager = SessionConnectionManager(
container_server_manager, SessionRedisRepo(redis_url)
)
else:
session_connection_manager = SessionConnectionManager(container_server_manager)
await session_connection_manager.initialize()
return session_connection_manager
@@ -0,0 +1,69 @@
import logging
import traceback
from typing import Callable, Tuple
from fastapi import Request
import httpx
logger = logging.getLogger(__name__)
def get_remote_addr(request: Request) -> str:
"""Get the request remote address"""
client_ip = request.client.host
try:
x_forwarded_for = request.headers.get("X-Forwarded-For")
if x_forwarded_for:
client_ip = x_forwarded_for.split(",")[0].strip()
except:
logger.error(f"Failed to get request remote addr: {traceback.format_exc()}")
return client_ip
async def get_mcp_operation(request: Request) -> Tuple[str, str | None, str | None]:
"""Get the MCP operation from the request
Returns:
Tuple[str, str | None, str | None]: (http_method, mcp_client_method, mcp_tool_method)
"""
try:
if request.method in ["GET", "DELETE"]:
return request.method, None, None
jsonrpc_body = await request.json()
mcp_client_method = jsonrpc_body.get("method")
mcp_tool_method = jsonrpc_body.get("params", {}).get("name")
return request.method, mcp_client_method, mcp_tool_method
except Exception as e:
logger.warning(
f"Failed to get mcp operation, remote.addr={get_remote_addr(request)} http_method={request.method}, {e}"
)
return request.method, None, None
async def check_container_server_health(ip: str, port: int) -> bool:
"""Check if the container server is healthy"""
return await check_server_health(
ip=ip,
port=port,
checker=lambda body: body["status"] == "success"
and body["message"] == "Container server is healthy",
)
async def check_server_health(
ip: str,
port: int,
timeout: float = 3.0,
checker: Callable[[dict], bool] = lambda body: body["status"] == "success",
) -> bool:
"""Check if the server is healthy"""
try:
async with httpx.AsyncClient(timeout=timeout) as client:
response = await client.get(f"http://{ip}:{port}/health")
response.raise_for_status()
body = response.json()
return checker(body)
except Exception as e:
logger.error(
f"Failed to check server health: {ip}:{port}, {e}"
)
return False
@@ -0,0 +1,68 @@
import logging
from fastapi.responses import StreamingResponse
import httpx
logger = logging.getLogger(__name__)
async def proxy_pass_lines(
client: httpx.AsyncClient, method: str, url: str, headers: dict, content: bytes
) -> StreamingResponse:
stream_response_context = client.stream(
method=method,
url=url,
headers=headers,
content=content,
)
stream_response = await stream_response_context.__aenter__()
content_type = stream_response.headers.get("content-type", "")
response_headers = dict(stream_response.headers)
status_code = stream_response.status_code
async def stream_lines():
try:
async for line in stream_response.aiter_lines():
yield f"{line}\n"
finally:
await stream_response_context.__aexit__(None, None, None)
await client.__aexit__(None, None, None)
return StreamingResponse(
content=stream_lines(),
status_code=status_code,
headers=response_headers,
media_type=content_type,
)
async def proxy_pass_bytes(
client: httpx.AsyncClient, method: str, url: str, headers: dict, content: bytes
) -> StreamingResponse:
stream_response_context = client.stream(
method=method,
url=url,
headers=headers,
content=content,
)
stream_response = await stream_response_context.__aenter__()
content_type = stream_response.headers.get("content-type", "")
response_headers = dict(stream_response.headers)
status_code = stream_response.status_code
async def stream_bytes():
try:
async for chunk in stream_response.aiter_raw(chunk_size=1024):
yield chunk
finally:
await stream_response_context.__aexit__(None, None, None)
return StreamingResponse(
content=stream_bytes(),
status_code=status_code,
headers=response_headers,
media_type=content_type,
)
@@ -0,0 +1,18 @@
import time
import jwt
def gen_auth_token(root_token: str, app: str):
pay_load = {"app": app, "version": 1, "time": time.time()}
token = jwt.encode(payload=pay_load, key=root_token, algorithm="HS256")
return token
def test_gen_token():
root_token = "123321"
token = gen_auth_token(root_token, "local_debug")
print(token)
if __name__ == "__main__":
test_gen_token()
@@ -0,0 +1,13 @@
import os, time, jwt
def gen_auth_token(app: str = "agiopenwebui-vnc-proxy"):
# novnc_server_secret = "123321"
novnc_server_secret = "AwOrld@0DF0-41F9-4d47-9730-35F706B76045@20250820"
pay_load = {"app": app, "version": 1, "time": time.time()}
token = jwt.encode(payload=pay_load, key=novnc_server_secret, algorithm="HS256")
return token
token = gen_auth_token(app="aworldcore-agent")
print(token)
File diff suppressed because it is too large Load Diff