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:
+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
Reference in New Issue
Block a user