mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
test(proxy): cover native Skills error paths
This commit is contained in:
parent
1c05c5afb5
commit
0324249f99
4 changed files with 210 additions and 4 deletions
8
.github/workflows/test-unit.yml
vendored
8
.github/workflows/test-unit.yml
vendored
|
|
@ -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: >-
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
{
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue