mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
feat(proxy): add native OpenAI Skills routing
This commit is contained in:
parent
d447be15b9
commit
1e4464d82c
4 changed files with 554 additions and 14 deletions
|
|
@ -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,111 @@ 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]:
|
||||
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)
|
||||
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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -855,7 +855,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,
|
||||
|
|
@ -1026,7 +1026,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.
|
||||
|
||||
|
|
@ -1035,7 +1035,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)
|
||||
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,116 @@ 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")),
|
||||
}
|
||||
)
|
||||
|
||||
# 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 _native_skill_request(
|
||||
operation: str,
|
||||
request_data: dict[str, Any],
|
||||
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 or {}), "Foundry-Features": "Skills=V1Preview"} # mutable-ok: SDK headers
|
||||
if custom_llm_provider == "azure"
|
||||
else 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)
|
||||
|
|
@ -195,6 +303,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 +424,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
|
||||
|
|
@ -371,6 +490,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),
|
||||
|
|
@ -543,6 +673,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),
|
||||
|
|
@ -707,6 +848,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),
|
||||
|
|
|
|||
274
tests/test_litellm/skills/test_native_skills.py
Normal file
274
tests/test_litellm/skills/test_native_skills.py
Normal file
|
|
@ -0,0 +1,274 @@
|
|||
import json
|
||||
from typing import Any
|
||||
from unittest.mock import patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
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.skills_endpoints import (
|
||||
_native_skill_data,
|
||||
)
|
||||
from litellm.proxy.anthropic_endpoints.skills_endpoints import (
|
||||
router as skills_router,
|
||||
)
|
||||
from litellm.router import Router
|
||||
from litellm.skills.main import _azure_skills_api_base
|
||||
|
||||
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)
|
||||
|
||||
|
||||
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_uses_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 requests[0].headers["foundry-features"] == "Skills=V1Preview"
|
||||
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"),
|
||||
(None, "model=query", "header", "query"),
|
||||
(None, "", "header", "header"),
|
||||
],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_native_endpoint_model_priority(
|
||||
body_model: str | None,
|
||||
query: str,
|
||||
header_model: str,
|
||||
expected: str,
|
||||
) -> None:
|
||||
body = json.dumps({"model": body_model} if body_model 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())],
|
||||
"query_string": query.encode(),
|
||||
"path_params": {"skill_id": "skill_1"},
|
||||
},
|
||||
receive,
|
||||
)
|
||||
|
||||
data = await _native_skill_data(request, "update")
|
||||
|
||||
assert data["model"] == expected
|
||||
assert data["skill_id"] == "skill_1"
|
||||
assert data["custom_llm_provider"] == "openai"
|
||||
assert data["_skill_operation"] == "update"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("api_base", "expected"),
|
||||
[
|
||||
("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, expected: str) -> 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