from __future__ import annotations import base64 import binascii import hashlib import json import os import re from dataclasses import dataclass from typing import Any, Literal from urllib.parse import quote import httpx from fastapi import FastAPI from fastapi.routing import APIRoute from mcp.server.fastmcp import FastMCP from pydantic import BaseModel, Field BodyType = Literal["json", "form", "text"] PATH_PARAMETER = re.compile(r"{(?P[^}:]+)(?::(?P[^}]+))?}") INLINE_RESPONSE_LIMIT = int(os.getenv("MCP_INLINE_RESPONSE_MAX_BYTES", str(1024 * 1024))) TEXT_PREVIEW_LIMIT = int(os.getenv("MCP_TEXT_PREVIEW_MAX_CHARS", "12000")) RESTRICTED_HEADERS = { "connection", "content-length", "host", "transfer-encoding", } class EncodedFile(BaseModel): field_name: str = Field(default="files", description="Multipart form field name.") filename: str content_base64: str = Field(description="Base64-encoded file content.") content_type: str = "application/octet-stream" @dataclass(frozen=True) class RouteSource: namespace: str app: FastAPI path_prefix: str = "" def create_mcp_server( target_app: FastAPI, route_sources: tuple[RouteSource, ...], ) -> tuple[FastMCP, list[dict[str, Any]]]: server = FastMCP( "Unified Media Studio", instructions=( "Each tool maps to one Unified Media Studio HTTP endpoint. " "Use path_params for URL placeholders, query for query parameters, " "body for JSON or form data, headers for authentication, and files " "for base64-encoded multipart uploads." ), host="0.0.0.0", stateless_http=True, json_response=True, streamable_http_path="/", ) catalog: list[dict[str, Any]] = [] used_names: set[str] = set() for source in route_sources: for route in source.app.routes: if not isinstance(route, APIRoute): continue full_path = _join_paths(source.path_prefix, route.path) for method in sorted(route.methods or []): if method in {"HEAD", "OPTIONS"}: continue tool_name = _tool_name(source.namespace, method, route, used_names) description = _tool_description(method, full_path, route) tool = _make_route_tool(target_app, method, full_path) server.tool(name=tool_name, description=description)(tool) catalog.append( { "name": tool_name, "method": method, "path": full_path, "namespace": source.namespace, "tags": list(route.tags or []), } ) return server, catalog def _make_route_tool(target_app: FastAPI, method: str, path_template: str): async def invoke_endpoint( path_params: dict[str, Any] | None = None, query: dict[str, Any] | None = None, body: Any = None, body_type: BodyType = "json", headers: dict[str, str] | None = None, files: list[EncodedFile] | None = None, ) -> dict[str, Any]: """ Invoke the mapped Studio endpoint. Path placeholders belong in path_params. Query-string values belong in query. Use body_type="form" for form fields and files for multipart uploads. Existing endpoint authentication headers are forwarded. """ request_path = _render_path(path_template, path_params or {}) request_headers = _clean_headers(headers or {}) request_kwargs = _request_kwargs(body, body_type, files or []) transport = httpx.ASGITransport(app=target_app, raise_app_exceptions=False) async with httpx.AsyncClient( transport=transport, base_url="http://studio.internal", follow_redirects=True, timeout=None, ) as client: async with client.stream( method, request_path, params=query or None, headers=request_headers or None, **request_kwargs, ) as response: return await _serialize_response(response, method, request_path) invoke_endpoint.__name__ = f"invoke_{method.lower()}_{_slug(path_template)}" return invoke_endpoint def _request_kwargs( body: Any, body_type: BodyType, files: list[EncodedFile], ) -> dict[str, Any]: if files: multipart_files = [] for item in files: try: content = base64.b64decode(item.content_base64, validate=True) except (ValueError, binascii.Error) as exc: raise ValueError(f"Invalid base64 content for {item.filename}") from exc multipart_files.append( ( item.field_name, (item.filename, content, item.content_type), ) ) return { "data": _form_items(body), "files": multipart_files, } if body is None: return {} if body_type == "json": return {"json": body} if body_type == "form": return {"data": _form_items(body)} if isinstance(body, str): return {"content": body} return {"content": json.dumps(body)} def _form_items(body: Any) -> dict[str, str | list[str]]: if body is None: return {} if not isinstance(body, dict): raise ValueError("Form bodies must be JSON objects") items: dict[str, str | list[str]] = {} for key, value in body.items(): values = value if isinstance(value, list) else [value] encoded_values: list[str] = [] for item in values: if isinstance(item, (dict, list)): encoded = json.dumps(item) elif item is None: encoded = "" elif isinstance(item, bool): encoded = str(item).lower() else: encoded = str(item) encoded_values.append(encoded) items[str(key)] = encoded_values if isinstance(value, list) else encoded_values[0] return items def _render_path(path_template: str, path_params: dict[str, Any]) -> str: expected = {match.group("name") for match in PATH_PARAMETER.finditer(path_template)} missing = sorted(expected - path_params.keys()) if missing: raise ValueError(f"Missing path parameters: {', '.join(missing)}") def replace(match: re.Match[str]) -> str: name = match.group("name") converter = match.group("converter") safe = "/" if converter == "path" else "" return quote(str(path_params[name]), safe=safe) return PATH_PARAMETER.sub(replace, path_template) async def _serialize_response( response: httpx.Response, method: str, request_path: str, ) -> dict[str, Any]: content_type = response.headers.get("content-type", "").split(";", 1)[0].lower() result: dict[str, Any] = { "status_code": response.status_code, "method": method, "path": request_path, "content_type": content_type or None, } for header in ("content-disposition", "location"): if header in response.headers: result[header.replace("-", "_")] = response.headers[header] declared_size = _content_length(response.headers.get("content-length")) if declared_size is not None and declared_size > INLINE_RESPONSE_LIMIT: result["size_bytes"] = declared_size result["body_omitted"] = True return result content = await _read_bounded(response, INLINE_RESPONSE_LIMIT + 1) if not content: result["body"] = None return result if len(content) > INLINE_RESPONSE_LIMIT: result["size_bytes_at_least"] = len(content) result["body_omitted"] = True if _is_text_content(content_type): result["preview"] = content[:TEXT_PREVIEW_LIMIT].decode("utf-8", errors="replace") return result if content_type == "application/json" or content_type.endswith("+json"): try: result["body"] = json.loads(content) except (UnicodeDecodeError, ValueError): result["body"] = content[:TEXT_PREVIEW_LIMIT].decode("utf-8", errors="replace") return result if _is_text_content(content_type): result["body"] = content[:TEXT_PREVIEW_LIMIT].decode("utf-8", errors="replace") return result result["size_bytes"] = len(content) result["body_base64"] = base64.b64encode(content).decode("ascii") return result async def _read_bounded(response: httpx.Response, limit: int) -> bytes: content = bytearray() async for chunk in response.aiter_bytes(): remaining = limit - len(content) if remaining <= 0: break content.extend(chunk[:remaining]) if len(content) >= limit: break return bytes(content) def _content_length(value: str | None) -> int | None: if value is None: return None try: return max(0, int(value)) except ValueError: return None def _tool_name( namespace: str, method: str, route: APIRoute, used_names: set[str], ) -> str: base = _slug(f"{namespace}_{method}_{route.name or route.path}") if len(base) > 96: digest = hashlib.sha1(base.encode("utf-8")).hexdigest()[:10] base = f"{base[:85]}_{digest}" candidate = base suffix = 2 while candidate in used_names: candidate = f"{base}_{suffix}" suffix += 1 used_names.add(candidate) return candidate def _tool_description(method: str, full_path: str, route: APIRoute) -> str: summary = route.summary or route.description or route.name or "Studio API endpoint" summary = " ".join(summary.split()) return ( f"{method} {full_path}. {summary} " "Arguments: path_params, query, body, body_type, headers, and base64 files." ) def _join_paths(prefix: str, path: str) -> str: if not prefix: return path or "/" if path == "/": return f"{prefix.rstrip('/')}/" return f"{prefix.rstrip('/')}/{path.lstrip('/')}" def _clean_headers(headers: dict[str, str]) -> dict[str, str]: return { str(key): str(value) for key, value in headers.items() if str(key).lower() not in RESTRICTED_HEADERS } def _is_text_content(content_type: str) -> bool: return ( content_type.startswith("text/") or content_type in {"application/javascript", "application/xml"} or content_type.endswith("+xml") ) def _slug(value: str) -> str: cleaned = re.sub(r"[^a-zA-Z0-9_-]+", "_", value).strip("_").lower() return cleaned or "endpoint"