365 lines
13 KiB
Python
365 lines
13 KiB
Python
#!/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)
|