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
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:
Vendored
+29
@@ -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"]
|
||||
|
||||
|
||||
+20
@@ -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"
|
||||
Vendored
Vendored
+17
@@ -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")
|
||||
+231
@@ -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")
|
||||
+229
@@ -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
|
||||
Vendored
+58
@@ -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)
|
||||
+45
@@ -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}",
|
||||
}
|
||||
+1320
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
|
||||
+39
@@ -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"
|
||||
Vendored
Vendored
+43
@@ -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
|
||||
+12
@@ -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"
|
||||
chapter9/gaia-experience/AWorld/env/virtualpc-mcp/mcp_gateway/src/mcp_gateway/containers/__init__.py
Vendored
+1
@@ -0,0 +1 @@
|
||||
from .container_server_manager import ContainerServer, ContainerServerManager, container_server_manager_builder
|
||||
+300
@@ -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
|
||||
+96
@@ -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)
|
||||
Vendored
+88
@@ -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'}"
|
||||
)
|
||||
Vendored
Vendored
+44
@@ -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,
|
||||
)
|
||||
chapter9/gaia-experience/AWorld/env/virtualpc-mcp/mcp_gateway/src/mcp_gateway/routers/gateway_api.py
Vendored
+25
@@ -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"}
|
||||
+105
@@ -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}")
|
||||
Vendored
+5
@@ -0,0 +1,5 @@
|
||||
from .session_connection import SessionId, VpcSession
|
||||
from .session_connection_manager import (
|
||||
SessionConnectionManager,
|
||||
session_connection_manager_builder,
|
||||
)
|
||||
+105
@@ -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
|
||||
+295
@@ -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
|
||||
Vendored
Vendored
+69
@@ -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
|
||||
Vendored
+68
@@ -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,
|
||||
)
|
||||
+18
@@ -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()
|
||||
+13
@@ -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)
|
||||
+1393
File diff suppressed because it is too large
Load Diff
+99
@@ -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"
|
||||
Vendored
+6
@@ -0,0 +1,6 @@
|
||||
#!/usr/bin/env bash
|
||||
set -e
|
||||
|
||||
umask 000
|
||||
|
||||
exec "$@"
|
||||
Vendored
+30
@@ -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>
|
||||
+7
@@ -0,0 +1,7 @@
|
||||
{
|
||||
"name": "playwright-remote",
|
||||
"version": "1.0.0",
|
||||
"dependencies": {
|
||||
"playwright": "1.51.1"
|
||||
}
|
||||
}
|
||||
Vendored
+39
@@ -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);
|
||||
});
|
||||
})();
|
||||
+8
@@ -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
|
||||
Vendored
+64
@@ -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
|
||||
+29
@@ -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"
|
||||
Vendored
+25
@@ -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"
|
||||
+16
@@ -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()),
|
||||
)
|
||||
+44
@@ -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())
|
||||
+192
@@ -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
|
||||
)
|
||||
+28
@@ -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
|
||||
+117
@@ -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]
|
||||
+102
@@ -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])
|
||||
+78
@@ -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())
|
||||
+1
@@ -0,0 +1 @@
|
||||
tool_test_cases = [{"tool_name": "read_url", "args": {"url": "https://www.baidu.com"}}]
|
||||
+112
@@ -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])
|
||||
+51
@@ -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())
|
||||
Generated
Vendored
+1439
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
|
||||
Reference in New Issue
Block a user