diff --git a/.github/workflows/test-unit.yml b/.github/workflows/test-unit.yml index 71eb0958bec..520d6f69dec 100644 --- a/.github/workflows/test-unit.yml +++ b/.github/workflows/test-unit.yml @@ -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: >- diff --git a/litellm/proxy/anthropic_endpoints/skills_endpoints.py b/litellm/proxy/anthropic_endpoints/skills_endpoints.py index 8e72754bbd1..62b21b7a9d4 100644 --- a/litellm/proxy/anthropic_endpoints/skills_endpoints.py +++ b/litellm/proxy/anthropic_endpoints/skills_endpoints.py @@ -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 diff --git a/litellm/skills/main.py b/litellm/skills/main.py index 61d21e7bb84..ef6f8176856 100644 --- a/litellm/skills/main.py +++ b/litellm/skills/main.py @@ -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, diff --git a/tests/test_litellm/skills/test_native_skills.py b/tests/test_litellm/skills/test_native_skills.py index 7edfcc067f6..461732c5118 100644 --- a/tests/test_litellm/skills/test_native_skills.py +++ b/tests/test_litellm/skills/test_native_skills.py @@ -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( {