mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-28 01:32:17 +00:00
Merge b0869568b0 into 860bc7811d
This commit is contained in:
commit
f813d3ecdf
3 changed files with 189 additions and 21 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue