test(proxy): cover native Skills error paths

This commit is contained in:
ymuichiro 2026-08-20 23:42:03 +09:00
parent 1c05c5afb5
commit 0324249f99
4 changed files with 210 additions and 4 deletions

View file

@ -124,6 +124,14 @@ jobs:
timeout-minutes: 20
job-timeout-minutes: 60
- shard: skills
artifact-name: skills
test-path: "tests/test_litellm/skills"
workers: 2
reruns: 2
timeout-minutes: 20
job-timeout-minutes: 60
- shard: proxy-auth
artifact-name: proxy-auth
test-path: >-

View file

@ -425,7 +425,9 @@ async def delete_skill(
SkillRouteType = Literal["acreate_skill", "alist_skills", "aget_skill", "adelete_skill"]
async def _native_skill_data(request: Request, operation: str) -> dict[str, object]:
async def _native_skill_data(
request: Request, operation: str
) -> dict[str, object]: # mutable-ok: proxy processing mutates routing data
body: Final = await convert_upload_files_to_file_data(await get_request_body(request))
model: Final = extract_model_param(request, body)
custom_llm_provider: Final = (
@ -434,6 +436,7 @@ async def _native_skill_data(request: Request, operation: str) -> dict[str, obje
data: Final = dict(request.query_params) # mutable-ok: proxy processing mutates route data
data.update(body)
data.update(request.path_params)
data.pop("model", None)
if model:
data["model"] = model
data["custom_llm_provider"] = custom_llm_provider

View file

@ -72,7 +72,7 @@ def _azure_skills_api_base(api_base: str | None) -> str | None:
def _native_skill_request(
operation: str,
request_data: dict[str, Any],
request_data: dict[str, Any], # mutable-ok: logging and SDK dispatch consume request data
custom_llm_provider: str,
litellm_params: GenericLiteLLMParams,
logging_obj: LiteLLMLoggingObj,

View file

@ -1,9 +1,12 @@
import json
import sys
from types import SimpleNamespace
from typing import Any
from unittest.mock import patch
import httpx
import pytest
from fastapi import Response
from fastapi.routing import APIRoute
from openai import AsyncOpenAI
from starlette.requests import Request
@ -11,15 +14,17 @@ from starlette.requests import Request
import litellm
from litellm.files.main import openai_files_instance
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from litellm.proxy.openai_files_endpoints.common_utils import extract_model_param
from litellm.proxy.anthropic_endpoints import skills_endpoints
from litellm.proxy.anthropic_endpoints.skills_endpoints import (
_native_skill_data,
)
from litellm.proxy.anthropic_endpoints.skills_endpoints import (
router as skills_router,
)
from litellm.proxy.openai_files_endpoints.common_utils import extract_model_param
from litellm.router import Router
from litellm.skills.main import _azure_skills_api_base
from litellm.skills.main import _azure_skills_api_base, _native_skill_request
from litellm.types.router import GenericLiteLLMParams
SKILL = {
"id": "skill_1",
@ -60,6 +65,75 @@ def _mock_response(request: httpx.Request) -> httpx.Response:
return httpx.Response(200, json=SKILL)
def _native_request(
body: dict[str, Any] | None = None,
*,
method: str = "POST",
path: str = "/v1/skills/skill_1",
query_string: bytes = b"",
path_params: dict[str, str] | None = None,
) -> Request:
payload = json.dumps(body).encode() if body is not None else b""
async def receive() -> dict[str, Any]:
return {"type": "http.request", "body": payload, "more_body": False}
return Request(
{
"type": "http",
"method": method,
"path": path,
"headers": [(b"content-type", b"application/json")],
"query_string": query_string,
"path_params": path_params or {"skill_id": "skill_1"},
},
receive,
)
@pytest.fixture
def native_endpoint_harness(monkeypatch: pytest.MonkeyPatch) -> tuple[type, dict[str, Any]]:
captured: dict[str, Any] = {}
class FakeProcessor:
result: Any = {"ok": True}
error: Exception | None = None
handled_error: Exception = RuntimeError("handled")
def __init__(self, data: dict[str, Any]) -> None:
captured["data"] = data
async def base_process_llm_request(self, **kwargs: Any) -> Any:
captured["kwargs"] = kwargs
if self.error is not None:
raise self.error
return self.result
async def _handle_llm_api_exception(self, **kwargs: Any) -> Exception:
captured["error"] = kwargs["e"]
return self.handled_error
monkeypatch.setattr(skills_endpoints, "ProxyBaseLLMRequestProcessing", FakeProcessor)
monkeypatch.setitem(
sys.modules,
"litellm.proxy.proxy_server",
SimpleNamespace(
general_settings={},
llm_router=None,
proxy_config=None,
proxy_logging_obj=None,
select_data_generator=None,
user_api_base=None,
user_max_tokens=None,
user_model=None,
user_request_timeout=None,
user_temperature=None,
version="test",
),
)
return FakeProcessor, captured
async def _call(operation: str, client: AsyncOpenAI, provider: str = "openai") -> Any:
common = {
"custom_llm_provider": provider,
@ -252,6 +326,127 @@ async def test_native_endpoint_model_priority(
assert data["_skill_operation"] == "update"
@pytest.mark.asyncio
async def test_native_endpoint_drops_invalid_body_model(native_endpoint_harness: tuple[type, dict[str, Any]]) -> None:
_, captured = native_endpoint_harness
endpoint = skills_endpoints._native_skill_route("update", "acreate_skill")
result = await endpoint(
_native_request({"model": {"invalid": "type"}}),
Response(),
object(),
)
assert result == {"ok": True}
assert "model" not in captured["data"]
assert captured["data"]["skill_id"] == "skill_1"
@pytest.mark.asyncio
async def test_native_endpoint_returns_provider_content_response(
native_endpoint_harness: tuple[type, dict[str, Any]],
) -> None:
processor, _ = native_endpoint_harness
processor.result = SimpleNamespace(
response=httpx.Response(
206,
content=b"skill archive",
headers={"content-type": "application/zip", "x-test": "present"},
)
)
result = await skills_endpoints._native_skill_endpoint(
_native_request(method="GET", path="/v1/skills/skill_1/content"),
Response(),
object(),
"content",
"aget_skill",
)
assert isinstance(result, Response)
assert result.status_code == 206
assert result.body == b"skill archive"
assert result.headers["content-type"] == "application/zip"
assert result.headers["x-test"] == "present"
@pytest.mark.asyncio
async def test_native_endpoint_maps_invalid_content_response(
native_endpoint_harness: tuple[type, dict[str, Any]],
) -> None:
processor, captured = native_endpoint_harness
processor.result = SimpleNamespace(response=object())
processor.handled_error = RuntimeError("invalid content response")
with pytest.raises(RuntimeError, match="invalid content response"):
await skills_endpoints._native_skill_endpoint(
_native_request(method="GET", path="/v1/skills/skill_1/content"),
Response(),
object(),
"content",
"aget_skill",
)
assert isinstance(captured["error"], TypeError)
@pytest.mark.asyncio
async def test_native_endpoint_maps_provider_error(
native_endpoint_harness: tuple[type, dict[str, Any]],
) -> None:
processor, captured = native_endpoint_harness
processor.error = ValueError("provider failure")
processor.handled_error = RuntimeError("mapped provider failure")
with pytest.raises(RuntimeError, match="mapped provider failure"):
await skills_endpoints._native_skill_endpoint(
_native_request({"skill_id": "skill_1"}),
Response(),
object(),
"update",
"acreate_skill",
)
assert isinstance(captured["error"], ValueError)
def test_native_skill_request_rejects_azure_without_api_base() -> None:
with patch(
"litellm.llms.azure.common_utils.get_azure_credentials",
return_value=SimpleNamespace(api_base=None, api_key="test"),
):
with pytest.raises(ValueError, match="api_base is required"):
_native_skill_request(
"get",
{"skill_id": "skill_1"},
"azure",
GenericLiteLLMParams(api_key="test"),
None, # type: ignore[arg-type]
None,
False,
)
def test_native_skill_request_rejects_uninitialized_openai_client() -> None:
with (
patch(
"litellm.llms.openai.common_utils.get_openai_credentials",
return_value=SimpleNamespace(api_base=None, api_key="test", organization=None),
),
patch.object(openai_files_instance, "get_openai_client", return_value=None),
):
with pytest.raises(ValueError, match="client is not initialized"):
_native_skill_request(
"get",
{"skill_id": "skill_1"},
"openai",
GenericLiteLLMParams(api_key="test"),
None, # type: ignore[arg-type]
None,
False,
)
def test_extract_model_param_ignores_non_string_body_model() -> None:
request = Request(
{