diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 20c114a2f3e..779c1bddc87 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -904,7 +904,7 @@ async def _openapi_spec_health( return "unhealthy", f"OpenAPI specification request failed (HTTP {exc.response.status_code})" except HTTPResponseLimitError as exc: return "unknown", f"OpenAPI specification probe refused: {exc}" - except (httpx.RequestError, ValueError, OSError) as exc: + except (httpx.RequestError, TypeError, ValueError, OSError) as exc: return "unhealthy", f"OpenAPI specification could not be loaded ({type(exc).__name__})" return "healthy", None diff --git a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py index 1247ff1ac28..c066fc689ec 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py @@ -9,6 +9,7 @@ import os import re from collections.abc import Mapping, Sequence from pathlib import PurePosixPath +from types import ModuleType from typing import Any, Final, TypedDict from urllib.parse import quote @@ -58,6 +59,30 @@ from litellm.proxy._experimental.mcp_server.tool_registry import ( from litellm.types.mcp import MCPAuthType, credential_redirect_hook, custom_credential_slot +def _import_yaml() -> ModuleType: + """Import and return the yaml module, raising a clear error if missing.""" + try: + import yaml as _yaml + + return _yaml + except ImportError: + raise ImportError( + "PyYAML is required to parse YAML OpenAPI specs. Install it with: pip install pyyaml" + ) from None + + +def _load_yaml_mapping(text: str) -> dict[str, Any]: + """Parse YAML text, requiring a mapping at the document root.""" + yaml_mod = _import_yaml() + try: + parsed = yaml_mod.safe_load(text) + except yaml_mod.YAMLError as exc: + raise ValueError(f"Invalid YAML OpenAPI spec: {exc}") from exc + if not isinstance(parsed, dict): + raise TypeError("Invalid OpenAPI spec: expected a JSON/YAML mapping at the document root") + return parsed + + class _OpenAPIJSONSchema(TypedDict, total=False): properties: Mapping[str, object] type: ReadOnly[str] @@ -164,6 +189,48 @@ def load_openapi_spec(filepath: str) -> dict[str, Any]: return asyncio.run(load_openapi_spec_async(filepath)) +def _is_yaml_content(filepath: str, content_type: str | None = None) -> bool: + """Determine if the content should be parsed as YAML.""" + # Check file extension + lower = filepath.lower() + if lower.endswith((".yaml", ".yml")): + return True + # Check Content-Type header + return bool(content_type and "yaml" in content_type) + + +def _load_local_openapi_spec(filepath: str) -> dict[str, Any]: + """Read a local OpenAPI spec file, parsing YAML or JSON.""" + if not os.path.exists(filepath): + raise FileNotFoundError(f"OpenAPI spec not found at {filepath}") + with open(filepath, "r", encoding="utf-8") as f: + content = f.read() + if _is_yaml_content(filepath): + return _load_yaml_mapping(content) + try: + return json.loads(content) + except ValueError as json_exc: + try: + return _load_yaml_mapping(content) + except (TypeError, ValueError): + raise json_exc from None + + +def _load_remote_openapi_spec(text: str, as_yaml: bool) -> dict[str, Any]: + """Parse a fetched spec body. Runs in a worker thread: YAML parsing is + synchronous CPU work with no nesting/alias limits, so it must not run on + the event loop where a pathological document could stall the proxy.""" + if as_yaml: + return _load_yaml_mapping(text) + try: + return json.loads(text) + except ValueError as json_exc: + try: + return _load_yaml_mapping(text) + except (TypeError, ValueError): + raise json_exc from None + + async def load_openapi_spec_async(filepath: str, *, max_bytes: int | None = None) -> dict[str, Any]: if filepath.startswith("http://") or filepath.startswith("https://"): client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP) @@ -173,14 +240,16 @@ async def load_openapi_spec_async(filepath: str, *, max_bytes: int | None = None else await async_safe_get(client, filepath, max_response_bytes=max_bytes) ) r.raise_for_status() - return r.json() - # fallback: local file - # Local filesystem path - if not os.path.exists(filepath): - raise FileNotFoundError(f"OpenAPI spec not found at {filepath}") - with open(filepath, "r", encoding="utf-8") as f: - return json.load(f) + content_type = r.headers.get("content-type", "") + # Try JSON first; fall back to YAML for specs served without + # proper Content-Type headers (common with raw GitHub URLs). + as_yaml = _is_yaml_content(filepath, content_type) + return await asyncio.to_thread(_load_remote_openapi_spec, r.text, as_yaml) + + # Local files go through a worker thread: the async path must not + # perform blocking disk I/O directly (ruff ASYNC230). + return await asyncio.to_thread(_load_local_openapi_spec, filepath) def get_base_url(spec: Mapping[str, Any], spec_path: str | None = None) -> str: diff --git a/tests/mcp_tests/test_openapi_spec_path_url.py b/tests/mcp_tests/test_openapi_spec_path_url.py index 17a0022046e..5b231aac09d 100644 --- a/tests/mcp_tests/test_openapi_spec_path_url.py +++ b/tests/mcp_tests/test_openapi_spec_path_url.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import Any, Dict +from typing import Any import httpx import pytest @@ -33,7 +33,7 @@ class _FakeAsyncHTTPHandler: def test_load_openapi_spec_supports_http_url(monkeypatch: pytest.MonkeyPatch) -> None: url = "http://example.local/openapi.json" - expected: Dict[str, Any] = { + expected: dict[str, Any] = { "openapi": "3.0.0", "info": {"title": "Test API", "version": "1.0.0"}, "paths": {}, @@ -44,7 +44,7 @@ def test_load_openapi_spec_supports_http_url(monkeypatch: pytest.MonkeyPatch) -> resp = httpx.Response(status_code=200, json=expected, request=req) calls = {"get_async_httpx_client": 0} - handler_holder: Dict[str, Any] = {} + handler_holder: dict[str, Any] = {} def fake_get_async_httpx_client(*args, **kwargs): calls["get_async_httpx_client"] += 1 @@ -56,9 +56,7 @@ def test_load_openapi_spec_supports_http_url(monkeypatch: pytest.MonkeyPatch) -> monkeypatch.setattr(gen, "get_async_httpx_client", fake_get_async_httpx_client) # Bypass SSRF validation in test (example.local doesn't resolve) - monkeypatch.setattr( - gen, "async_safe_get", lambda client, url, **kw: client.get(url) - ) + monkeypatch.setattr(gen, "async_safe_get", lambda client, url, **kw: client.get(url)) # Fail loudly if someone reintroduces direct httpx.get() def boom(*args, **kwargs): @@ -73,10 +71,8 @@ def test_load_openapi_spec_supports_http_url(monkeypatch: pytest.MonkeyPatch) -> assert handler_holder["handler"].calls == 1 -def test_load_openapi_spec_supports_local_file_path( - tmp_path, monkeypatch: pytest.MonkeyPatch -) -> None: - expected: Dict[str, Any] = { +def test_load_openapi_spec_supports_local_file_path(tmp_path, monkeypatch: pytest.MonkeyPatch) -> None: + expected: dict[str, Any] = { "openapi": "3.0.0", "info": {"title": "Local API", "version": "1.0.0"}, "paths": {}, @@ -90,11 +86,114 @@ def test_load_openapi_spec_supports_local_file_path( # For local files, shared client must NOT be used. def boom_client(*args, **kwargs): - raise AssertionError( - "get_async_httpx_client() must not be called for local file paths" - ) + raise AssertionError("get_async_httpx_client() must not be called for local file paths") monkeypatch.setattr(gen, "get_async_httpx_client", boom_client) spec = gen.load_openapi_spec(str(p)) assert spec == expected + + +def test_load_openapi_spec_supports_yaml_file(tmp_path, monkeypatch: pytest.MonkeyPatch) -> None: + """YAML OpenAPI spec files should be parsed correctly.""" + expected: dict[str, Any] = { + "openapi": "3.0.0", + "info": {"title": "YAML API", "version": "1.0.0"}, + "paths": {}, + } + + p = tmp_path / "openapi.yaml" + p.write_text( + "openapi: '3.0.0'\ninfo:\n title: YAML API\n version: '1.0.0'\npaths: {}", + encoding="utf-8", + ) + + def boom_client(*args, **kwargs): + raise AssertionError("get_async_httpx_client() must not be called for local file paths") + + monkeypatch.setattr(gen, "get_async_httpx_client", boom_client) + + spec = gen.load_openapi_spec(str(p)) + assert spec == expected + + +def test_load_openapi_spec_url_yaml_content_type( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """URL returning YAML with Content-Type: text/yaml should be parsed as YAML.""" + url = "http://example.local/openapi.yaml" + yaml_body = "openapi: '3.0.0'\ninfo:\n title: Remote YAML\n version: '1.0.0'\npaths: {}" + + req = httpx.Request("GET", url) + resp = httpx.Response( + status_code=200, + content=yaml_body.encode(), + headers={"content-type": "text/yaml"}, + request=req, + ) + + handler_holder: dict[str, Any] = {} + + def fake_get_async_httpx_client(*args, **kwargs): + h = _FakeAsyncHTTPHandler(resp, expected_url=url) + handler_holder["handler"] = h + return h + + monkeypatch.setattr(gen, "get_async_httpx_client", fake_get_async_httpx_client) + monkeypatch.setattr(gen, "async_safe_get", lambda client, url, **kw: client.get(url)) + + spec = gen.load_openapi_spec(url) + assert spec["info"]["title"] == "Remote YAML" + + +def test_load_openapi_spec_url_yaml_fallback( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """URL with .yaml extension but no YAML Content-Type should still parse as YAML.""" + url = "https://raw.githubusercontent.com/firefly-iii/api-docs/refs/heads/v6.6.6/dist/firefly-iii-v6.6.6-v1.yaml" + yaml_body = "openapi: '3.0.0'\ninfo:\n title: Firefly\n version: '6.6.6'\npaths: {}" + + req = httpx.Request("GET", url) + resp = httpx.Response( + status_code=200, + content=yaml_body.encode(), + headers={"content-type": "text/plain"}, + request=req, + ) + + handler_holder: dict[str, Any] = {} + + def fake_get_async_httpx_client(*args, **kwargs): + h = _FakeAsyncHTTPHandler(resp, expected_url=url) + handler_holder["handler"] = h + return h + + monkeypatch.setattr(gen, "get_async_httpx_client", fake_get_async_httpx_client) + monkeypatch.setattr(gen, "async_safe_get", lambda client, url, **kw: client.get(url)) + + spec = gen.load_openapi_spec(url) + assert spec["info"]["title"] == "Firefly" + + +def test_load_openapi_spec_url_plain_text_body_raises( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """A 200 body that is neither JSON nor a YAML mapping must raise, not return a string.""" + url = "https://example.local/openapi.json" + + req = httpx.Request("GET", url) + resp = httpx.Response( + status_code=200, + content=b"secret invalid JSON body", + headers={"content-type": "text/plain"}, + request=req, + ) + + def fake_get_async_httpx_client(*args, **kwargs): + return _FakeAsyncHTTPHandler(resp, expected_url=url) + + monkeypatch.setattr(gen, "get_async_httpx_client", fake_get_async_httpx_client) + monkeypatch.setattr(gen, "async_safe_get", lambda client, url, **kw: client.get(url)) + + with pytest.raises(ValueError, match="Expecting value"): + gen.load_openapi_spec(url)