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,29 @@
FROM docker:dind
# Configure mirror
ARG ENABLE_MIRROR=false
ENV PIP_INDEX_URL=${ENABLE_MIRROR:+https://mirrors.aliyun.com/pypi/simple}
RUN apk add --no-cache python3 py3-pip curl tini
COPY --from=ghcr.io/astral-sh/uv:latest /uv /uvx /bin/
WORKDIR /app/container_server
COPY container_server/pyproject.toml /app/container_server/pyproject.toml
COPY container_server/src /app/container_server/src
RUN uv sync --python-preference=only-system
EXPOSE 8000
HEALTHCHECK --interval=30s --timeout=10s --start-period=60s --retries=5 \
CMD curl -f http://localhost:8000/health || exit 1
ENV DOCKER_MODE=dind
ENTRYPOINT ["tini", "--"]
CMD ["sh", "-c", "(dockerd-entrypoint.sh) & (sleep 5 && uv run --no-sync -m container_server.main) & wait -n"]
@@ -0,0 +1,20 @@
[project]
name = "container-server"
version = "0.1.0"
description = "Container Server"
requires-python = ">=3.12"
dependencies = [
"docker",
"aiohttp",
"pydantic",
"pydantic-settings",
"fastapi[standard]",
"python-dotenv",
"websockets",
"pyyaml",
"mcp>=1.12.2",
]
[build-system]
requires = ["setuptools"]
build-backend = "setuptools.build_meta"
@@ -0,0 +1,17 @@
import os
debug_mode = os.getenv("DEBUG_MODE", "false").lower() == "true"
container_server_port = int(os.getenv("CONTAINER_SERVER_PORT", "9000"))
docker_registry_url = os.getenv("DOCKER_REGISTRY_URL")
docker_registry_user_name = os.getenv("DOCKER_REGISTRY_USER_NAME")
docker_registry_password = os.getenv("DOCKER_REGISTRY_PASSWORD")
gateway_server_addr = os.getenv("GATEWAY_SERVER_ADDR", "http://mcp-gateway:8000")
mcp_server_image_id = os.getenv(
"VIRTUALPC_MCP_SERVER_IMAGE_ID",
"aworld-registry-registry-vpc.ap-southeast-1.cr.aliyuncs.com/aworld/mcp-server",
)
docker_mode = os.getenv("DOCKER_MODE", "dind")
@@ -0,0 +1,231 @@
from functools import cache
import socket
import logging
import traceback
import asyncio
import uuid
import httpx
import threading
from .dockers import docker_helper
from .configs import (
container_server_port,
docker_registry_url,
docker_registry_user_name,
docker_registry_password,
mcp_server_image_id,
gateway_server_addr,
docker_mode,
debug_mode,
)
logger = logging.getLogger(__name__)
token = str(uuid.uuid4())
async def wait_docker_ready(timeout: int = 30):
for i in range(timeout):
try:
cmd = ["docker", "ps"]
p = await asyncio.subprocess.create_subprocess_exec(
*cmd, stdout=asyncio.subprocess.PIPE, stderr=asyncio.subprocess.STDOUT
)
stdout, _ = await p.communicate()
if p.returncode == 0:
logger.info(f"Docker daemon is ready! \n{stdout.decode()}")
return
else:
logger.warning(f"Docker daemon is not ready! {stdout.decode()}")
except:
logger.error(f"Docker daemon is not ready! {traceback.format_exc()}")
await asyncio.sleep(1)
else:
logger.error(f"Docker daemon is not ready after {timeout} seconds!")
raise Exception(f"Docker daemon is not ready after {timeout} seconds!")
async def start_container_server_register_task():
register_url = f"{gateway_server_addr}/api/container_server/register"
local_ip_addr = get_local_ip()
async def register():
try:
async with httpx.AsyncClient() as client:
response = await client.post(
register_url,
json={
"token": token,
"ip_addr": local_ip_addr,
"port": container_server_port,
"cpu_load": [],
"memory_usage": [],
},
timeout=10.0,
)
response.raise_for_status()
except Exception as e:
logger.error(f"Register container server failed: {register_url}, {e}")
raise
async def init_register():
for _ in range(20):
try:
await register()
logger.info(f"Init register success")
break
except:
await asyncio.sleep(3)
else:
logger.error(f"Init register failed after 20 times!")
raise Exception(f"Init register failed after 20 times!")
await init_register()
async def update_register():
while True:
try:
await register()
except:
logger.error(f"Update register failed: {traceback.format_exc()}")
await asyncio.sleep(10)
asyncio.create_task(update_register())
async def load_mcp_server_image():
if docker_registry_url and docker_registry_user_name and docker_registry_password:
await docker_helper.login_async(
registry_url=docker_registry_url,
username=docker_registry_user_name,
password=docker_registry_password,
)
if not debug_mode:
await docker_helper.pull_async(mcp_server_image_id)
async def start_mcp_server_life_cycle_manager():
pass
async def clean_mcp_server_container():
pass
async def create_mcp_server_container():
mcp_port = docker_helper.get_available_port()
novnc_port = docker_helper.get_available_port()
ip_addr = get_local_ip()
container_name = f"mcp_server_{str(uuid.uuid4()).replace('-', '')}"
logger.info(
f"Create mcp server container: {container_name}, image_id: {mcp_server_image_id}, mcp_port: {mcp_port}, novnc_port: {novnc_port}"
)
try:
ports = {4242: f"{mcp_port}", 5901: f"{novnc_port}"}
network = ""
if docker_mode == "host":
network = "visualvirtualpc_virtualpc-network"
container = await docker_helper.run_async(
image_id=mcp_server_image_id,
container_name=container_name,
ports=ports,
network=network,
)
logger.info(
f"Create mcp server container success, waiting for Ready: {container.id}, ip_addr: {ip_addr}, mcp_port: {mcp_port}, novnc_port: {novnc_port}"
)
def tail_logs():
logs = container.logs(stream=True, tail=100, follow=True)
try:
buffer = []
for line in logs:
buffer.append(line.decode())
if len(buffer) >= 20:
logger.info(f"VPC[{container.name}] >>> \n{'> '.join(buffer)}\n")
buffer.clear()
if buffer:
logger.info(f"VPC[{container.name}] >>> \n{'> '.join(buffer)}\n")
logger.info(f"VPC [{container.name}] logs end!")
except Exception as e:
logger.error(f"Error in tail_logs: {e}")
# Start log tailing in background thread
log_thread = threading.Thread(
target=tail_logs, name=f"VPC_{container.name}_logs", daemon=True
)
log_thread.start()
async def health_check(timeout: float = 3.0):
try:
# async with httpx.AsyncClient() as client:
# response = await client.get(
# f"http://{ip_addr}:{mcp_port}/health",
# timeout=httpx.Timeout(timeout),
# )
# response.raise_for_status()
# return True
return await docker_helper.check_health(container.id)
except Exception as e:
logger.error(f"Check mcp server health error! {e}")
return False
max_check = 30
for i in range(max_check):
if await health_check():
logger.info(
f"MCP server {ip_addr}:{mcp_port} is ready: {i+1}/{max_check}"
)
break
else:
logger.warning(
f"MCP server {ip_addr}:{mcp_port} is not ready: {i+1}/{max_check}"
)
await asyncio.sleep(3)
else:
logger.error(
f"MCP server {ip_addr}:{mcp_port} is not ready after {max_check} times!"
)
raise Exception(
f"MCP server {ip_addr}:{mcp_port} is not ready after {max_check} times!"
)
if docker_mode == "host":
ip_addr = await docker_helper.get_container_ip(container.id)
mcp_port = 4242
novnc_port = 5901
return container.id, ip_addr, mcp_port, novnc_port
except:
logger.error(f"Create mcp server container failed: {traceback.format_exc()}")
raise
async def shutdown_mcp_server_container(container_id: str):
try:
await docker_helper.stop_async(container_id)
except:
logger.error(f"Shutdown mcp server container failed: {traceback.format_exc()}")
raise
@cache
def get_local_ip() -> str | None:
try:
host_name = socket.gethostname()
_, _, ip_list = socket.gethostbyname_ex(host_name)
for ip in ip_list:
if not ip.startswith("127."):
return ip
except Exception as e:
logger.error(f"Get local ip failed: {traceback.format_exc()}")
raise RuntimeError("Get local ip failed")
@@ -0,0 +1,229 @@
import asyncio
import traceback
from typing import Tuple
import docker
import logging
import socket
logger = logging.getLogger(__name__)
client = docker.from_env(timeout=600)
async def login_async(registry_url: str, username: str, password: str):
try:
return await asyncio.to_thread(login, registry_url, username, password)
except Exception as e:
logger.error(f"Error in login_async: {e}")
raise
async def pull_async(image_id: str):
try:
return await asyncio.to_thread(pull, image_id)
except Exception as e:
logger.error(f"Error in pull_async: {e}")
raise
async def run_async(
image_id: str,
container_name: str,
ports: dict[int, str] = {},
network: str = "",
volumes: dict[str, str] = {},
environments: dict[str, str] = {},
):
try:
return await asyncio.to_thread(
run, image_id, container_name, ports, network, volumes, environments
)
except Exception as e:
logger.error(f"Error in run_async: {e}")
raise
async def exec_async(container_id: str, cmd: list[str]) -> Tuple[int, str]:
try:
return await asyncio.to_thread(exec, container_id, cmd)
except Exception as e:
logger.error(f"Error in exec_async: {e}")
raise
async def stop_async(container_id: str):
try:
return await asyncio.to_thread(stop, container_id)
except Exception as e:
logger.error(f"Error in stop_async: {e}")
raise
async def check_health(container_id: str):
try:
container = client.containers.get(container_id)
container.reload()
return container.health == "healthy"
except Exception as e:
logger.error(f"Error in check_health: {e}")
return False
async def get_container_ip(container_id: str):
try:
container = client.containers.get(container_id)
container.reload()
nets = container.attrs["NetworkSettings"]["Networks"]
net = list(nets.values())[0]
return net["IPAddress"]
except Exception as e:
logger.error(f"Error in get_container_ip: {e}")
return None
async def build_image_async(image_id: str, dockerfile: str, context_path: str):
async def build():
try:
cmd = [
"docker",
"build",
"--platform",
"linux/amd64",
"-t",
image_id,
"-f",
dockerfile,
context_path,
]
p = await asyncio.subprocess.create_subprocess_exec(
*cmd,
cwd=context_path,
stdout=asyncio.subprocess.PIPE,
stderr=asyncio.subprocess.STDOUT,
)
assert p.stdout is not None
while True:
line = await p.stdout.readline()
if not line:
break
logger.info(line.decode(errors="ignore").rstrip())
rc = await p.wait()
if rc == 0:
logger.info("Build mcp server image success!")
return
else:
raise Exception(f"Build mcp server image error! return code: {rc}")
except Exception as e:
logger.error(f"Build mcp server image error! {e}")
raise
for _ in range(3):
try:
await build()
break
except Exception as e:
await asyncio.sleep(1)
else:
logger.info(f"Build mcp server image failed after 3 times!")
raise Exception("Build mcp server image failed after 3 times!")
def pull(image_id: str):
logger.info(f"Pulling image {image_id}")
try:
if client.images.get(image_id):
logger.info(f"Image {image_id} already exists, skipping pull")
return
img = client.images.pull(image_id)
logger.info(f"Pulled image {image_id}, {img}")
except:
logger.error(f"Failed to pull image {image_id}\n{traceback.format_exc()}")
raise
def run(
image_id: str,
container_name: str,
ports: dict[int, str] = {},
network: str = "",
volumes: dict[str, str] = {},
environments: dict[str, str] = {},
):
logger.info(
f"Creating container {container_name} with args: {{'name': {container_name}, 'image': {image_id}, 'ports': {ports}, 'volumes': {volumes}, 'environments': {environments}}}"
)
try:
container = client.containers.run(
name=container_name,
image=image_id,
detach=True,
auto_remove=True,
ports=ports,
network=network,
volumes=volumes,
environment=environments,
cpu_period=100000,
cpu_quota=90000,
mem_limit="2G",
)
logger.info(f"Created container {container_name} response: {container}")
return container
except:
logger.error(
f"Failed to create container {container_name} with image {image_id}\n{traceback.format_exc()}"
)
raise
def exec(container_id: str, cmd: list[str]) -> Tuple[int, str]:
logger.info(f"Executing command {cmd} on container {container_id}")
try:
container = client.containers.get(container_id)
exit_code, output = container.exec_run(cmd)
logger.info(
f"Command {cmd} executed on container {container_id} with result: {exit_code} {output}"
)
return exit_code, output
except:
logger.error(
f"Failed to execute command {cmd} on container {container_id}\n{traceback.format_exc()}"
)
raise
def stop(container_id: str):
logger.info(f"Stop container {container_id}")
try:
container = client.containers.get(container_id)
container.stop()
logger.info(f"Stopped container {container_id}")
except:
logger.error(
f"Failed to stop container {container_id}\n{traceback.format_exc()}"
)
raise
def login(registry_url: str, username: str, password: str):
logger.info(f"Logging in to {registry_url} with username {username}")
try:
result = client.login(
registry=registry_url, username=username, password=password
)
logger.info(
f"Logged in to {registry_url} with username {username} result: {result}"
)
except:
logger.error(
f"Failed to login to {registry_url} with username {username}\n{traceback.format_exc()}"
)
raise
def get_available_port() -> int:
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
s.bind(("127.0.0.1", 0))
port = s.getsockname()[1]
return port
@@ -0,0 +1,58 @@
from contextlib import asynccontextmanager
from fastapi import FastAPI, Request
import logging
from .configs import container_server_port
from .routers import api_server
from . import container_server_manager
# Configure logging
logging.basicConfig(
level=logging.INFO, format="%(asctime)s - %(name)s - %(levelname)s - %(message)s"
)
# Suppress httpx INFO logs
logging.getLogger("httpx").setLevel(logging.WARNING)
logger = logging.getLogger(__name__)
@asynccontextmanager
async def lifespan(app: FastAPI):
"""FastAPI lifespan context manager for startup and shutdown events"""
# Startup
await container_server_manager.wait_docker_ready()
await container_server_manager.load_mcp_server_image()
await container_server_manager.start_mcp_server_life_cycle_manager()
await container_server_manager.start_container_server_register_task()
try:
yield
finally:
# Shutdown
# await mcp_server_manager.clean_mcp_server_container()
pass
# FastAPI application setup with lifespan
app = FastAPI(
title="MCP Container Server",
description="MCP Container Server",
version="1.0.0",
lifespan=lifespan,
)
app.include_router(api_server.router, prefix="/api")
@app.get("/health")
async def health(request: Request):
return {"status": "success", "message": "Container server is healthy"}
if __name__ == "__main__":
import uvicorn
# Run the server
uvicorn.run(app, host="0.0.0.0", port=container_server_port)
@@ -0,0 +1,45 @@
import logging
from fastapi import APIRouter, Request
from .. import container_server_manager
logger = logging.getLogger(__name__)
router = APIRouter()
@router.post("/container/create")
async def create_container(request: Request, body: dict):
token = body.get("token")
logger.info(f"Create container: token={token}")
container_id, ip_addr, mcp_port, novnc_port = (
await container_server_manager.create_mcp_server_container()
)
logger.info(
f"Container created: container_id={container_id}, ip_addr={ip_addr}, mcp_port={mcp_port}, novnc_port={novnc_port}"
)
return {
"status": "success",
"message": f"MCP server created: {ip_addr}:{mcp_port}",
"data": {
"ip_addr": ip_addr,
"mcp_port": mcp_port,
"novnc_port": novnc_port,
"container_id": container_id,
},
}
@router.post("/container/shutdown")
async def shutdown_container(request: Request, body: dict):
token = body.get("token")
container_id = body.get("container_id")
logger.info(f"Shutdown container: token={token}, container_id={container_id}")
await container_server_manager.shutdown_mcp_server_container(container_id)
logger.info(f"Container shutdown: container_id={container_id}")
return {
"status": "success",
"message": f"MCP server shutdown: {container_id}",
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,50 @@
services:
mcp-gateway:
build:
context: mcp_gateway
dockerfile: Dockerfile_mcp_gateway
platform: linux/amd64
ports:
- 8000:8000
environment:
MCP_GATEWAY_TOKEN_SECRET: "${MCP_GATEWAY_TOKEN_SECRET:-123321}"
VNC_AUTH: false
MCP_GATEWAY_REDIS_URL: redis://redis-server:6379/0
restart: on-failure:3
depends_on:
- redis-server
deploy:
resources:
limits:
memory: 2G
container-server:
build:
context: .
dockerfile: container_server/Dockerfile_container_server
platform: linux/amd64
restart: on-failure:3
privileged: true
volumes:
- /var/run/docker.sock:/var/run/docker.sock
command: ["uv", "run", "--no-sync", "-m", "container_server.main"]
environment:
VIRTUALPC_MCP_SERVER_IMAGE_ID: gaia-mcp-server
GATEWAY_SERVER_ADDR: http://mcp-gateway:8000
deploy:
resources:
limits:
memory: 8G
redis-server:
image: redis:7.2-alpine
platform: linux/amd64
ports:
- 6379:6379
volumes:
- ./.data/redis-data:/data
restart: on-failure:3
deploy:
resources:
limits:
memory: 2G
@@ -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
@@ -0,0 +1,99 @@
FROM python:3.13-slim-bookworm as base
# 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 \
nodejs \
npm \
git \
xvfb \
x11vnc \
openbox \
supervisor \
novnc \
websockify \
procps \
xdg-utils \
python3-xdg \
x11-xserver-utils \
curl \
--no-install-recommends
# Install noVNC
RUN git clone --depth 1 --branch v1.6.0 https://github.com/novnc/noVNC.git /usr/local/novnc \
&& git clone --depth 1 --branch v0.13.0 https://github.com/novnc/websockify /usr/local/novnc/utils/websockify
# Set up working directory
WORKDIR /app/view_server
# Install Playwright and browsers with dependencies
RUN npm install playwright@1.54
RUN npx playwright install chromium --with-deps
# Set up supervisord configuration
COPY docker/resources/supervisord.conf /etc/supervisor/supervisord.conf
# Copy scripts
COPY docker/resources/start.sh start.sh
COPY docker/resources/playwright-server.js playwright-server.js
COPY docker/resources/x11-setup.sh x11-setup.sh
# Make scripts executable
RUN chmod +x start.sh x11-setup.sh
# Create a simple openbox configuration to only show the browser window
RUN mkdir -p /root/.config/openbox
COPY docker/resources/openbox-rc.xml /root/.config/openbox/rc.xml
ENV PLAYWRIGHT_WS_PATH="default"
ENV PLAYWRIGHT_PORT=37367
ENV NO_VNC_PORT=5901
# Set the display environment variable
ENV DISPLAY=:99
COPY docker/resources/entrypoint.sh /usr/local/bin/entrypoint.sh
RUN chmod +x /usr/local/bin/entrypoint.sh
ENTRYPOINT ["/usr/local/bin/entrypoint.sh"]
# MCP Server Container
FROM base
RUN pip install -U --break-system-packages uv
COPY mcp_server_proxy /app/mcp_server_proxy
RUN cd /app/mcp_server_proxy && uv sync --python-preference=only-system
RUN apt-get install -y --no-install-recommends \
wget \
unzip \
libterm-readline-perl-perl \
libmupdf-dev \
vim \
libmagic1
WORKDIR /root/workspace
EXPOSE 4242
HEALTHCHECK --interval=10s --timeout=10s --start-period=10s --retries=12 \
CMD curl -f http://localhost:4242/health || exit 1
CMD ["sh", "-c", "/app/view_server/start.sh && cd /app/mcp_server_proxy && uv run --no-sync -m mcp_server_proxy.main"]
@@ -0,0 +1,7 @@
#!/bin/sh
cd "$(dirname "$0")"
docker build -t mcp-server-base -f Dockerfile_mcp_server . && \
echo "✅ Build image success: mcp-server-base"
@@ -0,0 +1,6 @@
#!/usr/bin/env bash
set -e
umask 000
exec "$@"
@@ -0,0 +1,30 @@
<?xml version="1.0" encoding="UTF-8"?>
<openbox_config xmlns="http://openbox.org/3.4/rc" xmlns:xi="http://www.w3.org/2001/XInclude">
<desktops>
<number>1</number>
</desktops>
<margins>
<top>0</top>
<bottom>0</bottom>
<left>0</left>
<right>0</right>
</margins>
<applications>
<application class="*">
<decor>no</decor>
<maximized>yes</maximized>
<fullscreen>yes</fullscreen>
<position>
<x>0</x>
<y>0</y>
</position>
<size>
<width>100%</width>
<height>100%</height>
</size>
<focus>yes</focus>
<desktop>1</desktop>
<layer>normal</layer>
</application>
</applications>
</openbox_config>
@@ -0,0 +1,7 @@
{
"name": "playwright-remote",
"version": "1.0.0",
"dependencies": {
"playwright": "1.51.1"
}
}
@@ -0,0 +1,39 @@
const { chromium } = require("playwright");
// Read ws path from environment variable
const wsPath = process.env.WS_PATH || "default";
const port = process.env.PLAYWRIGHT_PORT || 37367;
(async () => {
console.log("Starting Playwright server...");
// Start the remote debugging server
const browserServer = await chromium.launchServer({
headless: false,
port: port,
wsPath: wsPath,
args: [
"--start-fullscreen",
"--start-maximized",
"--window-size=1280,1280",
"--window-position=0,0",
"--disable-infobars",
"--no-default-browser-check",
"--kiosk",
"--disable-session-crashed-bubble",
"--noerrdialogs",
"--force-device-scale-factor=1.0",
"--disable-features=DefaultViewportMetaTag",
"--force-device-width=1280",
],
});
console.log(`Playwright server running: ${browserServer.wsEndpoint()}`);
// Keep the process running
process.on("SIGINT", async () => {
console.log("Shutting down Playwright server...");
// await browser.close();
await browserServer.close();
process.exit(0);
});
})();
@@ -0,0 +1,8 @@
#!/bin/bash
set -e
echo "Starting services..."
echo "DISPLAY=$DISPLAY"
# Start supervisord to manage all processes
exec supervisord -c /etc/supervisor/supervisord.conf
@@ -0,0 +1,64 @@
[supervisord]
logfile=/var/log/supervisord.log
logfile_maxbytes=50MB
loglevel=info
[include]
files = /etc/supervisor/conf.d/*.conf
[program:xvfb]
command=Xvfb :99 -screen 0 2560x2560x24 -dpi 192 -ac -nolisten tcp
autorestart=true
stdout_logfile=/dev/stdout
stdout_logfile_maxbytes=0
stderr_logfile=/dev/stderr
stderr_logfile_maxbytes=0
[program:openbox]
command=openbox-session
environment=DISPLAY=:99
autorestart=true
stdout_logfile=/dev/stdout
stdout_logfile_maxbytes=0
stderr_logfile=/dev/stderr
stderr_logfile_maxbytes=0
[program:x11setup]
command=/app/view_server/x11-setup.sh
environment=DISPLAY=:99
autorestart=false
startsecs=0
startretries=0
priority=10
stdout_logfile=/dev/stdout
stdout_logfile_maxbytes=0
stderr_logfile=/dev/stderr
stderr_logfile_maxbytes=0
[program:x11vnc]
command=x11vnc -display :99 -forever -shared -nopw -geometry 1280x1280 -scale 1:1 -nomodtweak -noxdamage
autorestart=true
priority=20
stdout_logfile=/dev/stdout
stdout_logfile_maxbytes=0
stderr_logfile=/dev/stderr
stderr_logfile_maxbytes=0
[program:novnc]
command=/usr/local/novnc/utils/novnc_proxy --vnc localhost:5900 --listen %(ENV_NO_VNC_PORT)s
autorestart=true
stdout_logfile=/dev/stdout
stdout_logfile_maxbytes=0
stderr_logfile=/dev/stderr
stderr_logfile_maxbytes=0
[program:playwright-server]
command=node /app/view_server/playwright-server.js
environment=DISPLAY=:99,WS_PATH=%(ENV_PLAYWRIGHT_WS_PATH)s,PLAYWRIGHT_PORT=%(ENV_PLAYWRIGHT_PORT)s
autorestart=true
priority=30
startsecs=1
stdout_logfile=/dev/stdout
stdout_logfile_maxbytes=0
stderr_logfile=/dev/stderr
stderr_logfile_maxbytes=0
@@ -0,0 +1,29 @@
#!/bin/bash
# Make sure DISPLAY is set
if [ -z "$DISPLAY" ]; then
export DISPLAY=:99
fi
# Set background to black to make black bars less obvious
xsetroot -solid "#000000"
# Force X11 to use the exact screen dimensions without any offsets
xrandr --output default --mode 1280x1280 --pos 0x0
# Set proper DPI settings for the display
echo "Xft.dpi: 96" | xrdb -merge
echo "Xft.antialias: 1" | xrdb -merge
echo "Xft.hinting: 1" | xrdb -merge
echo "Xft.hintstyle: hintfull" | xrdb -merge
echo "Xft.rgba: rgb" | xrdb -merge
# Disable any screen savers or power management
xset s off
xset -dpms
xset s noblank
# Ensure consistent scaling
xrandr --dpi 96
echo "X11 environment configured for optimal display"
@@ -0,0 +1,25 @@
[project]
name = "mcp-proxy"
version = "0.1.0"
description = "MCP Proxy"
requires-python = ">=3.12"
dependencies = [
"docker",
"aiohttp",
"playwright==1.52",
"pydantic",
"pydantic-settings",
"fastapi[standard]",
"typer",
"aiofiles",
"python-dotenv",
"websockets",
"pyyaml",
"mcp==1.12.4",
"httpx[http2]",
"requests>=2.32.5",
]
[build-system]
requires = ["setuptools"]
build-backend = "setuptools.build_meta"
@@ -0,0 +1,16 @@
import os
from pathlib import Path
mcp_servers_path = os.getenv(
"MCP_SERVERS_PATH",
str((Path(__file__).parent.parent.parent.parent / "mcp_servers").resolve()),
)
mcp_servers_config_path = os.getenv(
"MCP_SERVERS_CONFIG_PATH", str((Path(mcp_servers_path) / "mcp_config.py").resolve())
)
mcp_tool_schema_path = os.getenv(
"MCP_TOOL_SCHEMA_PATH",
str((Path(mcp_servers_path) / "mcp_tool_schema.json").resolve()),
)
@@ -0,0 +1,44 @@
import asyncio
import datetime
import logging
from fastapi import Request, Response
from fastapi.responses import JSONResponse
from .mcp_server_proxy import MCPServerProxy
# Configure logging
logging.basicConfig(
level=logging.INFO, format="%(asctime)s - %(name)s - %(levelname)s - %(message)s"
)
logger = logging.getLogger(__name__)
mcp = MCPServerProxy(
name="MCP Server",
stateless_http=False,
host="0.0.0.0",
port=4242,
log_level="DEBUG",
)
@mcp.custom_route("/health", methods=["GET"])
async def health(request: Request) -> Response:
return JSONResponse(
{
"status": "success",
"message": "MCP Server is healthy",
"last_active": datetime.datetime.now().strftime("%Y/%m/%d %H:%M:%S.%f"),
}
)
async def main():
logger.info("Starting MCP Server Proxy...")
await mcp.initialize()
await mcp.run_streamable_http_async()
if __name__ == "__main__":
asyncio.run(main())
@@ -0,0 +1,192 @@
from contextlib import AsyncExitStack
from datetime import timedelta
import json
import os
from pathlib import Path
import traceback
from typing import Any
import asyncio
from mcp import ClientSession, StdioServerParameters
from mcp.server.fastmcp import Context
from mcp.server.session import ServerSessionT
from mcp.shared.context import LifespanContextT, RequestT
import logging
from mcp.client.stdio import stdio_client
from mcp.client.sse import sse_client
from mcp.client.streamable_http import streamablehttp_client
from mcp.types import LoggingMessageNotificationParams
from .configs import mcp_servers_path
logger = logging.getLogger(__name__)
class MCPServerExecutor:
def __init__(self, name: str, config: dict):
self._name = name
self._config = config
self._session = None
self._exit_stack = None
self._lock = asyncio.Lock()
self._init_event = asyncio.Event()
self._terminate_event = asyncio.Event()
async def call_tool(
self,
name: str,
arguments: dict[str, Any],
context: Context[ServerSessionT, LifespanContextT, RequestT] | None = None,
convert_result: bool = False,
) -> Any:
"""Call a tool by name with arguments."""
await self._ensure_server_ready()
async def progress_callback_adapter(
progress: float, total: float | None, message: str | None
):
logger.info(
f"progress_callback: tool={name}, {progress}, {total}, {message}"
)
await self.progress_callback(
progress=progress,
total=total,
message=message,
context=context,
)
result = await self._session.call_tool(
name, arguments, progress_callback=progress_callback_adapter
)
return result.content, result.structuredContent
async def _ensure_server_ready(self):
if not self._session:
asyncio.create_task(self._start_tool_server())
await self._init_event.wait()
async def _start_tool_server(
self,
):
if self._session is not None:
return
async with self._lock:
if self._session is not None:
return
name: str = self._name
config: dict = self._config
try:
logger.info(f"Starting tool server {name} with config {config}")
exit_stack = AsyncExitStack()
await exit_stack.__aenter__()
# Create client context and enter it
if config.get("type") == "sse":
read_stream, write_stream = await exit_stack.enter_async_context(
sse_client(
url=config.get("url", ""),
headers=config.get("headers", {}),
timeout=config.get("timeout", 60 * 10),
sse_read_timeout=config.get("sse_read_timeout", 60 * 10),
auth=config.get("auth", None),
)
)
elif config.get("type") == "streamable_http":
read_stream, write_stream, _ = await exit_stack.enter_async_context(
streamablehttp_client(
url=config.get("url", ""),
headers=config.get("headers", {}),
timeout=config.get("timeout", 60 * 10),
sse_read_timeout=config.get("sse_read_timeout", 60 * 10),
auth=config.get("auth", None),
)
)
else: # stdio
env = config.get("env", {})
env.update(os.environ)
server_params = StdioServerParameters(
command=config.get("command", ""),
args=config.get("args", []),
env=env,
cwd=str(Path(mcp_servers_path) / config.get("cwd", "")),
)
read_stream, write_stream = await exit_stack.enter_async_context(
stdio_client(server=server_params)
)
async def log_callback(params: LoggingMessageNotificationParams):
logger.info(f"MCP Server {name} >>> {params}")
# Create session and tool manager
session = await exit_stack.enter_async_context(
ClientSession(
read_stream,
write_stream,
logging_callback=log_callback,
read_timeout_seconds=timedelta(
seconds=config.get("read_timeout", 60 * 10)
),
)
)
await session.initialize()
self._session = session
self._exit_stack = exit_stack
logger.info(f"Starting tool server success! {name}: {config}")
self._init_event.set()
await self._terminate_event.wait()
except Exception as e:
logger.error(
f"Error starting tool server {name}: {config}\n{traceback.format_exc()}"
)
self._init_event.set()
try:
await self._exit_stack.aclose()
except Exception:
pass
raise e
async def cleanup(self):
self._terminate_event.set()
self._session = None
self._exit_stack = None
async def progress_callback(
self,
progress: float,
total: float | None,
message: str | None,
context: Context[ServerSessionT, LifespanContextT, RequestT] | None = None,
):
def get_session_id() -> str:
try:
return context.request_context.request.headers.get("Mcp-Session-Id")
except Exception as e:
logger.error(f"Error getting session id: {e}")
return None
if "tool_call_card_novnc_window" == message:
"""Show the VNC window"""
vnc_tool_card = {
"type": "tool_call_card_novnc_window",
"card_data": {
"title": "VNC Window",
"url": f"/novnc/{get_session_id()}/vnc.html?autoconnect=true&reconnect=true&quality=9&compression=9&show_dot=0&resize=local",
"token": get_session_id(),
},
}
message = f"""\
\n\n
```tool_card
{json.dumps(vnc_tool_card, indent=2, ensure_ascii=False)}
```
\n\n
"""
if context:
await context.report_progress(
progress=progress, total=total, message=message
)
@@ -0,0 +1,28 @@
import json
from importlib.util import spec_from_file_location, module_from_spec
from pathlib import Path
from .configs import mcp_servers_config_path, mcp_tool_schema_path
class MCPServerLoader:
def __init__(self):
pass
def load_mcp_servers_config(self):
mcp_config = self._load_mcp_config()
return mcp_config.get("mcpServers", {})
def load_mcp_tool_schema(self):
with open(mcp_tool_schema_path, "r") as f:
return json.load(f)
def _load_mcp_config(self):
path = Path(mcp_servers_config_path).resolve()
assert path.exists(), f"MCP servers config file not found: {path}"
spec = spec_from_file_location("mcp_servers_config", path)
module = module_from_spec(spec)
spec.loader.exec_module(module)
mcp_config = getattr(module, "mcp_config")
return mcp_config
@@ -0,0 +1,117 @@
import traceback
from typing import Any, Callable, Sequence
from mcp.server.fastmcp import Context, FastMCP
from mcp.server.fastmcp.exceptions import ToolError
from mcp.server.fastmcp.resources import Resource
from mcp.server.fastmcp.tools import Tool, ToolManager
from mcp.server.session import ServerSessionT
from mcp.shared.context import LifespanContextT, RequestT
from mcp.types import ContentBlock, ToolAnnotations
from mcp.server.fastmcp.utilities.func_metadata import ArgModelBase, FuncMetadata
import logging
from mcp.server.fastmcp.tools.base import Tool as ServerTool
from mcp.types import Tool as ClientTool
from .mcp_server_executor import MCPServerExecutor
from .mcp_server_loader import MCPServerLoader
logger = logging.getLogger(__name__)
class MCPServerProxy(FastMCP):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self._mcp_tool_schema: dict[str, list[dict[str, Any]]] = {}
self._mcp_server_executors: dict[str, MCPServerExecutor] = {}
self._mcp_server_loader = MCPServerLoader()
async def initialize(self):
self._load_tool_schema()
self._load_mcp_servers()
async def call_tool(
self, name: str, arguments: dict[str, Any]
) -> Sequence[ContentBlock] | dict[str, Any]:
"""Call a tool by name with arguments."""
try:
context = self.get_context()
request_mcp_server_executor = self._get_request_mcp_server_executor(
tool_name=name
)
return await request_mcp_server_executor.call_tool(
name, arguments, context=context, convert_result=True
)
except:
logger.error(f"Error calling tool {name}: {traceback.format_exc()}")
raise
async def list_tools(self) -> list[ClientTool]:
"""List all available tools."""
try:
request_mcp_servers = self._get_request_mcp_servers()
request_tools = [
tool
for server_name, server_tools in self._mcp_tool_schema.items()
if server_name in request_mcp_servers
for tool in server_tools
]
return [
ClientTool(
name=tool.get("name", ""),
title=tool.get("title", ""),
description=tool.get("description", ""),
inputSchema=tool.get("inputSchema", {}),
outputSchema=tool.get("outputSchema", {}),
annotations=tool.get("annotations", {}),
_meta=tool.get("_meta", {}),
)
for tool in request_tools
]
except:
logger.error(f"Error listing tools: {traceback.format_exc()}")
raise
def _load_tool_schema(self):
self._mcp_tool_schema = self._mcp_server_loader.load_mcp_tool_schema()
mcp_tool_schema = ""
for server_name, server_tools in self._mcp_tool_schema.items():
mcp_tool_schema += f" {server_name}:\n"
for tool in server_tools:
mcp_tool_schema += f" - {tool.get('name', '')}\n"
logger.info(f"Loaded MCP tool schema: mcp_tool_schema={mcp_tool_schema}")
def _load_mcp_servers(self):
for name, config in self._mcp_server_loader.load_mcp_servers_config().items():
self._mcp_server_executors[name] = MCPServerExecutor(name, config)
logger.info(f"Added MCP server executor: {name}")
def _get_request_mcp_servers(self) -> list[str]:
context = self.get_context()
request_servers = context.request_context.request.headers.get("MCP_SERVERS")
if request_servers:
return [server.strip() for server in request_servers.split(",")]
return []
def _get_request_mcp_server_executor(self, tool_name: str) -> MCPServerExecutor:
request_mcp_servers = self._get_request_mcp_servers()
request_tools = {
server_name: tool
for server_name, server_tools in self._mcp_tool_schema.items()
if server_name in request_mcp_servers
for tool in server_tools
if tool.get("name", "") == tool_name
}
if not request_tools:
raise ToolError(f"Tool {tool_name} not found")
else:
if len(request_tools) > 1:
logger.warning(
f"Tool {tool_name} found in multiple MCP servers: {request_tools}"
)
server_name = list(request_tools.keys())[0]
return self._mcp_server_executors[server_name]
@@ -0,0 +1,102 @@
import asyncio
import json
import subprocess
import logging
from pathlib import Path
from typing import Any, AsyncGenerator, List
from mcp import ClientSession
from mcp.types import (
LoggingMessageNotificationParams,
ElicitResult,
ElicitRequestParams,
)
from mcp.client.streamable_http import streamablehttp_client
from mcp.shared.context import RequestContext
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
)
logger = logging.getLogger(__name__)
async def mcp_client(
url: str,
token: str,
session_id: str = None,
mcp_servers: List[str] = [
"readweb-server",
"browser-server",
"browseruse-server",
"documents-csv-server",
"documents-docx-server",
"documents-pptx-server",
"documents-pdf-server",
"documents-txt-server",
"download-server",
"intelligence-code-server",
"intelligence-think-server",
"intelligence-guard-server",
"media-audio-server",
"media-image-server",
"media-video-server",
"parxiv-server",
"terminal-server",
"wayback-server",
"wiki-server",
"googlesearch-server",
],
) -> AsyncGenerator[ClientSession, None]:
headers = {
"Authorization": f"Bearer {token}",
"MCP_SERVERS": ",".join(mcp_servers),
}
if session_id:
headers["SESSION_ID"] = session_id
async with streamablehttp_client(
url=url,
headers=headers,
) as (
read_stream,
write_stream,
get_session_id,
):
async def logging_callback(params: LoggingMessageNotificationParams):
logger.info(f"Receive logging callback: {params}")
async def elicitation_callback(
context: RequestContext["ClientSession", Any],
params: ElicitRequestParams,
) -> ElicitResult:
logger.info(f"Receive elicitation callback: {params}")
return ElicitResult(action="accept", content={"user_name": "John"})
async with ClientSession(
read_stream=read_stream,
write_stream=write_stream,
logging_callback=logging_callback,
elicitation_callback=elicitation_callback,
) as session:
logger.info(f"MCP client connected: url={url}")
await session.initialize()
logger.info(
f"MCP client session initialized: url={url}, session_id={get_session_id()}"
)
yield session
async def progress_callback(progress: float, total: float | None, message: str | None):
logger.info(
f"Receive progress callback: progress={progress}, total={total}, message={message}"
)
if "```tool_card" in message:
data = json.loads(message.split("```tool_card")[1].split("```")[0])
vnc_url = f"{base_url}{data.get('card_data').get('url')}"
logger.info(f"VNC URL: {vnc_url}")
subprocess.run(["open", vnc_url])
@@ -0,0 +1,78 @@
import asyncio
from contextlib import AsyncExitStack
from datetime import timedelta
import json
import logging
from pathlib import Path
from mcp import ClientSession, StdioServerParameters, stdio_client
from mcp.client.sse import sse_client
from mcp.client.streamable_http import streamablehttp_client
logger = logging.getLogger(__name__)
config = {
"type": "stdio",
"command": "npx",
"args": ["-y", "@modelcontextprotocol/server-filesystem", "~/workspace"],
}
async def server_session():
async with AsyncExitStack() as exit_stack:
# Create client context and enter it
if config.get("type") == "sse":
read_stream, write_stream = await exit_stack.enter_async_context(
sse_client(
url=config.get("url", ""),
headers=config.get("headers", {}),
timeout=config.get("timeout", 60),
sse_read_timeout=config.get("sse_read_timeout", 60 * 5),
auth=config.get("auth", None),
)
)
elif config.get("type") == "streamable_http":
read_stream, write_stream, _ = await exit_stack.enter_async_context(
streamablehttp_client(
url=config.get("url", ""),
headers=config.get("headers", {}),
timeout=config.get("timeout", 120),
sse_read_timeout=config.get("sse_read_timeout", 60 * 5),
auth=config.get("auth", None),
)
)
else: # stdio
base_folder = Path(__file__).parent
server_params = StdioServerParameters(
command=config.get("command", ""),
args=config.get("args", []),
env=config.get("env", {}),
cwd=str(base_folder / config.get("cwd", "")),
)
read_stream, write_stream = await exit_stack.enter_async_context(
stdio_client(server=server_params)
)
# Create session and tool manager
session = await exit_stack.enter_async_context(
ClientSession(
read_stream,
write_stream,
read_timeout_seconds=timedelta(seconds=config.get("read_timeout", 120)),
)
)
await session.initialize()
yield session
async def test():
async for session in server_session():
ls = await session.list_tools()
assert ls and ls.tools, "list_tools return null"
tools = ls.tools
logger.info(f"list_tools return:\n - {'\n - '.join([t.name for t in tools])}")
print(tools[0])
if __name__ == "__main__":
asyncio.run(test())
@@ -0,0 +1 @@
tool_test_cases = [{"tool_name": "read_url", "args": {"url": "https://www.baidu.com"}}]
@@ -0,0 +1,112 @@
import asyncio
import base64
import hashlib
import hmac
import json
import subprocess
import logging
import os
import time
from pathlib import Path
from typing import Any, AsyncGenerator
from mcp import ClientSession
from mcp.types import (
LoggingMessageNotificationParams,
ElicitResult,
ElicitRequestParams,
)
from mcp.client.streamable_http import streamablehttp_client
from mcp.shared.context import RequestContext
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
)
logger = logging.getLogger(__name__)
LOCAL_MCP_TOKEN_SECRET = "123321"
def _jwt_part(value: dict) -> str:
raw = json.dumps(value, separators=(",", ":")).encode()
return base64.urlsafe_b64encode(raw).rstrip(b"=").decode()
def gen_local_mcp_token(app: str = "mcp-gateway-debug") -> str:
secret = os.getenv("MCP_GATEWAY_TOKEN_SECRET", LOCAL_MCP_TOKEN_SECRET)
header = {"alg": "HS256", "typ": "JWT"}
payload = {"app": app, "version": 1, "time": time.time()}
signing_input = f"{_jwt_part(header)}.{_jwt_part(payload)}"
signature = hmac.new(
secret.encode(),
signing_input.encode(),
hashlib.sha256,
).digest()
encoded_signature = base64.urlsafe_b64encode(signature).rstrip(b"=").decode()
return f"{signing_input}.{encoded_signature}"
if __name__ == "__main__":
base_url, token = (
"http://localhost:8000",
gen_local_mcp_token(),
)
asyncio.run(McpClient.mcp_test_client(base_url, token))
class McpClient:
async def mcp_test_client(
base_url: str, token: str
) -> AsyncGenerator[ClientSession, None]:
url = f"{base_url}/mcp"
async with streamablehttp_client(
url=url,
headers={
"Authorization": f"Bearer {token}",
"MCP_SERVERS": "readweb-server,browser-server,browseruse-server,documents-csv-server,documents-docx-server,documents-pptx-server,documents-pdf-server,documents-txt-server,download-server,intelligence-code-server,intelligence-think-server,intelligence-guard-server,media-audio-server,media-image-server,media-video-server,parxiv-server,terminal-server,wayback-server,wiki-server,googlesearch-server",
# "SESSION_ID": "CHAT_WLDEV",
},
) as (
read_stream,
write_stream,
get_session_id,
):
async def logging_callback(params: LoggingMessageNotificationParams):
logger.info(f"Receive logging callback: {params}")
async def elicitation_callback(
context: RequestContext["ClientSession", Any],
params: ElicitRequestParams,
) -> ElicitResult:
logger.info(f"Receive elicitation callback: {params}")
return ElicitResult(action="accept", content={"user_name": "John"})
async with ClientSession(
read_stream=read_stream,
write_stream=write_stream,
logging_callback=logging_callback,
elicitation_callback=elicitation_callback,
) as session:
logger.info(f"MCP client connected: url={url}")
await session.initialize()
logger.info(
f"MCP client session initialized: url={url}, session_id={get_session_id()}"
)
yield session
async def progress_callback(progress: float, total: float | None, message: str | None):
logger.info(
f"Receive progress callback: progress={progress}, total={total}, message={message}"
)
if "```tool_card" in message:
data = json.loads(message.split("```tool_card")[1].split("```")[0])
vnc_url = f"{base_url}{data.get('card_data').get('url')}"
logger.info(f"VNC URL: {vnc_url}")
subprocess.run(["open", vnc_url])
@@ -0,0 +1,51 @@
import asyncio
import json
import logging
from pathlib import Path
import sys
import os
from dotenv import load_dotenv
# Add the project root to Python path
sys.path.insert(0, str(Path(__file__).parent.parent))
from core.mcp_client import mcp_client, progress_callback
from core.test_data import tool_test_cases
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
)
logger = logging.getLogger(__name__)
load_dotenv()
read_arg = lambda e: (os.getenv(f"URL_{e}"), os.getenv(f"TOKEN_{e}"))
url, token = read_arg("REMOTE")
# url, token = read_arg("GW_DEBUG")
# url, token = read_arg("MCP_DEBUG")
async def main():
async for session in mcp_client(url, token):
ls = await session.list_tools()
assert ls and ls.tools, "list_tools return null"
tools = ls.tools
logger.info(f"list_tools return:\n - {'\n - '.join([t.name for t in tools])}")
t = tools[0]
assert t.name, "tool.name is null"
assert t.inputSchema, "tool.inputSchema is null"
assert t.outputSchema, "tool.outputSchema is null"
for t in tool_test_cases:
tool_name = t["tool_name"]
args = t["args"]
logger.info(f"call tool: {tool_name}")
result = await session.call_tool(tool_name, args, progress_callback=progress_callback)
logger.info(f"tool result: {result.content[0].text[:300]}")
input("Press Enter to continue...")
if __name__ == "__main__":
asyncio.run(main())
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,4 @@
#!/bin/sh
cd "$(dirname "$0")"
docker compose up --build --force-recreate -d