diff --git a/.github/workflows/test-unit.yml b/.github/workflows/test-unit.yml index c23678c51ae..467b3d8ef28 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 9390bf4c537..62b21b7a9d4 100644 --- a/litellm/proxy/anthropic_endpoints/skills_endpoints.py +++ b/litellm/proxy/anthropic_endpoints/skills_endpoints.py @@ -2,8 +2,10 @@ Anthropic Skills API endpoints - /v1/skills """ -from typing import Final +from collections.abc import Awaitable, Callable +from typing import Annotated, Final, Literal +import httpx import orjson from fastapi import APIRouter, Depends, Request, Response @@ -13,11 +15,10 @@ from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessin from litellm.proxy.common_utils.http_parsing_utils import ( convert_upload_files_to_file_data, get_form_data, + get_request_body, ) -from litellm.types.llms.anthropic_skills import ( - DeleteSkillResponse, - ListSkillsResponse, - Skill, +from litellm.proxy.openai_files_endpoints.common_utils import ( + extract_model_param, ) router: Final = APIRouter() @@ -27,7 +28,6 @@ router: Final = APIRouter() "/v1/skills", tags=["[beta] Anthropic Skills API"], dependencies=[Depends(user_api_key_auth)], - response_model=Skill, ) async def create_skill( fastapi_response: Response, @@ -87,6 +87,7 @@ async def create_skill( model: Final = data.get("model") or request.query_params.get("model") or request.headers.get("x-litellm-model") if model: data["model"] = model + data["_skill_operation"] = "create" if "custom_llm_provider" not in data: data["custom_llm_provider"] = custom_llm_provider @@ -125,7 +126,6 @@ async def create_skill( "/v1/skills", tags=["[beta] Anthropic Skills API"], dependencies=[Depends(user_api_key_auth)], - response_model=ListSkillsResponse, ) async def list_skills( fastapi_response: Response, @@ -176,7 +176,7 @@ async def list_skills( # Read request body body: Final = await request.body() - data: Final = orjson.loads(body) if body else {} + data: Final = {**dict(request.query_params), **(orjson.loads(body) if body else {})} # mutable-ok: pagination data # Use query params if not in body if "limit" not in data and limit is not None: @@ -190,6 +190,7 @@ async def list_skills( model: Final = data.get("model") or request.query_params.get("model") or request.headers.get("x-litellm-model") if model: data["model"] = model + data["_skill_operation"] = "list" # Set custom_llm_provider: body > query param > default if "custom_llm_provider" not in data: @@ -229,7 +230,6 @@ async def list_skills( "/v1/skills/{skill_id}", tags=["[beta] Anthropic Skills API"], dependencies=[Depends(user_api_key_auth)], - response_model=Skill, ) async def get_skill( skill_id: str, @@ -287,6 +287,7 @@ async def get_skill( model: Final = data.get("model") or request.query_params.get("model") or request.headers.get("x-litellm-model") if model: data["model"] = model + data["_skill_operation"] = "get" # Set custom_llm_provider: body > query param > default if "custom_llm_provider" not in data: @@ -326,7 +327,6 @@ async def get_skill( "/v1/skills/{skill_id}", tags=["[beta] Anthropic Skills API"], dependencies=[Depends(user_api_key_auth)], - response_model=DeleteSkillResponse, ) async def delete_skill( skill_id: str, @@ -386,6 +386,7 @@ async def delete_skill( model: Final = data.get("model") or request.query_params.get("model") or request.headers.get("x-litellm-model") if model: data["model"] = model + data["_skill_operation"] = "delete" # Set custom_llm_provider: body > query param > default if "custom_llm_provider" not in data: @@ -419,3 +420,114 @@ async def delete_skill( proxy_logging_obj=proxy_logging_obj, version=version, ) + + +SkillRouteType = Literal["acreate_skill", "alist_skills", "aget_skill", "adelete_skill"] + + +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 = ( + body.get("custom_llm_provider") or request.query_params.get("custom_llm_provider") or "openai" + ) + 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 + data["_skill_operation"] = operation + return data + + +async def _native_skill_endpoint( + request: Request, + fastapi_response: Response, + user_api_key_dict: UserAPIKeyAuth, + operation: str, + route_type: SkillRouteType, +) -> object: + from litellm.proxy.proxy_server import ( + general_settings, + llm_router, + proxy_config, + proxy_logging_obj, + select_data_generator, + user_api_base, + user_max_tokens, + user_model, + user_request_timeout, + user_temperature, + version, + ) + + data: Final = await _native_skill_data(request, operation) + processor: Final = ProxyBaseLLMRequestProcessing(data=data) + try: + result: Final = await processor.base_process_llm_request( + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=user_api_key_dict, + route_type=route_type, + proxy_logging_obj=proxy_logging_obj, + llm_router=llm_router, + general_settings=general_settings, + proxy_config=proxy_config, + select_data_generator=select_data_generator, + model=data.get("model"), + user_model=user_model, + user_temperature=user_temperature, + user_request_timeout=user_request_timeout, + user_max_tokens=user_max_tokens, + user_api_base=user_api_base, + version=version, + ) + if operation not in ("content", "version_content"): + return result + response: Final = getattr(result, "response", None) + if not isinstance(response, httpx.Response): + raise TypeError("Skills content response did not contain an HTTP response") + return Response(content=response.content, status_code=response.status_code, headers=response.headers) + except Exception as e: # noqa: BLE001 # proxy maps provider errors to the public exception contract + raise await processor._handle_llm_api_exception( + e=e, + user_api_key_dict=user_api_key_dict, + proxy_logging_obj=proxy_logging_obj, + version=version, + ) + + +def _native_skill_route(operation: str, route_type: SkillRouteType) -> Callable[..., Awaitable[object]]: + async def endpoint( + request: Request, + fastapi_response: Response, + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + ) -> object: + return await _native_skill_endpoint(request, fastapi_response, user_api_key_dict, operation, route_type) + + return endpoint + + +_NATIVE_SKILL_ROUTES: Final[tuple[tuple[str, str, str, SkillRouteType], ...]] = ( + ("POST", "/v1/skills/{skill_id}", "update", "acreate_skill"), + ("GET", "/v1/skills/{skill_id}/content", "content", "aget_skill"), + ("POST", "/v1/skills/{skill_id}/versions", "create_version", "acreate_skill"), + ("GET", "/v1/skills/{skill_id}/versions", "list_versions", "alist_skills"), + ("GET", "/v1/skills/{skill_id}/versions/{version}", "version", "aget_skill"), + ("DELETE", "/v1/skills/{skill_id}/versions/{version}", "delete_version", "adelete_skill"), + ("GET", "/v1/skills/{skill_id}/versions/{version}/content", "version_content", "aget_skill"), +) + +for method, path, operation, route_type in _NATIVE_SKILL_ROUTES: + router.add_api_route( + path, + _native_skill_route(operation, route_type), + methods=[method], # mutable-ok: FastAPI contract + name=f"{operation}_skill", + response_model=None, + tags=["[beta] OpenAI Skills API"], # mutable-ok: FastAPI contract + ) diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index ddfdb56ac2c..dd0e21ecf08 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -864,7 +864,7 @@ async def extract_file_creation_params( target_model_names = await _extract_target_model_names_from_form(request) # Extract model parameter - model: Final = _extract_model_param(request, request_body) + model: Final = extract_model_param(request, request_body) return FileCreationParams( target_storage=target_storage, @@ -1035,7 +1035,7 @@ async def validate_managed_id_requirement( ) -def _extract_model_param(request: "Request", request_body: dict) -> str | None: +def extract_model_param(request: "Request", request_body: Mapping[str, object]) -> str | None: """ Extract model parameter from request. @@ -1044,7 +1044,12 @@ def _extract_model_param(request: "Request", request_body: dict) -> str | None: 2. Query parameter (?model=) 3. Header (x-litellm-model) """ - return request_body.get("model") or request.query_params.get("model") or request.headers.get("x-litellm-model") + body_model: Final = request_body.get("model") + return ( + body_model + if isinstance(body_model, str) and body_model + else request.query_params.get("model") or request.headers.get("x-litellm-model") + ) # ============================================================================ diff --git a/litellm/skills/main.py b/litellm/skills/main.py index ae1ce150368..885d2cb2d76 100644 --- a/litellm/skills/main.py +++ b/litellm/skills/main.py @@ -5,8 +5,11 @@ Provides create, list, get, and delete operations for skills import asyncio import contextvars +import inspect from collections.abc import Coroutine from functools import partial +from operator import attrgetter +from types import MappingProxyType from typing import Any, Final import httpx @@ -30,11 +33,118 @@ from litellm.utils import ProviderConfigManager, client # Initialize HTTP handler base_llm_http_handler = BaseLLMHTTPHandler() DEFAULT_ANTHROPIC_API_BASE: Final = "https://api.anthropic.com/v1" +_NATIVE_SKILL_PROVIDERS: Final = frozenset({"openai", "azure"}) +_NATIVE_SKILL_OPERATIONS: Final = MappingProxyType( + { + "create": ("skills.create", ("files",)), + "list": ("skills.list", ("after", "limit", "order")), + "get": ("skills.retrieve", ("skill_id",)), + "update": ("skills.update", ("skill_id", "default_version")), + "delete": ("skills.delete", ("skill_id",)), + "content": ("skills.content.retrieve", ("skill_id",)), + "create_version": ("skills.versions.create", ("skill_id", "default", "files")), + "list_versions": ("skills.versions.list", ("skill_id", "after", "limit", "order")), + "version": ("skills.versions.retrieve", ("skill_id", "version")), + "delete_version": ("skills.versions.delete", ("skill_id", "version")), + "version_content": ("skills.versions.content.retrieve", ("skill_id", "version")), + } +) +_NATIVE_ONLY_SKILL_OPERATIONS: Final = frozenset(_NATIVE_SKILL_OPERATIONS) - {"create", "list", "get", "delete"} # Initialize LiteLLM skills handler (lazy - only used when custom_llm_provider="litellm") _litellm_skills_handler = None +def _azure_skills_api_base(api_base: str | None) -> str | None: + if api_base is None: + return None + url: Final = httpx.URL(api_base) + path: Final = url.path.rstrip("/") + suffix: Final = next( + ( + item + for item in ("/openai/v1/responses", "/openai/responses", "/openai/v1", "/openai") + if path.endswith(item) + ), + "", + ) + return str(url.copy_with(path=path[: -len(suffix)] if suffix else path, query=None)).rstrip("/") + + +def _validate_skill_operation(operation: str, custom_llm_provider: str) -> None: + if operation in _NATIVE_ONLY_SKILL_OPERATIONS and custom_llm_provider not in _NATIVE_SKILL_PROVIDERS: + raise ValueError(f"{operation} skills operation is only supported for OpenAI and Azure OpenAI") + + +def _native_skill_request( + operation: str, + request_data: dict[str, Any], # mutable-ok: logging and SDK dispatch consume request data + custom_llm_provider: str, + litellm_params: GenericLiteLLMParams, + logging_obj: LiteLLMLoggingObj, + litellm_call_id: str | None, + is_async: bool, +) -> object: + from litellm.files.main import azure_files_instance, openai_files_instance + from litellm.llms.azure.common_utils import get_azure_credentials + from litellm.llms.openai.common_utils import get_openai_credentials + + method_path, request_fields = _NATIVE_SKILL_OPERATIONS[operation] + extra_headers: Final = request_data.get("extra_headers") + headers: Final = extra_headers + params: Final = { # mutable-ok: SDK request parameters + field: value + for field, value in ( + *((field, request_data.get(field)) for field in request_fields), + ("extra_headers", headers), + ("extra_query", request_data.get("extra_query")), + ("extra_body", request_data.get("extra_body")), + ("timeout", request_data.get("timeout")), + ) + if value is not None + } + if custom_llm_provider == "openai": + openai_credentials: Final = get_openai_credentials( + api_base=litellm_params.api_base, + api_key=litellm_params.api_key, + organization=litellm_params.organization, + ) + sdk_client = openai_files_instance.get_openai_client( + api_key=openai_credentials.api_key, + api_base=openai_credentials.api_base, + timeout=request_data.get("timeout") or request_timeout, + max_retries=litellm_params.max_retries, + organization=openai_credentials.organization, + client=request_data.get("client"), + _is_async=is_async, + ) + else: + azure_credentials: Final = get_azure_credentials( + api_base=litellm_params.api_base, api_key=litellm_params.api_key + ) + api_base: Final = _azure_skills_api_base(azure_credentials.api_base) + if api_base is None: + raise ValueError("api_base is required for Azure OpenAI Skills") + sdk_client = azure_files_instance.get_azure_openai_client( + api_key=azure_credentials.api_key, + api_base=api_base, + api_version="v1", + client=request_data.get("client"), + litellm_params=litellm_params.model_dump(exclude_none=True), + _is_async=is_async, + ) + if sdk_client is None: + raise ValueError(f"{custom_llm_provider} client is not initialized") + logging_obj.update_from_kwargs( + kwargs=request_data, + model=None, + optional_params=params, + litellm_params={"litellm_call_id": litellm_call_id}, # mutable-ok: logging consumes request data + custom_llm_provider=custom_llm_provider, + ) + return attrgetter(method_path)(sdk_client)(**params) + + def _get_user_api_key_auth_from_kwargs(kwargs: dict[str, Any]) -> Any | None: for metadata_key in ("metadata", "litellm_metadata"): metadata = kwargs.get(metadata_key) @@ -170,6 +280,7 @@ def create_skill( # Determine provider if custom_llm_provider is None: custom_llm_provider = "anthropic" + _validate_skill_operation(kwargs.get("_skill_operation", "create"), custom_llm_provider) # Build create request create_request: Final[CreateSkillRequest] = {} @@ -195,6 +306,17 @@ def create_skill( litellm_call_id=litellm_call_id, ) + if custom_llm_provider in _NATIVE_SKILL_PROVIDERS: + return _native_skill_request( + kwargs.get("_skill_operation", "create"), + {**local_vars, **kwargs, **create_request}, # mutable-ok: Skills handlers consume request data + custom_llm_provider, + litellm_params, + litellm_logging_obj, + litellm_call_id, + _is_async, + ) + # Get provider config for external providers (Anthropic, etc.) skills_api_provider_config: BaseSkillsAPIConfig | None = ProviderConfigManager.get_provider_skills_api_config( provider=litellm.LlmProviders(custom_llm_provider), @@ -305,7 +427,7 @@ async def alist_skills( func_with_context: Final = partial(ctx.run, func) init_response: Final = await loop.run_in_executor(None, func_with_context) - if asyncio.iscoroutine(init_response): + if inspect.isawaitable(init_response): response = await init_response else: response = init_response @@ -359,6 +481,7 @@ def list_skills( # Determine provider if custom_llm_provider is None: custom_llm_provider = "anthropic" + _validate_skill_operation(kwargs.get("_skill_operation", "list"), custom_llm_provider) # Route to LiteLLM DB if custom_llm_provider="litellm_proxy" if custom_llm_provider == LlmProviders.LITELLM_PROXY.value: @@ -371,6 +494,17 @@ def list_skills( litellm_call_id=litellm_call_id, ) + if custom_llm_provider in _NATIVE_SKILL_PROVIDERS: + return _native_skill_request( + kwargs.get("_skill_operation", "list"), + {**local_vars, **kwargs}, # mutable-ok: Skills handlers consume request data + custom_llm_provider, + litellm_params, + litellm_logging_obj, + litellm_call_id, + _is_async, + ) + # Get provider config for external providers (Anthropic, etc.) skills_api_provider_config: BaseSkillsAPIConfig | None = ProviderConfigManager.get_provider_skills_api_config( provider=litellm.LlmProviders(custom_llm_provider), @@ -532,6 +666,7 @@ def get_skill( # Determine provider if custom_llm_provider is None: custom_llm_provider = "anthropic" + _validate_skill_operation(kwargs.get("_skill_operation", "get"), custom_llm_provider) # Route to LiteLLM DB if custom_llm_provider="litellm_proxy" if custom_llm_provider == LlmProviders.LITELLM_PROXY.value: @@ -543,6 +678,17 @@ def get_skill( litellm_call_id=litellm_call_id, ) + if custom_llm_provider in _NATIVE_SKILL_PROVIDERS: + return _native_skill_request( + kwargs.get("_skill_operation", "get"), + {**local_vars, **kwargs}, # mutable-ok: Skills handlers consume request data + custom_llm_provider, + litellm_params, + litellm_logging_obj, + litellm_call_id, + _is_async, + ) + # Get provider config for external providers (Anthropic, etc.) skills_api_provider_config: BaseSkillsAPIConfig | None = ProviderConfigManager.get_provider_skills_api_config( provider=litellm.LlmProviders(custom_llm_provider), @@ -696,6 +842,7 @@ def delete_skill( # Determine provider if custom_llm_provider is None: custom_llm_provider = "anthropic" + _validate_skill_operation(kwargs.get("_skill_operation", "delete"), custom_llm_provider) # Route to LiteLLM DB if custom_llm_provider="litellm_proxy" if custom_llm_provider == LlmProviders.LITELLM_PROXY.value: @@ -707,6 +854,17 @@ def delete_skill( litellm_call_id=litellm_call_id, ) + if custom_llm_provider in _NATIVE_SKILL_PROVIDERS: + return _native_skill_request( + kwargs.get("_skill_operation", "delete"), + {**local_vars, **kwargs}, # mutable-ok: Skills handlers consume request data + custom_llm_provider, + litellm_params, + litellm_logging_obj, + litellm_call_id, + _is_async, + ) + # Get provider config for external providers (Anthropic, etc.) skills_api_provider_config: BaseSkillsAPIConfig | None = ProviderConfigManager.get_provider_skills_api_config( provider=litellm.LlmProviders(custom_llm_provider), diff --git a/tests/test_litellm/skills/test_native_skills.py b/tests/test_litellm/skills/test_native_skills.py new file mode 100644 index 00000000000..61602822c3e --- /dev/null +++ b/tests/test_litellm/skills/test_native_skills.py @@ -0,0 +1,520 @@ +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 + +import litellm +from litellm.files.main import openai_files_instance +from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER +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, + _native_skill_request, + _validate_skill_operation, +) +from litellm.types.router import GenericLiteLLMParams + +SKILL = { + "id": "skill_1", + "created_at": 1, + "default_version": "1", + "description": "description", + "latest_version": "1", + "name": "test-skill", + "object": "skill", +} +VERSION = { + "id": "version_1", + "created_at": 1, + "description": "description", + "name": "test-skill", + "object": "skill.version", + "skill_id": "skill_1", + "version": "1", +} + + +def _mock_response(request: httpx.Request) -> httpx.Response: + path = request.url.path + if path.endswith("/content"): + return httpx.Response(200, content=b"skill archive", headers={"content-type": "application/zip"}) + if request.method == "DELETE" and "/versions/" in path: + return httpx.Response( + 200, json={"id": "skill_1", "deleted": True, "object": "skill.version.deleted", "version": "1"} + ) + if request.method == "DELETE": + return httpx.Response(200, json={"id": "skill_1", "deleted": True, "object": "skill.deleted"}) + if path.endswith("/versions") and request.method == "GET": + return httpx.Response(200, json={"object": "list", "data": [VERSION], "has_more": False}) + if "/versions/" in path or path.endswith("/versions"): + return httpx.Response(200, json=VERSION) + if path.endswith("/skills") and request.method == "GET": + return httpx.Response(200, json={"object": "list", "data": [SKILL], "has_more": False}) + 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, + "client": client, + "api_base": "https://resource.openai.azure.com/openai/v1" if provider == "azure" else None, + "extra_headers": {"x-test-header": "present"} if provider == "azure" else None, + "_skill_operation": operation, + } + if operation == "create": + return await litellm.acreate_skill(files=[("SKILL.md", b"skill")], **common) + if operation == "update": + return await litellm.acreate_skill(skill_id="skill_1", default_version=2, **common) + if operation == "create_version": + return await litellm.acreate_skill(skill_id="skill_1", files=[("SKILL.md", b"skill")], default=True, **common) + if operation in {"list", "list_versions"}: + return await litellm.alist_skills(skill_id="skill_1", after="cursor", limit=2, order="desc", **common) + if operation in {"get", "content", "version", "version_content"}: + return await litellm.aget_skill(skill_id="skill_1", version="1", **common) + return await litellm.adelete_skill(skill_id="skill_1", version="1", **common) + + +@pytest.mark.asyncio +async def test_openai_sdk_handles_every_native_skill_operation() -> None: + requests: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + return _mock_response(request) + + client = AsyncOpenAI( + api_key="test", + base_url="https://api.openai.test/v1", + http_client=httpx.AsyncClient(transport=httpx.MockTransport(handler)), + ) + operations = ( + "create", + "list", + "get", + "update", + "delete", + "content", + "create_version", + "list_versions", + "version", + "delete_version", + "version_content", + ) + try: + results = [await _call(operation, client) for operation in operations] + finally: + await GLOBAL_LOGGING_WORKER.flush() + await client.close() + + actual_requests = [(request.method, request.url.path) for request in requests] + expected_requests = [ + ("POST", "/v1/skills"), + ("GET", "/v1/skills"), + ("GET", "/v1/skills/skill_1"), + ("POST", "/v1/skills/skill_1"), + ("DELETE", "/v1/skills/skill_1"), + ("GET", "/v1/skills/skill_1/content"), + ("POST", "/v1/skills/skill_1/versions"), + ("GET", "/v1/skills/skill_1/versions"), + ("GET", "/v1/skills/skill_1/versions/1"), + ("DELETE", "/v1/skills/skill_1/versions/1"), + ("GET", "/v1/skills/skill_1/versions/1/content"), + ] + assert actual_requests == expected_requests, actual_requests + assert dict(requests[1].url.params) == {"after": "cursor", "limit": "2", "order": "desc"} + assert json.loads(requests[3].content) == {"default_version": 2} + assert b'name="files[]"' in requests[0].content + assert b'name="default"' in requests[6].content + assert b"SKILL.md" in requests[0].content + assert dict(requests[7].url.params) == {"after": "cursor", "limit": "2", "order": "desc"} + assert results[5].response.content == b"skill archive" + assert results[10].response.content == b"skill archive" + assert results[0].id == "skill_1" + assert results[1].data[0].id == "skill_1" + assert results[4].deleted is True + assert results[6].skill_id == "skill_1" + assert results[9].deleted is True + + +@pytest.mark.asyncio +async def test_azure_does_not_use_foundry_preview_header_with_existing_sdk_client() -> None: + requests: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + return _mock_response(request) + + client = AsyncOpenAI( + api_key="test", + base_url="https://resource.openai.azure.com/openai/v1", + http_client=httpx.AsyncClient(transport=httpx.MockTransport(handler)), + ) + try: + await _call("get", client, provider="azure") + finally: + await GLOBAL_LOGGING_WORKER.flush() + await client.close() + + assert "foundry-features" not in requests[0].headers + assert requests[0].headers["x-test-header"] == "present" + + +@pytest.mark.asyncio +async def test_router_model_configuration_overrides_request_provider() -> None: + requests: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + return _mock_response(request) + + client = AsyncOpenAI( + api_key="test", + base_url="https://api.openai.test/v1", + http_client=httpx.AsyncClient(transport=httpx.MockTransport(handler)), + ) + router = Router( + model_list=[ + { + "model_name": "skills-openai", + "litellm_params": { + "model": "openai/gpt-5", + "api_key": "test", + }, + } + ] + ) + try: + with patch.object(openai_files_instance, "get_openai_client", return_value=client) as get_client: + await router.acreate_skill( + model="skills-openai", + files=[("SKILL.md", b"skill")], + custom_llm_provider="anthropic", + ) + finally: + await GLOBAL_LOGGING_WORKER.flush() + await client.close() + + assert requests[0].url.path == "/v1/skills" + assert get_client.call_args.kwargs["api_key"] == "test" + + +@pytest.mark.parametrize( + ("body_model", "query", "header_model", "expected"), + [ + ("body", "model=query", "header", "body"), + ("", "model=query", "header", "query"), + (None, "model=query", "header", "query"), + (None, "", "header", "header"), + (None, "", None, None), + ], +) +@pytest.mark.asyncio +async def test_native_endpoint_model_priority( + body_model: object | None, + query: str, + header_model: str | None, + expected: str | None, +) -> None: + body = json.dumps({"model": body_model} if body_model is not None else {}).encode() + + async def receive() -> dict[str, Any]: + return {"type": "http.request", "body": body, "more_body": False} + + request = Request( + { + "type": "http", + "method": "POST", + "path": "/v1/skills/skill_1", + "headers": [ + (b"content-type", b"application/json"), + *(([(b"x-litellm-model", header_model.encode())]) if header_model else []), + ], + "query_string": query.encode(), + "path_params": {"skill_id": "skill_1"}, + }, + receive, + ) + + data = await _native_skill_data(request, "update") + + if expected is None: + assert "model" not in data + else: + assert data["model"] == expected + assert data["skill_id"] == "skill_1" + assert data["custom_llm_provider"] == "openai" + 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, + ) + + +@pytest.mark.parametrize( + "operation", + ["update", "content", "create_version", "list_versions", "version", "delete_version", "version_content"], +) +def test_native_only_skill_operations_reject_non_native_providers(operation: str) -> None: + with pytest.raises(ValueError, match="only supported for OpenAI and Azure OpenAI"): + _validate_skill_operation(operation, "anthropic") + + +def test_extract_model_param_ignores_non_string_body_model() -> None: + request = Request( + { + "type": "http", + "method": "POST", + "path": "/v1/skills/skill_1", + "headers": [], + "query_string": b"", + } + ) + + assert extract_model_param(request, {"model": {"unexpected": "type"}}) is None + + +def test_extract_model_param_falls_back_from_empty_body_model() -> None: + request = Request( + { + "type": "http", + "method": "POST", + "path": "/v1/skills/skill_1", + "headers": [(b"x-litellm-model", b"header-model")], + "query_string": b"model=query-model", + } + ) + + assert extract_model_param(request, {"model": ""}) == "query-model" + + +@pytest.mark.parametrize( + ("api_base", "expected"), + [ + (None, None), + ("https://resource.openai.azure.com", "https://resource.openai.azure.com"), + ("https://resource.openai.azure.com/openai", "https://resource.openai.azure.com"), + ("https://resource.openai.azure.com/openai/v1", "https://resource.openai.azure.com"), + ("https://resource.openai.azure.com/openai/responses?api-version=preview", "https://resource.openai.azure.com"), + ("https://resource.openai.azure.com/openai/v1/responses", "https://resource.openai.azure.com"), + ], +) +def test_azure_skills_api_base(api_base: str | None, expected: str | None) -> None: + assert _azure_skills_api_base(api_base) == expected + + +def test_native_routes_are_registered_without_anthropic_response_coercion() -> None: + expected = { + ("POST", "/v1/skills/{skill_id}"), + ("GET", "/v1/skills/{skill_id}/content"), + ("POST", "/v1/skills/{skill_id}/versions"), + ("GET", "/v1/skills/{skill_id}/versions"), + ("GET", "/v1/skills/{skill_id}/versions/{version}"), + ("DELETE", "/v1/skills/{skill_id}/versions/{version}"), + ("GET", "/v1/skills/{skill_id}/versions/{version}/content"), + } + routes = [route for route in skills_router.routes if isinstance(route, APIRoute)] + + assert expected <= {(method, route.path) for route in routes for method in route.methods} + assert all(route.response_model is None for route in routes)