mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
Merge 80f05a8e7b into 77765fd302
This commit is contained in:
commit
862d1e6669
5 changed files with 817 additions and 14 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: >-
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
)
|
||||
|
||||
|
||||
# ============================================================================
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
520
tests/test_litellm/skills/test_native_skills.py
Normal file
520
tests/test_litellm/skills/test_native_skills.py
Normal file
|
|
@ -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)
|
||||
Loading…
Add table
Reference in a new issue