Files
mcp-linkgen/server.py
T
2026-09-04 18:30:15 +01:00

365 lines
13 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/usr/bin/env python3
"""
mcp-linkgen — Lightweight MCP server for generating ephemeral download links.
Exposes a single MCP tool (`generate_download_link`) that mints a single-use,
time-limited HTTPS download URL for any file inside the configured workspace
directory. An HTTP route (`/download`) serves the file once and invalidates
the token.
Designed to run standalone via `python server.py` or inside a minimal Docker
container.
"""
import os
import sys
import json
import logging
import secrets
import time
from collections import defaultdict
from typing import Dict
from fastmcp import FastMCP
from fastmcp.server.middleware import Middleware, MiddlewareContext
from starlette.middleware import Middleware as StarletteMiddleware
from starlette.middleware.cors import CORSMiddleware
from starlette.responses import FileResponse, JSONResponse
# ---------------------------------------------------------------------------
# Optional .env loading (convenience for native/local development)
# ---------------------------------------------------------------------------
try:
from dotenv import load_dotenv
load_dotenv()
except ImportError:
pass
# ---------------------------------------------------------------------------
# Logging — structured JSON to stdout
# ---------------------------------------------------------------------------
class JSONLogFormatter(logging.Formatter):
"""Emit one JSON object per log line for machine-parseable output."""
def format(self, record):
log_record = {
"timestamp": self.formatTime(record, self.datefmt),
"level": record.levelname,
"logger": record.name,
"message": record.getMessage(),
"file": f"{record.filename}:{record.lineno}",
}
if record.exc_info:
log_record["exception"] = self.formatException(record.exc_info)
return json.dumps(log_record)
_json_handler = logging.StreamHandler(sys.stdout)
_json_handler.setFormatter(JSONLogFormatter())
logging.basicConfig(
level=logging.INFO,
handlers=[_json_handler],
force=True,
)
logger = logging.getLogger("mcp-linkgen")
# ---------------------------------------------------------------------------
# Configuration (all values overridable via environment variables)
# ---------------------------------------------------------------------------
HOST = os.environ.get("HOST", "0.0.0.0")
PORT = int(os.environ.get("PORT", "8000"))
WORKSPACE_ROOT = os.path.abspath(os.environ.get("WORKSPACE_ROOT", "/workspace"))
PUBLIC_BASE_URL = os.environ.get("PUBLIC_BASE_URL", "http://localhost:8000")
DOWNLOAD_DEFAULT_TTL = int(os.environ.get("DOWNLOAD_DEFAULT_TTL", "10")) # minutes
DOWNLOAD_MIN_TTL = int(os.environ.get("DOWNLOAD_MIN_TTL", "1")) # minutes
DOWNLOAD_MAX_TTL = int(os.environ.get("DOWNLOAD_MAX_TTL", "60")) # minutes
# CORS origins — comma-separated list, or "*" for all.
_cors_raw = os.environ.get("CORS_ORIGINS", "*").strip()
CORS_ORIGINS: list[str] = (
["*"] if _cors_raw == "*" else [o.strip() for o in _cors_raw.split(",") if o.strip()]
)
# Rate limiting (per-IP sliding window; 0 = disabled).
MCP_RATE_LIMIT = int(os.environ.get("MCP_RATE_LIMIT", "30")) # req/min per IP
DOWNLOAD_RATE_LIMIT = int(os.environ.get("DOWNLOAD_RATE_LIMIT", "10")) # req/min per IP
RATE_LIMIT_WINDOW = int(os.environ.get("RATE_LIMIT_WINDOW", "60")) # sliding window (s)
MCP_ENDPOINT = os.environ.get("MCP_ENDPOINT", "/mcp") # MCP HTTP path (FastMCP default)
logger.info("Workspace root: %s", WORKSPACE_ROOT)
logger.info("Public base URL: %s", PUBLIC_BASE_URL)
logger.info("CORS origins: %s", CORS_ORIGINS)
logger.info("Rate limits — MCP: %d/min, download: %d/min, window: %ds",
MCP_RATE_LIMIT, DOWNLOAD_RATE_LIMIT, RATE_LIMIT_WINDOW)
# ---------------------------------------------------------------------------
# FastMCP application
# ---------------------------------------------------------------------------
mcp = FastMCP("mcp-linkgen")
# In-memory token store — tokens never touch disk.
TOKEN_STORE: Dict[str, dict] = {}
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _validate_download_path(raw_path: str) -> str:
"""Resolve *raw_path* and ensure it lives inside ``WORKSPACE_ROOT``.
Returns the canonical absolute path on success.
Raises
------
PermissionError
If the resolved path escapes the workspace root.
FileNotFoundError
If the file does not exist.
IsADirectoryError
If the target is a directory (archives must be used for directories).
"""
canonical_path = os.path.realpath(raw_path)
if not canonical_path.startswith(WORKSPACE_ROOT + "/"):
raise PermissionError("Access denied: path escapes WORKSPACE_ROOT")
if not os.path.exists(canonical_path):
raise FileNotFoundError(f"File not found: {raw_path}")
if os.path.isdir(canonical_path):
raise IsADirectoryError(
"Target path is a directory. Please compress it to an archive first."
)
return canonical_path
def _purge_expired_tokens():
"""Remove all tokens whose TTL has elapsed."""
now = time.time()
expired = [k for k, v in TOKEN_STORE.items() if now > v["expires_at"]]
for k in expired:
TOKEN_STORE.pop(k, None)
# ---------------------------------------------------------------------------
# HTTP routes
# ---------------------------------------------------------------------------
@mcp.custom_route("/health", methods=["GET"])
async def health_check(request):
"""Liveness probe — returns 200 when the server is running."""
return JSONResponse({"status": "healthy", "service": "mcp-linkgen"})
@mcp.custom_route("/download", methods=["GET"])
async def download_file(request):
"""Serve a file identified by a single-use token, then invalidate it."""
_purge_expired_tokens()
token = request.query_params.get("t")
if not token:
return JSONResponse({"detail": "Missing token parameter."}, status_code=400)
record = TOKEN_STORE.get(token)
if not record:
return JSONResponse(
{"detail": "Invalid or previously used download link."}, status_code=403
)
# Immediately consume the token (single-use).
del TOKEN_STORE[token]
if time.time() > record["expires_at"]:
return JSONResponse({"detail": "Download link has expired."}, status_code=410)
file_path = record["file_path"]
if not os.path.exists(file_path):
return JSONResponse(
{"detail": "File no longer exists on host."}, status_code=404
)
return FileResponse(
path=file_path,
filename=os.path.basename(file_path),
media_type="application/octet-stream",
)
# ---------------------------------------------------------------------------
# Middleware — rate limiting, structured logging, CORS
# ---------------------------------------------------------------------------
class RateLimitMiddleware:
"""Per-IP sliding-window rate limiter for selected routes.
Wraps a Starlette application. Routes matching ``MCP_ENDPOINT`` (POST)
or ``/download`` (GET) are subject to their respective limits. All other
requests pass through unconditionally.
"""
def __init__(self, app, mcp_limit: int, download_limit: int, window: int, mcp_path: str):
self.app = app
self.mcp_limit = mcp_limit
self.download_limit = download_limit
self.window = window
self.mcp_path = mcp_path
self._hits: Dict[str, list[float]] = defaultdict(list)
self._last_cleanup = time.time()
# -- internal helpers ---------------------------------------------------
def _client_ip(self, request) -> str:
forwarded = request.headers.get("x-forwarded-for")
if forwarded:
return forwarded.split(",")[0].strip()
return request.client.host if request.client else "unknown"
def _prune(self, timestamps: list[float], now: float) -> list[float]:
cutoff = now - self.window
return [t for t in timestamps if t > cutoff]
def _cleanup_if_stale(self, now: float):
"""Periodically purge stale entries to bound memory usage."""
if now - self._last_cleanup > self.window:
stale_keys = [
ip for ip, ts in self._hits.items()
if not ts or ts[-1] <= now - self.window
]
for ip in stale_keys:
del self._hits[ip]
self._last_cleanup = now
# -- ASGI interface -----------------------------------------------------
async def __call__(self, scope, receive, send):
if scope["type"] != "http":
return await self.app(scope, receive, send)
now = time.time()
self._cleanup_if_stale(now)
path = scope.get("path", "")
method = scope.get("method", "")
limit = None
if path == self.mcp_path and method == "POST" and self.mcp_limit > 0:
limit = self.mcp_limit
elif path == "/download" and method == "GET" and self.download_limit > 0:
limit = self.download_limit
if limit is not None:
ip = self._client_ip_from_scope(scope)
self._hits[ip] = self._prune(self._hits[ip], now)
if len(self._hits[ip]) >= limit:
retry_after = int(self.window - (now - self._hits[ip][0])) + 1
response = JSONResponse(
{"detail": "Rate limit exceeded. Try again later."},
status_code=429,
headers={"Retry-After": str(max(retry_after, 1))},
)
return await response(scope, receive, send)
self._hits[ip].append(now)
return await self.app(scope, receive, send)
def _client_ip_from_scope(self, scope) -> str:
headers = dict(scope.get("headers", []))
forwarded = headers.get(b"x-forwarded-for")
if forwarded:
return forwarded.decode().split(",")[0].strip()
client = scope.get("client")
return client[0] if client else "unknown"
class ToolObserverMiddleware(Middleware):
"""Log every MCP tool invocation and its result."""
async def on_call_tool(self, context: MiddlewareContext, call_next):
tool_name = getattr(context.message, "name", "unknown")
arguments = getattr(context.message, "arguments", {})
logger.info("LLM REQUEST | Tool: '%s' | Args: %s", tool_name, arguments)
try:
result = await call_next(context)
content = getattr(result, "content", result)
preview = str(content)[:250].replace("\n", " ")
logger.info("SERVER RESPONSE | Tool: '%s' completed | Preview: %s", tool_name, preview)
return result
except Exception as exc:
logger.error("SERVER ERROR | Tool '%s' crashed: %s", tool_name, exc)
raise
mcp.add_middleware(ToolObserverMiddleware())
middleware_config = [
StarletteMiddleware(
CORSMiddleware,
allow_origins=CORS_ORIGINS,
allow_methods=["GET", "POST", "DELETE", "OPTIONS"],
allow_headers=[
"mcp-protocol-version",
"mcp-session-id",
"Authorization",
"Content-Type",
],
expose_headers=["mcp-session-id"],
),
StarletteMiddleware(
RateLimitMiddleware,
mcp_limit=MCP_RATE_LIMIT,
download_limit=DOWNLOAD_RATE_LIMIT,
window=RATE_LIMIT_WINDOW,
mcp_path=MCP_ENDPOINT,
),
]
# ---------------------------------------------------------------------------
# MCP tool — generate_download_link
# ---------------------------------------------------------------------------
@mcp.tool()
def generate_download_link(path: str, ttl_minutes: int = DOWNLOAD_DEFAULT_TTL) -> str:
"""Generate a secure, single-use, temporary download link for a file.
The link points at the ``/download`` HTTP endpoint on this server.
Each link is valid for exactly one request; after that the underlying token
is destroyed. Tokens also expire after *ttl_minutes* even if unused.
Parameters
----------
path : str
Absolute path to the file. Must reside inside the configured
workspace directory (default ``/workspace``).
ttl_minutes : int, optional
How many minutes the link stays valid (default: 10).
Clamped to the range ``DOWNLOAD_MIN_TTL`` – ``DOWNLOAD_MAX_TTL``.
Returns
-------
str
A full URL string on success, or an error message prefixed with
``Error:``.
"""
try:
abs_path = _validate_download_path(path)
except (PermissionError, FileNotFoundError, IsADirectoryError) as exc:
return f"Error: {exc}"
_purge_expired_tokens()
effective_ttl = max(DOWNLOAD_MIN_TTL, min(ttl_minutes, DOWNLOAD_MAX_TTL))
token = secrets.token_urlsafe(32)
expires_at = time.time() + (effective_ttl * 60)
TOKEN_STORE[token] = {
"file_path": abs_path,
"expires_at": expires_at,
}
url = f"{PUBLIC_BASE_URL}/download?t={token}"
logger.info("Generated download link for %s (ttl=%dm)", abs_path, effective_ttl)
return url
# ---------------------------------------------------------------------------
# Entry point
# ---------------------------------------------------------------------------
if __name__ == "__main__":
mcp.run(transport="http", host=HOST, port=PORT, middleware=middleware_config)