diff --git a/cookbook/laya/README.md b/cookbook/laya/README.md new file mode 100644 index 00000000000..28447d32bd6 --- /dev/null +++ b/cookbook/laya/README.md @@ -0,0 +1,68 @@ +# Laya gateway and classifier + +[Laya](https://github.com/NandhaKishorM/laya) serves typed decisions using the System One protocol. LiteLLM exposes its native API at `POST /laya/v1/systemone` and offers Laya under the auto-router's Decision Model classifier + +## Start Laya + +Install the server in a separate environment from LiteLLM, then start it on a port reachable from the gateway + +```bash +python -m pip install 'laya[serve]==0.3.21' +LAYA_HOST=127.0.0.1 LAYA_PORT=8000 LAYA_MODELS=english laya-serve +``` + +Set `LAYA_API_BASE=http://127.0.0.1:8000` in the gateway environment. If the Laya server requires authentication, set the same `LAYA_API_KEY` on both processes. The server URL must use HTTP or HTTPS and contain no embedded credentials, query, or fragment + +## Call the native API + +Use a LiteLLM virtual key with access to `laya/english`. The request names the bare checkpoint; LiteLLM uses `laya/english` for authorization and usage logs + +```bash +curl http://localhost:4000/laya/v1/systemone \ + -H "Authorization: Bearer $LITELLM_API_KEY" \ + -H 'Content-Type: application/json' \ + -d '{ + "model": "english", + "state": "My invoice has two identical charges", + "questions": { + "department": { + "type": "choice", + "instructions": "Choose the department that should help", + "criteria": { + "billing": "Invoices, payments, and refunds", + "technical": "Bugs and connectivity problems" + } + } + } + }' +``` + +The response preserves Laya's `answers`, `usage`, and `routing` fields. Every request must explicitly choose `english`, `multilingual`, or `typed-decisions`; unknown names and automatic selection are rejected so model permissions match the checkpoint being called. This integration does not expose `/v1/evaluate` or chat completions + +## Use Laya as a classifier + +Add this router alongside your existing `small-solver` and `large-solver` deployments + +```yaml +model_list: + - model_name: laya-router + litellm_params: + model: auto_router/complexity_router + complexity_router_config: + classifier_type: jev + jev_classifier_config: + provider: laya + model: english + timeout_ms: 15000 + tiers: + SIMPLE: small-solver + MEDIUM: small-solver + COMPLEX: large-solver + REASONING: large-solver +``` + +In the dashboard, choose Decision Model, then Laya and a checkpoint. `classifier_type: jev` remains the stored type for both Jev and Laya + +Omitting `api_base` uses the gateway's `LAYA_API_BASE` and optional `LAYA_API_KEY`. Administrators may instead set `api_base` and optional `api_key` in `jev_classifier_config`; an explicit base without a key is keyless and does not inherit an environment key. Switching providers or changing the saved server URL discards the previous hidden key. Team members can select a centrally configured Laya checkpoint but cannot set connection credentials or URLs + +Laya's catalog token prices are zero because self-hosting has no provider API fee. Hosting costs are separate; operators may register their own token rates. Logs use the selected checkpoint, such as `laya/english`, and record the usage returned by Laya. Evaluate classification accuracy on your own prompts before enabling routing diff --git a/litellm/llms/laya/__init__.py b/litellm/llms/laya/__init__.py new file mode 100644 index 00000000000..8b137891791 --- /dev/null +++ b/litellm/llms/laya/__init__.py @@ -0,0 +1 @@ + diff --git a/litellm/llms/laya/common_utils.py b/litellm/llms/laya/common_utils.py new file mode 100644 index 00000000000..5b658732c50 --- /dev/null +++ b/litellm/llms/laya/common_utils.py @@ -0,0 +1,60 @@ +from collections.abc import Mapping +from dataclasses import dataclass, field +from typing import Final, Literal, TypeAlias + +from pydantic import AnyHttpUrl, BaseModel, TypeAdapter, ValidationError + +from litellm.secret_managers.main import get_secret_str + +LayaCheckpoint: TypeAlias = Literal["english", "multilingual", "typed-decisions"] + + +def validate_laya_model(value: object) -> LayaCheckpoint: + try: + return TypeAdapter(LayaCheckpoint).validate_python(value) + except ValidationError as exc: + raise ValueError("Laya model must be 'english', 'multilingual', or 'typed-decisions'") from exc + + +def validate_laya_request(body: Mapping[str, object]) -> LayaCheckpoint: + if "custom_body" in body: + raise ValueError("custom_body is not supported for Laya requests") + if body.get("stream"): + raise ValueError("Streaming is not supported for Laya requests") + return validate_laya_model(body.get("model")) + + +@dataclass(frozen=True, slots=True) +class LayaConnection: + api_base: str + api_key: str | None = field(repr=False) + + +def validate_laya_api_base(value: str) -> str: + try: + url: Final = TypeAdapter(AnyHttpUrl).validate_python(value) + except ValidationError as exc: + raise ValueError("Laya api_base must be an HTTP or HTTPS server URL") from exc + if url.username or url.password or url.query or url.fragment: + raise ValueError("Laya api_base must not contain credentials, a query, or a fragment") + return str(url).rstrip("/") + + +def laya_connection(api_base: str | None = None, api_key: str | None = None) -> LayaConnection: + base: Final = api_base if api_base is not None else get_secret_str("LAYA_API_BASE") + if not base: + raise ValueError("Laya requires api_base or LAYA_API_BASE pointing to a self-hosted server") + key: Final = api_key if api_base is not None else api_key or get_secret_str("LAYA_API_KEY") + return LayaConnection(api_base=validate_laya_api_base(base), api_key=key) + + +class _LayaRouting(BaseModel): + model: str | None = None + + +def laya_response_model(response: Mapping[str, object], requested_model: str | None) -> str: + try: + routing: Final = TypeAdapter(_LayaRouting).validate_python(response.get("routing") or {}) + except ValidationError: + return requested_model or "unknown" + return routing.model or requested_model or "unknown" diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 4c64b59190d..5b351b41168 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -72077,6 +72077,45 @@ "supports_audio_input": true, "supports_video_input": true }, + "laya/english": { + "input_cost_per_token": 0.0, + "litellm_provider": "laya", + "mode": "evaluation", + "output_cost_per_token": 0.0, + "source": "https://github.com/NandhaKishorM/laya", + "supported_endpoints": [ + "/v1/systemone" + ], + "metadata": { + "notes": "Self-hosted decision model; infrastructure costs are paid separately" + } + }, + "laya/multilingual": { + "input_cost_per_token": 0.0, + "litellm_provider": "laya", + "mode": "evaluation", + "output_cost_per_token": 0.0, + "source": "https://github.com/NandhaKishorM/laya", + "supported_endpoints": [ + "/v1/systemone" + ], + "metadata": { + "notes": "Self-hosted decision model; infrastructure costs are paid separately" + } + }, + "laya/typed-decisions": { + "input_cost_per_token": 0.0, + "litellm_provider": "laya", + "mode": "evaluation", + "output_cost_per_token": 0.0, + "source": "https://github.com/NandhaKishorM/laya", + "supported_endpoints": [ + "/v1/systemone" + ], + "metadata": { + "notes": "Self-hosted decision model; infrastructure costs are paid separately" + } + }, "typesafe/jev-1.13.0": { "input_cost_per_token": 4.2e-08, "litellm_provider": "typesafe", diff --git a/litellm/proxy/_lazy_features.py b/litellm/proxy/_lazy_features.py index 5be87a8bf4d..6b9b2587cfc 100644 --- a/litellm/proxy/_lazy_features.py +++ b/litellm/proxy/_lazy_features.py @@ -214,6 +214,7 @@ LAZY_FEATURES: Final[tuple[LazyFeature, ...]] = ( "/tinyfish/", "/transcribe", "/typesafe/", + "/laya/", "/openrouter/", "/vertex-ai/", "/vertex_ai/", diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index db57aa4f046..a9ee23e6a50 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -26343,6 +26343,30 @@ ] } }, + "/laya/v1/systemone": { + "post": { + "operationId": "laya_proxy_route_laya_v1_systemone_post", + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Laya Proxy Route", + "tags": [ + "llm_passthrough" + ] + } + }, "/milvus/{endpoint}": { "delete": { "description": "Enable using Milvus `/vectors` endpoint as a pass-through endpoint.", diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 3d8d701d15b..aadbd4dc63c 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -499,6 +499,7 @@ class LiteLLMRoutes(enum.Enum): "/vllm", "/mistral", "/typesafe", + "/laya", "/openrouter", "/milvus", "/gigachat", diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index c0123ae45a3..3653e519f41 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -1883,6 +1883,14 @@ def _extract_model_candidates_from_request( llm_router: Router | None = None, team_id: str | None = None, ) -> list[str]: + if route.rstrip("/") == "/laya/v1/systemone": + from litellm.llms.laya.common_utils import validate_laya_model + + try: + laya_model: Final = validate_laya_model(request_data.get("model")) + except ValueError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + return [f"laya/{laya_model}"] if route == "/cost/predict-cache": prediction_models: Final = _cache_prediction_model_candidates(request_data, llm_router, team_id) # pyright: ignore[reportUnknownArgumentType] # the typed reader validates each deployment ID from this legacy payload return _dedupe_model_candidates(prediction_models) diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index ae294871afc..31ff9559741 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -39,6 +39,7 @@ from litellm.litellm_core_utils.ptu_pricing import ( SEARCH_CONTEXT_SIZES, ptu_config_error, ) +from litellm.litellm_core_utils.sensitive_data_masker import redact_credentials_in_payload from litellm.proxy._types import ( BlockModelRequest, CommonProxyErrors, @@ -114,6 +115,7 @@ from litellm.router_strategy.complexity_router import ( normalize_classification_examples, normalize_classification_prompt, ) +from litellm.router_strategy.complexity_router.config import resolve_complexity_router_config_write from litellm.router_utils.auto_router_model_naming import ( GATED_AUTO_ROUTER_CAPABILITIES, STRATEGY_ROUTER_PARAM_FIELDS, @@ -186,6 +188,25 @@ class _ProxyModelRow(Protocol): def model_dump_json(self, *, exclude_none: bool = False) -> str: ... +def _model_write_response( + row: _ProxyModelRow, member_write: MemberAutoRouterWrite | None +) -> _ProxyModelRow | Mapping[str, object]: + if member_write is None: + return row + payload: Final = TypeAdapter(dict[str, object]).validate_json(row.model_dump_json()) + stored_params: Final = payload.get("litellm_params") + params: Final = ( + TypeAdapter(dict[str, object]).validate_json(stored_params) + if isinstance(stored_params, str) + else TypeAdapter(dict[str, object]).validate_python(stored_params) + ) + redacted: Final = redact_credentials_in_payload(params) + return { + **payload, + "litellm_params": json.dumps(redacted) if isinstance(stored_params, str) else redacted, + } + + class _ProxyModelTable(Protocol): def find_unique(self, *, where: Mapping[str, object]) -> Awaitable[BaseModel | None]: ... @@ -406,34 +427,10 @@ WHERE model_id <> $1 def _effective_complexity_router_config( incoming_params: GenericLiteLLMParams | None, existing_params: GenericLiteLLMParams | None -) -> object: +) -> Mapping[str, object] | None: incoming: Final = None if incoming_params is None else incoming_params.complexity_router_config existing: Final = None if existing_params is None else existing_params.complexity_router_config - if incoming is None: - return existing - if existing is None or incoming.get("classifier_type") != "jev" or existing.get("classifier_type") != "jev": - return incoming - incoming_jev: Final[object] = incoming.get("jev_classifier_config") - existing_jev: Final[object] = existing.get("jev_classifier_config") - if not isinstance(incoming_jev, Mapping) or not isinstance(existing_jev, Mapping): - return incoming - supplied: Final = TypeAdapter(dict[str, object]).validate_python(incoming_jev) - stored: Final = TypeAdapter(dict[str, object]).validate_python(existing_jev) - same_base: Final = "api_base" not in supplied or supplied["api_base"] == stored.get("api_base") - transport: Final = MappingProxyType( - { - key: value - for key, value in stored.items() - if key in ("api_key", "api_base") and (key != "api_key" or same_base) - } - ) - return { # mutable-ok: persisted JSON requires concrete nested dicts - **incoming, - "jev_classifier_config": { # mutable-ok: json.dumps cannot serialize MappingProxyType - **transport, - **supplied, - }, - } + return resolve_complexity_router_config_write(incoming, existing).effective def _effective_model( @@ -1303,7 +1300,7 @@ async def patch_model( live_after=reload_outcome.live_after, ) - return updated_model + return _model_write_response(updated_model, member_write) except Exception as e: verbose_proxy_logger.exception("Error in patch_model: %s", e) @@ -2538,7 +2535,7 @@ async def add_new_model( live_after=reload_outcome.live_after, ) - return model_response + return _model_write_response(model_response, member_write) except Exception as e: verbose_proxy_logger.exception("litellm.proxy.proxy_server.add_new_model(): Exception occured - %s", e) @@ -2764,7 +2761,7 @@ async def update_model( live_after=reload_outcome.live_after, ) - return model_response + return _model_write_response(model_response, member_write) except Exception as e: verbose_proxy_logger.exception("litellm.proxy.proxy_server.update_model(): Exception occured - %s", e) if isinstance(e, HTTPException): diff --git a/litellm/proxy/management_helpers/auto_router_permissions.py b/litellm/proxy/management_helpers/auto_router_permissions.py index 449a1032b35..64b752d428f 100644 --- a/litellm/proxy/management_helpers/auto_router_permissions.py +++ b/litellm/proxy/management_helpers/auto_router_permissions.py @@ -33,6 +33,10 @@ from litellm.repositories.prisma_protocols import DatabaseClient from litellm.repositories.project_repository import ProjectRepository from litellm.repositories.table_repositories import TeamMembershipRepository from litellm.router import Router +from litellm.router_strategy.complexity_router.config import ( + ComplexityRouterConfigWrite, + resolve_complexity_router_config_write, +) from litellm.router_utils.auto_router_model_naming import classify_strategy_router_model, strategy_router_dependencies from litellm.types.management_endpoints.auto_router_endpoints import RequestComplexityRouterConfig from litellm.types.router import Deployment, updateDeployment @@ -71,6 +75,7 @@ class _MemberJevClassifierConfig(BaseModel): model_config = ConfigDict(extra="forbid") + provider: Literal["typesafe", "laya"] = "typesafe" model: str api_key: None = None api_base: None = None @@ -123,14 +128,21 @@ def authorize_member_auto_router_team( def validate_member_auto_router_config(config: Mapping[str, object]) -> RequestComplexityRouterConfig: + return _validate_member_auto_router_config_write(resolve_complexity_router_config_write(config, None)) + + +def _validate_member_auto_router_config_write(write: ComplexityRouterConfigWrite) -> RequestComplexityRouterConfig: + if write.effective is None: + raise HTTPException(status_code=400, detail="A complexity_router_config is required.") try: - validated: Final = _MemberComplexityRouterConfig.model_validate(config) - for entries in validated.tier_model_configs.values(): - for entry in entries: - _MemberRouterGenerationParams.model_validate(entry.litellm_params) - if validated.jev_classifier_config is not None: - _MemberJevClassifierConfig.model_validate(validated.jev_classifier_config.model_dump()) - return validated + if write.submitted is not None: + validated: Final = _MemberComplexityRouterConfig.model_validate(write.submitted) + for entries in validated.tier_model_configs.values(): + for entry in entries: + _MemberRouterGenerationParams.model_validate(entry.litellm_params) + if validated.jev_classifier_config is not None: + _MemberJevClassifierConfig.model_validate(validated.jev_classifier_config.model_dump()) + return RequestComplexityRouterConfig.model_validate(write.effective) except ValidationError as exc: location: Final = ".".join(str(part) for part in exc.errors()[0]["loc"]) raise HTTPException(status_code=400, detail=f"Invalid member auto-router configuration at {location}.") from exc @@ -334,16 +346,15 @@ async def authorize_member_auto_router_write( if existing is not None and incoming.model_name not in (None, public_name, existing.model_name): raise HTTPException(status_code=403, detail="Team members cannot rename an auto router.") supplied_config: Final = _RouterConfigSource.model_validate(params.model_dump()).complexity_router_config - raw_config: Final = ( - supplied_config - if supplied_config is not None - else _RouterConfigSource.model_validate(existing.litellm_params.model_dump()).complexity_router_config + stored_config: Final = ( + _RouterConfigSource.model_validate(existing.litellm_params.model_dump()).complexity_router_config if existing is not None else None ) - if raw_config is None: - raise HTTPException(status_code=400, detail="A complexity_router_config is required.") - config: Final = validate_member_auto_router_config(raw_config) + resolved_config: Final = resolve_complexity_router_config_write(supplied_config, stored_config) + if resolved_config.supplied_connection_fields: + raise HTTPException(status_code=403, detail="Team members cannot change classifier connections.") + config: Final = _validate_member_auto_router_config_write(resolved_config) stored_default: Final = existing.litellm_params.complexity_router_default_model if existing is not None else None default_model: Final = ( params.complexity_router_default_model diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index eaa03b67b40..5d71aeb501a 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -24,6 +24,7 @@ from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, cast import httpx from fastapi import APIRouter, Depends, HTTPException, Request, Response, WebSocket from fastapi.responses import StreamingResponse +from pydantic import TypeAdapter from starlette.websockets import WebSocketState from typing_extensions import ReadOnly, TypedDict @@ -54,6 +55,7 @@ from litellm.llms.deepgram.common_utils import ( deepgram_listen_websocket_target, ) from litellm.llms.fal_ai.cost_calculator import fal_ai_passthrough_cost, fal_ai_queue_base +from litellm.llms.laya.common_utils import laya_connection, validate_laya_request from litellm.llms.nvidia_nim.passthrough.transformation import nvidia_nim_model_group_in_path from litellm.llms.vertex_ai.vertex_llm_base import VertexBase from litellm.passthrough.main import AsyncPassthroughStreamingResponse @@ -633,6 +635,43 @@ async def typesafe_proxy_route( return await endpoint_func(request, fastapi_response, user_api_key_dict) +@router.post( + "/laya/v1/systemone", + tags=["Laya Pass-through", "pass-through"], # mutable-ok: FastAPI route metadata requires a list +) +async def laya_proxy_route( + request: Request, + fastapi_response: Response, + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], +) -> Response: + body: Final = TypeAdapter(dict[str, object]).validate_python(await _read_request_body(request)) + try: + _ = validate_laya_request(body) + except ValueError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + try: + connection: Final = laya_connection() + except ValueError as exc: + raise HTTPException( + status_code=503, detail="Laya server is not configured correctly; check LAYA_API_BASE" + ) from exc + base_url: Final = httpx.URL(connection.api_base) + updated_url: Final = base_url.copy_with( + path=HttpPassThroughEndpointHelpers.join_base_and_endpoint_path(base_url, "/v1/systemone"), + ) + endpoint_func: Final = create_pass_through_route( + endpoint="v1/systemone", + target=str(updated_url), + custom_headers={ # mutable-ok: pass-through request headers require a mutable mapping + **({"Authorization": f"Bearer {connection.api_key}"} if connection.api_key else {}), + "Content-Type": "application/json", + }, + custom_llm_provider="laya", + is_streaming_request=False, + ) + return cast(Response, await endpoint_func(request, fastapi_response, user_api_key_dict)) + + @router.api_route( "/openrouter/{endpoint:path}", methods=["GET", "POST", "PUT", "DELETE", "PATCH"], # mutable-ok: FastAPI route metadata requires a list diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/typesafe_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/typesafe_passthrough_logging_handler.py index 887d17a7a20..e8e8e66df28 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/typesafe_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/typesafe_passthrough_logging_handler.py @@ -10,6 +10,7 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.litellm_core_utils.litellm_logging import ( get_standard_logging_object_payload, # pyright: ignore[reportUnknownVariableType] # legacy helper has an untyped signature ) +from litellm.llms.laya.common_utils import laya_response_model from litellm.proxy._types import PassThroughEndpointLoggingTypedDict from litellm.types.utils import ModelResponse, StandardPassThroughResponseObject, Usage @@ -69,9 +70,11 @@ class TypeSafePassthroughLoggingHandler: **kwargs: object, ) -> PassThroughEndpointLoggingTypedDict: response: Final = _parse_typesafe_response(response_body) - response_model: Final = response.model request_model_value: Final = request_body.get("model") request_model: Final = request_model_value if isinstance(request_model_value, str) else None + response_model: Final = ( + laya_response_model(response_body, request_model) if custom_llm_provider == "laya" else response.model + ) logged_model: Final = response_model or request_model or "unknown" model_name: Final = f"{custom_llm_provider}/{logged_model}" usage: Final = response.usage or _TypeSafeUsage() diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index e0a4184291e..fd6c23fa712 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -74,7 +74,11 @@ from litellm.proxy._types import ( ProxyException, UserAPIKeyAuth, ) -from litellm.proxy.auth.auth_utils import request_dispatched_to_pass_through_endpoint +from litellm.proxy.auth.auth_utils import ( + get_model_from_request, + get_request_route, + request_dispatched_to_pass_through_endpoint, +) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_request_processing import ( ProxyBaseLLMRequestProcessing, @@ -100,6 +104,8 @@ from litellm.proxy.common_utils.sse_keepalive import ( from litellm.proxy.litellm_pre_call_utils import ( LiteLLMProxyRequestSetup, _get_dynamic_logging_metadata, # pyright: ignore[reportPrivateUsage] # shared proxy helper, same import style as _read_request_body above + _key_or_team_allows_client_pricing_override, # pyright: ignore[reportPrivateUsage] # reuse the proxy's pricing trust policy + _strip_client_pricing_overrides, # pyright: ignore[reportPrivateUsage] # sanitize before trusted hooks add guardrail costs ) from litellm.proxy.route_llm_request import ProxyModelNotFoundError from litellm.proxy.utils import normalize_route_for_root_path @@ -583,7 +589,18 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): """ Filter out litellm params from the request body """ + from litellm.proxy.proxy_server import llm_router + _parsed_body = _parsed_body or {} + managed_model: Final = get_model_from_request( + request_data=_parsed_body, + route=get_request_route(request), + request_headers=request.headers, + request_query_params=request.query_params, + llm_router=llm_router, + request=request, + team_id=user_api_key_dict.team_id, + ) litellm_keys_in_body: Final = MappingProxyType( {k: _parsed_body.pop(k) for k in types_utils.all_litellm_params if k in _parsed_body} @@ -629,10 +646,19 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): # would attribute it to a budget the operator scoped to a LiteLLM model that # merely shares the name. if not request_dispatched_to_pass_through_endpoint(request): + _metadata["model_group"] = managed_model if isinstance(managed_model, str) else None _metadata["user_api_key_model_max_budget"] = user_api_key_dict.model_max_budget _metadata["user_api_key_team_model_max_budget"] = user_api_key_dict.team_model_max_budget _metadata["user_api_key_user_model_max_budget"] = user_api_key_dict.user_model_max_budget _metadata["user_api_key_end_user_model_max_budget"] = user_api_key_dict.end_user_model_max_budget + else: + for field in ( + "user_api_key_model_max_budget", + "user_api_key_team_model_max_budget", + "user_api_key_user_model_max_budget", + "user_api_key_end_user_model_max_budget", + ): + _metadata.pop(field, None) _metadata.update( LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict=user_api_key_dict) ) @@ -1127,6 +1153,14 @@ async def pass_through_request( _parsed_body, ) + if not _key_or_team_allows_client_pricing_override(user_api_key_dict): + _strip_client_pricing_overrides(_parsed_body) + if custom_llm_provider == "laya": + from litellm.llms.laya.common_utils import validate_laya_request + + checkpoint: Final = validate_laya_request(_parsed_body) + _parsed_body["model"] = f"laya/{checkpoint}" + ### COLLECT GUARDRAILS FOR PASSTHROUGH ENDPOINT ### # Passthrough endpoints are opt-in only for guardrails # When enabled, collect guardrails from org/team/key levels + passthrough-specific @@ -1181,6 +1215,13 @@ async def pass_through_request( data=_parsed_body, call_type="pass_through_endpoint", ) + if custom_llm_provider == "laya": + hook_model: Final = _parsed_body.get("model") + laya_body: Final = { + **_parsed_body, + "model": hook_model.removeprefix("laya/") if isinstance(hook_model, str) else hook_model, + } + _parsed_body = {**laya_body, "model": validate_laya_request(laya_body)} resolved_timeout: Final = resolve_pass_through_request_timeout(timeout) async_client_obj: Final = get_async_httpx_client( llm_provider=httpxSpecialProvider.PassThroughEndpoint, @@ -2384,7 +2425,9 @@ async def websocket_passthrough_request( # with the existing _init_kwargs_for_pass_through_endpoint function class DummyRequest: def __init__(self, url: str, method: str = "WEBSOCKET", headers: dict | None = None): - self.url = url + self.url = httpx.URL(url) + self.scope = websocket.scope + self.query_params = websocket.query_params self.method = method self.headers = headers or {} diff --git a/litellm/proxy/pass_through_endpoints/success_handler.py b/litellm/proxy/pass_through_endpoints/success_handler.py index 6bba879b6c1..3c4733d0bf0 100644 --- a/litellm/proxy/pass_through_endpoints/success_handler.py +++ b/litellm/proxy/pass_through_endpoints/success_handler.py @@ -334,8 +334,10 @@ class PassThroughEndpointLogging: ) standard_logging_response_object = transcribe_handler_result["result"] # rebind-ok: elif-chain kwargs = transcribe_handler_result["kwargs"] # rebind-ok: elif-chain contract - elif self.is_typesafe_route(custom_llm_provider) or self.is_openrouter_decisions_route( - url_route, custom_llm_provider + elif ( + self.is_typesafe_route(custom_llm_provider) + or custom_llm_provider == "laya" + or self.is_openrouter_decisions_route(url_route, custom_llm_provider) ): from .llm_provider_handlers.typesafe_passthrough_logging_handler import ( TypeSafePassthroughLoggingHandler, diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 76ee977bf28..dd9434f61a8 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -1309,6 +1309,16 @@ class ComplexityRouter(CustomLogger): @staticmethod def _build_jev_client(config: JevClassifierConfig) -> JevClassifierClient: + if config.provider == "laya": + from litellm.llms.laya.common_utils import laya_connection + + connection: Final = laya_connection(config.api_base, config.api_key) + return HttpJevClassifierClient( + api_key=connection.api_key, + api_base=connection.api_base, + http_client=get_async_httpx_client(httpxSpecialProvider.PassThroughEndpoint), + provider="laya", + ) api_key: Final = config.api_key or get_secret_str("TYPESAFE_API_KEY") if not api_key: raise ValueError("jev_classifier_config.api_key or TYPESAFE_API_KEY is required for classifier_type 'jev'") @@ -2217,7 +2227,8 @@ class ComplexityRouter(CustomLogger): probabilities=answer.probabilities, confidence=answer.confidence, model=model, - cost=jev_classifier_cost(response, config.model), + cost=jev_classifier_cost(response, config.model, config.provider), + provider=config.provider, ) if breaker is not None and permit is not None: breaker.record_success(permit) @@ -2225,8 +2236,8 @@ class ComplexityRouter(CustomLogger): tier=tier, score=None, signals=( - f"jev-classifier:{tier_name}", - f"jev-confidence={answer.confidence:.6f}", + f"{'laya' if config.provider == 'laya' else 'jev'}-classifier:{tier_name}", + f"{'laya' if config.provider == 'laya' else 'jev'}-confidence={answer.confidence:.6f}", *( f"tier-probability:{label}={probability:.6f}" for label, probability in answer.probabilities.items() @@ -4765,7 +4776,7 @@ class ComplexityRouter(CustomLogger): tier_litellm_params: Final = self._litellm_params_for_model(tier, routed_model) classifier_model: Final = ( - f"typesafe/{outcome.jev_verdict.model}" + f"{outcome.jev_verdict.provider}/{outcome.jev_verdict.model}" if outcome.cause == "jev_classifier" and outcome.jev_verdict is not None else self.config.classifier_llm_config.model if outcome.cause in ("llm_classifier", "capability_classifier", "llm_v2_classifier", "llm_v2_fallback") diff --git a/litellm/router_strategy/complexity_router/config.py b/litellm/router_strategy/complexity_router/config.py index e0427f89fe3..fd24ab8a81b 100644 --- a/litellm/router_strategy/complexity_router/config.py +++ b/litellm/router_strategy/complexity_router/config.py @@ -9,6 +9,7 @@ import math import re import warnings from collections.abc import Iterable, Mapping +from dataclasses import dataclass from enum import Enum from types import MappingProxyType from typing import Annotated, Final, Literal, NamedTuple @@ -19,6 +20,7 @@ from pydantic import ( Field, SkipValidation, StrictFloat, + TypeAdapter, field_serializer, field_validator, model_validator, @@ -681,11 +683,12 @@ class CapabilityClassifierConfig(BaseModel): class JevClassifierConfig(BaseModel): model_config = ConfigDict(extra="forbid", frozen=True) + provider: Literal["typesafe", "laya"] = "typesafe" model: str = "jev-latest" - api_key: str | None = Field(default=None, description="TypeSafe API key, falling back to TYPESAFE_API_KEY") + api_key: str | None = Field(default=None, description="Provider API key; optional for self-hosted Laya") api_base: str | None = Field( default=None, - description="TypeSafe API base, falling back to TYPESAFE_API_BASE and then https://api.typesafe.ai", + description="Provider API base; defaults to TYPESAFE_API_BASE or LAYA_API_BASE for the selected provider", ) timeout_ms: int = Field(default=3000, ge=1) instructions: str | None = Field( @@ -711,6 +714,13 @@ class JevClassifierConfig(BaseModel): @model_validator(mode="after") def _keep_the_environment_key_on_the_environment_base(self) -> "JevClassifierConfig": + if self.provider == "laya": + from litellm.llms.laya.common_utils import validate_laya_api_base, validate_laya_model + + _ = validate_laya_model(self.model) + if self.api_base is not None: + _ = validate_laya_api_base(self.api_base) + return self if self.api_base is not None and self.api_key is None: raise ValueError( "jev_classifier_config.api_base requires jev_classifier_config.api_key: TYPESAFE_API_KEY is only sent " @@ -719,6 +729,60 @@ class JevClassifierConfig(BaseModel): return self +@dataclass(frozen=True, slots=True) +class ComplexityRouterConfigWrite: + submitted: Mapping[str, object] | None + effective: Mapping[str, object] | None + + @property + def supplied_connection_fields(self) -> frozenset[str]: + classifier: Final = self.submitted.get("jev_classifier_config") if self.submitted is not None else None + return frozenset( + field for field in ("api_base", "api_key") if isinstance(classifier, Mapping) and field in classifier + ) + + +def resolve_complexity_router_config_write( + incoming: Mapping[str, object] | None, stored: Mapping[str, object] | None +) -> ComplexityRouterConfigWrite: + if incoming is None: + return ComplexityRouterConfigWrite(submitted=None, effective=stored) + if stored is None or incoming.get("classifier_type") != "jev" or stored.get("classifier_type") != "jev": + return ComplexityRouterConfigWrite(submitted=incoming, effective=incoming) + incoming_classifier: Final = incoming.get("jev_classifier_config") + stored_classifier: Final = stored.get("jev_classifier_config") + if not isinstance(incoming_classifier, Mapping) or not isinstance(stored_classifier, Mapping): + return ComplexityRouterConfigWrite(submitted=incoming, effective=incoming) + existing: Final = TypeAdapter(dict[str, object]).validate_python(stored_classifier) + supplied: Final = TypeAdapter(dict[str, object]).validate_python(incoming_classifier) + classifier: Final = ( + MappingProxyType({**supplied, "provider": existing["provider"]}) + if "provider" not in supplied and "provider" in existing + else supplied + ) + same_provider: Final = classifier.get("provider", "typesafe") == existing.get("provider", "typesafe") + same_base: Final = "api_base" not in classifier or ( + classifier["api_base"] is not None and classifier["api_base"] == existing.get("api_base") + ) + transport: Final = MappingProxyType( + { + key: value + for key, value in existing.items() + if same_provider and key in ("api_key", "api_base") and (key != "api_key" or same_base) + } + ) + return ComplexityRouterConfigWrite( + submitted=MappingProxyType({**incoming, "jev_classifier_config": classifier}), + effective={ # mutable-ok: persisted JSON requires concrete nested dicts + **incoming, + "jev_classifier_config": { # mutable-ok: json.dumps cannot serialize MappingProxyType + **transport, + **classifier, + }, + }, + ) + + MAX_CUSTOM_PATTERN_REPEAT: Final[int] = 64 MAX_CUSTOM_PATTERN_WORK: Final[int] = 2048 MAX_CUSTOM_DIMENSIONS_WORK: Final[int] = 8192 diff --git a/litellm/router_strategy/complexity_router/jev_classifier.py b/litellm/router_strategy/complexity_router/jev_classifier.py index 02e57975626..bd2b1701b38 100644 --- a/litellm/router_strategy/complexity_router/jev_classifier.py +++ b/litellm/router_strategy/complexity_router/jev_classifier.py @@ -18,6 +18,7 @@ from litellm.litellm_core_utils.internal_call_metadata import ( from litellm.litellm_core_utils.litellm_logging import Logging from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler +from litellm.llms.laya.common_utils import laya_response_model from litellm.proxy.pass_through_endpoints.llm_provider_handlers.typesafe_passthrough_logging_handler import ( TypeSafePassthroughLoggingHandler, ) @@ -78,10 +79,17 @@ class JevClassifierClient(Protocol): class HttpJevClassifierClient: - def __init__(self, api_key: str, api_base: str, http_client: AsyncHTTPHandler) -> None: + def __init__( + self, + api_key: str | None, + api_base: str, + http_client: AsyncHTTPHandler, + provider: Literal["typesafe", "laya"] = "typesafe", + ) -> None: self._api_key = api_key self._api_base = api_base.rstrip("/") self._http_client = http_client + self._provider = provider async def evaluate( self, @@ -95,21 +103,25 @@ class HttpJevClassifierClient: json=request.model_dump(mode="json"), headers=MappingProxyType( { - "Authorization": f"Bearer {self._api_key}", + **({"Authorization": f"Bearer {self._api_key}"} if self._api_key else {}), "Content-Type": "application/json", } ), # pyright: ignore[reportArgumentType] # HTTP headers are not mutated by AsyncHTTPHandler timeout=timeout_s, ) response.raise_for_status() + body: Final = TypeAdapter(dict[str, object]).validate_json(response.content) + normalized_body: Final = ( + {**body, "model": laya_response_model(body, request.model)} if self._provider == "laya" else body + ) try: self._log_response(request, response, request_kwargs, start_time) except Exception as exc: # noqa: BLE001 # logging integrations must not discard a provider verdict verbose_router_logger.warning("JEV response logging failed (%s)", type(exc).__name__) - return TypeAdapter(JevSystemOneResponse).validate_python(response.json()) + return TypeAdapter(JevSystemOneResponse).validate_python(normalized_body) - @staticmethod def _log_response( + self, request: JevSystemOneRequest, response: httpx.Response, request_kwargs: Mapping[str, object] | None, @@ -139,7 +151,7 @@ class HttpJevClassifierClient: "turn_off_message_logging": effective_turn_off_message_logging(request_kwargs), } logging_obj: Final = Logging( - model=f"typesafe/{request.model}", + model=f"{self._provider}/{request.model}", messages=[{"role": "user", "content": request.state}], # mutable-ok: callbacks require JSON message lists stream=False, call_type="pass_through_endpoint", @@ -150,7 +162,7 @@ class HttpJevClassifierClient: kwargs=params, ) logging_obj.update_environment_variables( - model=f"typesafe/{request.model}", + model=f"{self._provider}/{request.model}", user=parent_user if isinstance(parent_user := parent.get("user"), str) else None, optional_params={}, # mutable-ok: Logging's optional_params contract requires a dict litellm_params=params, @@ -165,7 +177,7 @@ class HttpJevClassifierClient: end_time=end_time, cache_hit=False, request_body=MappingProxyType({"model": request.model}), - custom_llm_provider="typesafe", + custom_llm_provider=self._provider, litellm_params=params, ) success_handlers: Final = logging_obj.dispatch_success_handlers( @@ -189,6 +201,7 @@ class JevVerdict(NamedTuple): confidence: float model: str cost: float | None + provider: Literal["typesafe", "laya"] = "typesafe" class _RegistryPricing(BaseModel): @@ -211,12 +224,14 @@ def build_jev_request( return JevSystemOneRequest(state=state, model=model, questions=MappingProxyType({"tier": question})) -def jev_classifier_cost(response: JevSystemOneResponse, configured_model: str) -> float | None: +def jev_classifier_cost( + response: JevSystemOneResponse, configured_model: str, provider: Literal["typesafe", "laya"] = "typesafe" +) -> float | None: usage: Final = response.usage if usage is None: return None model: Final = response.model or configured_model - model_key: Final = f"typesafe/{model}" + model_key: Final = f"{provider}/{model}" if model_key not in litellm.model_cost: # pyright: ignore[reportUnknownMemberType] # registry is dynamically typed return None try: diff --git a/litellm/router_utils/auto_router_model_naming.py b/litellm/router_utils/auto_router_model_naming.py index 6b589c3bfc0..416aa10dd4e 100644 --- a/litellm/router_utils/auto_router_model_naming.py +++ b/litellm/router_utils/auto_router_model_naming.py @@ -153,6 +153,7 @@ def strategy_router_dependencies( ) complexity: Final = _mapping(litellm_params.get("complexity_router_config")) classifier: Final = _mapping(complexity.get("classifier_llm_config")) + decision_classifier: Final = _mapping(complexity.get("jev_classifier_config")) return tuple( dict.fromkeys( tuple(dep for tier in _mapping(complexity.get("tiers")).values() for dep in _pool(tier, "tier")) @@ -165,7 +166,7 @@ def strategy_router_dependencies( ) + ( _named( - f"typesafe/{_mapping(complexity.get('jev_classifier_config')).get('model', 'jev-latest')}", + f"{decision_classifier.get('provider', 'typesafe')}/{decision_classifier.get('model', 'jev-latest')}", "evaluation", ) if complexity.get("classifier_type") == "jev" diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 4c64b59190d..5b351b41168 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -72077,6 +72077,45 @@ "supports_audio_input": true, "supports_video_input": true }, + "laya/english": { + "input_cost_per_token": 0.0, + "litellm_provider": "laya", + "mode": "evaluation", + "output_cost_per_token": 0.0, + "source": "https://github.com/NandhaKishorM/laya", + "supported_endpoints": [ + "/v1/systemone" + ], + "metadata": { + "notes": "Self-hosted decision model; infrastructure costs are paid separately" + } + }, + "laya/multilingual": { + "input_cost_per_token": 0.0, + "litellm_provider": "laya", + "mode": "evaluation", + "output_cost_per_token": 0.0, + "source": "https://github.com/NandhaKishorM/laya", + "supported_endpoints": [ + "/v1/systemone" + ], + "metadata": { + "notes": "Self-hosted decision model; infrastructure costs are paid separately" + } + }, + "laya/typed-decisions": { + "input_cost_per_token": 0.0, + "litellm_provider": "laya", + "mode": "evaluation", + "output_cost_per_token": 0.0, + "source": "https://github.com/NandhaKishorM/laya", + "supported_endpoints": [ + "/v1/systemone" + ], + "metadata": { + "notes": "Self-hosted decision model; infrastructure costs are paid separately" + } + }, "typesafe/jev-1.13.0": { "input_cost_per_token": 4.2e-08, "litellm_provider": "typesafe", diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index 790a050a878..f53da81a271 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -6,6 +6,7 @@ "url": "Link to provider documentation", "endpoints": { "chat_completions": "Supports /chat/completions endpoint", + "systemone": "Supports native System One typed decisions", "messages": "Supports /messages endpoint (Anthropic format)", "responses": "Supports /responses endpoint (OpenAI/Anthropic unified)", "embeddings": "Supports /embeddings endpoint", @@ -1458,6 +1459,13 @@ "rerank": false } }, + "laya": { + "display_name": "Laya (`laya`)", + "url": "https://github.com/BerriAI/litellm/tree/main/cookbook/laya", + "endpoints": { + "systemone": true + } + }, "lambda_ai": { "display_name": "Lambda AI (`lambda_ai`)", "url": "https://docs.litellm.ai/docs/providers/lambda_ai", @@ -3319,6 +3327,13 @@ "provider_json_field": "skills", "url": "https://docs.litellm.ai/docs/skills" }, + "systemone": { + "docs_label": "systemone", + "display_name": "System One Decision API", + "leftnav_label": "/laya/v1/systemone", + "provider_json_field": "systemone", + "url": "https://github.com/BerriAI/litellm/tree/main/cookbook/laya" + }, "text_completion": { "docs_label": "text_completion", "display_name": "OpenAI Completions API", diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index 83ac56c4c85..6cf4456a0ff 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -8,7 +8,7 @@ from typing import Optional from unittest.mock import MagicMock, patch import pytest -from fastapi import Request +from fastapi import HTTPException, Request from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.auth_utils import ( @@ -463,6 +463,24 @@ def test_get_model_from_request_no_request_extracts_model(): ) +@pytest.mark.parametrize("model", ["english", "multilingual", "typed-decisions"]) +@pytest.mark.parametrize("route", ["/laya/v1/systemone", "/laya/v1/systemone/"]) +def test_laya_native_model_uses_the_classifier_permission_identity(model: str, route: str) -> None: + assert get_model_from_request(request_data={"model": model}, route=route) == f"laya/{model}" + + +@pytest.mark.parametrize("model", [None, "", "auto", "laya/english", "unknown", ["english"], 7]) +def test_laya_native_model_cannot_implicitly_select_an_unauthorized_checkpoint(model: object) -> None: + with pytest.raises(HTTPException) as denied: + get_model_from_request(request_data={"model": model}, route="/laya/v1/systemone") + assert denied.value.status_code == 400 + + +def test_laya_model_normalization_does_not_change_other_provider_routes() -> None: + assert get_model_from_request(request_data={"model": "jev-latest"}, route="/typesafe/v1/systemone") == "jev-latest" + assert get_model_from_request(request_data={}, route="/laya/health") is None + + def _cache_prediction_router(): from litellm.router import Router diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index 5f7807650e1..912c056055a 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -8,6 +8,8 @@ from typing import Dict, Final, Optional from unittest.mock import AsyncMock, MagicMock, patch import pytest +from fastapi import HTTPException +from fastapi.encoders import jsonable_encoder from fastapi.testclient import TestClient from litellm._uuid import uuid @@ -7603,6 +7605,148 @@ class TestTeamMemberAutoRouterWrites: assert row.litellm_params["complexity_router_config"] == stored_config assert request.litellm_params.complexity_router_config == config + @pytest.mark.asyncio + @pytest.mark.parametrize("endpoint", ["patch", "legacy"]) + @pytest.mark.parametrize( + "stored_provider,stored_base,supplied,expected_transport", + [ + ("laya", "https://decision.test", {"provider": "laya", "model": "english"}, {"api_base": "https://decision.test", "api_key": "stored-secret"}), + ("laya", "https://decision.test", {"provider": "laya", "model": "english", "api_base": "https://decision.test"}, {"api_base": "https://decision.test", "api_key": "stored-secret"}), + ( + "laya", + "https://decision.test", + {"provider": "laya", "model": "english", "api_key": None}, + {"api_base": "https://decision.test"}, + ), + ("laya", "https://decision.test", {"provider": "laya", "model": "english", "api_base": "https://new.test"}, {}), + ("laya", "https://decision.test", {"provider": "laya", "model": "english", "api_base": None}, {}), + ("laya", None, {"provider": "laya", "model": "english", "api_base": None}, {}), + ("laya", "https://decision.test", {"provider": "typesafe", "model": "jev-latest"}, {}), + ( + "laya", "https://decision.test", {"model": "english", "timeout_ms": 8100}, + {"provider": "laya", "api_base": "https://decision.test", "api_key": "stored-secret"}, + ), + ("typesafe", "https://decision.test", {"provider": "laya", "model": "english"}, {}), + ], + ) + async def test_decision_provider_changes_cannot_reuse_a_stored_key( + self, endpoint: str, stored_provider: str, stored_base: str | None, + supplied: dict[str, object], expected_transport: dict[str, object] + ) -> None: + original: Final = self._row() + row: Final = original.model_copy(update={"litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_config": { + "classifier_type": "jev", "tiers": {"SIMPLE": "allowed"}, + "jev_classifier_config": { + "provider": stored_provider, "model": "english" if stored_provider == "laya" else "jev-latest", + "api_base": stored_base, "api_key": "stored-secret", + }, + }, + }}) + database: Final = self._database(self._team(), row) + config: Final = {"classifier_type": "jev", "tiers": {"SIMPLE": "allowed"}, "jev_classifier_config": supplied} + request: Final = updateDeployment( + litellm_params=updateLiteLLMParams(complexity_router_config=config), model_info=ModelInfo(id=row.model_id), + ) + actor: Final = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + with self._environment(database, row): + await (patch_model(row.model_id, request, actor) if endpoint == "patch" else update_model(request, actor)) + written: Final = database.db.litellm_proxymodeltable.update.await_args.kwargs["data"] + saved: Final = json.loads(written["litellm_params"])["complexity_router_config"] + assert saved == {**config, "jev_classifier_config": {**expected_transport, **supplied}} + + @pytest.mark.asyncio + @pytest.mark.parametrize("endpoint", ["patch", "legacy"]) + @pytest.mark.parametrize( + "string_params,reset_field,config_shape", + [ + (False, None, "full"), (True, None, "full"), (False, "api_key", "full"), + (False, "api_base", "full"), (False, None, "omit-provider"), + (False, None, "omit-config"), (False, None, "null-config"), + ], + ) + async def test_member_save_protects_stored_classifier_connection( + self, endpoint: str, string_params: bool, reset_field: str | None, config_shape: str + ) -> None: + original: Final = self._row() + config: Final = { + "classifier_type": "jev", "tiers": {"SIMPLE": "allowed"}, + "jev_classifier_config": {"provider": "laya", "model": "english"}, + } + secret_params: Final = { + "model": "auto_router/complexity_router", + "complexity_router_config": { + **config, "jev_classifier_config": { + **config["jev_classifier_config"], "api_key": "retained-laya-secret", "api_base": "https://laya.test", + }, + }, + } + row: Final = original.model_copy(update={"litellm_params": secret_params}) + team: Final = self._team().model_copy(update={"models": ["allowed", "laya/english"]}) + database: Final = self._database(team, row) + database.transaction.litellm_proxymodeltable.update.return_value = row.model_copy( + update={"litellm_params": json.dumps(secret_params) if string_params else secret_params} + ) + supplied_config: Final = { + **config, "jev_classifier_config": { + **{ + key: value for key, value in config["jev_classifier_config"].items() + if key != "provider" or config_shape != "omit-provider" + }, + **({reset_field: None} if reset_field is not None else {}), + }, + } + patch_params: Final = ( + {"complexity_router_default_model": "allowed"} + if config_shape == "omit-config" + else {"complexity_router_config": None, "complexity_router_default_model": "allowed"} + if config_shape == "null-config" + else {"complexity_router_config": supplied_config} + ) + request: Final = updateDeployment( + litellm_params=updateLiteLLMParams.model_validate(patch_params), + model_info=ModelInfo(id=row.model_id, team_id="member-team"), + ) + actor: Final = UserAPIKeyAuth( + user_id="owner", user_role=LitellmUserRoles.INTERNAL_USER, models=["allowed", "laya/english"], config={"timeout": 60}, + ) + with self._environment(database, row): + if reset_field is not None: + expected_error: Final = HTTPException if endpoint == "patch" else ProxyException + with pytest.raises(expected_error, match="Team members cannot change classifier connections") as denied: + await ( + patch_model(row.model_id, request, actor) if endpoint == "patch" else update_model(request, actor) + ) + assert ( + denied.value.status_code if isinstance(denied.value, HTTPException) else int(denied.value.code) + ) == 403 + database.transaction.litellm_proxymodeltable.update.assert_not_awaited() + assert row.litellm_params == secret_params + return + response: Final = await (patch_model(row.model_id, request, actor) if endpoint == "patch" else update_model(request, actor)) + written: Final = database.transaction.litellm_proxymodeltable.update.await_args.kwargs["data"] + saved: Final = json.loads(written["litellm_params"])["complexity_router_config"]["jev_classifier_config"] + assert saved == secret_params["complexity_router_config"]["jev_classifier_config"] + if config_shape in ("omit-config", "null-config"): + from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper + + assert decrypt_value_helper( + json.loads(written["litellm_params"])["complexity_router_default_model"], + key="complexity_router_default_model", return_original_value=True, + ) == "allowed" + response_payload: Final = jsonable_encoder(response) + assert "retained-laya-secret" not in json.dumps(response_payload) + response_params: Final = json.loads(response_payload["litellm_params"]) if string_params else response_payload["litellm_params"] + assert response_params == { + **secret_params, "complexity_router_config": { + **config, "jev_classifier_config": { + **config["jev_classifier_config"], "api_key": "REDACTED", "api_base": "https://laya.test", + }, + }, + } + assert "retained-laya-secret" in row.model_dump_json() + @pytest.mark.asyncio @pytest.mark.parametrize("endpoint", ["patch", "legacy"]) @pytest.mark.parametrize("access", ["owner", "peer", "limited-key"]) diff --git a/tests/test_litellm/proxy/management_helpers/test_auto_router_permissions.py b/tests/test_litellm/proxy/management_helpers/test_auto_router_permissions.py index e16271a5189..9d3d079bbca 100644 --- a/tests/test_litellm/proxy/management_helpers/test_auto_router_permissions.py +++ b/tests/test_litellm/proxy/management_helpers/test_auto_router_permissions.py @@ -1,3 +1,4 @@ +import json from collections.abc import Mapping from dataclasses import dataclass from typing import Final @@ -141,6 +142,8 @@ def test_tier_config_is_normalized_and_unknown_router_extras_are_rejected() -> N ({"api_key": "sk-member"}, "api_key"), ({"api_base": "https://collector.invalid", "api_key": "sk-member"}, "api_key"), ({"api_base": "https://collector.invalid", "api_key": ""}, "jev_classifier_config.api_key"), + ({"provider": "laya", "model": "english", "api_base": "https://collector.invalid"}, "api_base"), + ({"provider": "laya", "model": "english", "api_key": "sk-member"}, "api_key"), ], ) def test_members_cannot_move_the_jev_classifier_off_the_proxys_typesafe_account( @@ -154,16 +157,21 @@ def test_members_cannot_move_the_jev_classifier_off_the_proxys_typesafe_account( assert denied.value.detail == f"Invalid member auto-router configuration at {rejected_at}." -def test_members_can_still_tune_the_jev_classifier() -> None: +@pytest.mark.parametrize(("provider", "model"), [("typesafe", "jev-preview"), ("laya", "english")]) +def test_members_can_still_tune_the_jev_classifier(provider: str, model: str) -> None: validated: Final = validate_member_auto_router_config( { "tiers": {"SIMPLE": "allowed"}, "classifier_type": "jev", - "jev_classifier_config": {"model": "jev-preview", "timeout_ms": 500}, + "jev_classifier_config": {"provider": provider, "model": model, "timeout_ms": 500}, } ) assert validated.jev_classifier_config is not None - assert (validated.jev_classifier_config.model, validated.jev_classifier_config.timeout_ms) == ("jev-preview", 500) + assert ( + validated.jev_classifier_config.provider, + validated.jev_classifier_config.model, + validated.jev_classifier_config.timeout_ms, + ) == (provider, model, 500) assert validate_member_auto_router_config(validated.model_dump()).jev_classifier_config is not None @@ -217,6 +225,90 @@ async def test_member_updates_restrict_fields_and_preserve_an_inherited_default( assert granted.default_model == "allowed" +@pytest.mark.asyncio +@pytest.mark.parametrize( + "nested,expected_identity,restricted", + [ + ("omit-config", "laya/english", False), + ("omit-config", "laya/english", True), + ("omit-block", None, False), + (None, None, False), + ({}, None, False), + ({"timeout_ms": 500}, None, False), + ({"model": "english", "timeout_ms": 500}, "laya/english", False), + ({"model": "english", "timeout_ms": 500}, "laya/english", True), + ({"model": "multilingual"}, "laya/multilingual", False), + ({"provider": "typesafe", "model": "jev-latest"}, "typesafe/jev-latest", False), + ({"provider": "typesafe", "model": "jev-latest"}, "typesafe/jev-latest", True), + ], +) +async def test_member_authorization_and_persistence_resolve_the_same_classifier( + catalog: Router, monkeypatch: pytest.MonkeyPatch, nested: object, expected_identity: str | None, restricted: bool +) -> None: + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _strategy_router_write_violation, + update_db_model, + ) + from litellm.types.management_endpoints.auto_router_endpoints import RequestComplexityRouterConfig + + monkeypatch.setenv("LITELLM_SALT_KEY", "member-router-test-salt") + stored_config: Final = { + "classifier_type": "jev", "tiers": {"SIMPLE": "allowed"}, + "jev_classifier_config": { + "provider": "laya", "model": "english", "timeout_ms": 12000, + "api_base": "https://laya.test", "api_key": "stored-classifier-key", + }, + } + existing: Final = Deployment( + model_name="member-router", + litellm_params=LiteLLM_Params(model="auto_router/complexity_router", complexity_router_config=stored_config), + model_info=ModelInfo(id="router-a", team_id="team-a"), created_by="owner", + ) + incoming_config: Final = ( + None if nested == "omit-config" else { + "classifier_type": "jev", "tiers": {"SIMPLE": "allowed"}, + **({} if nested == "omit-block" else {"jev_classifier_config": nested}), + } + ) + patch: Final = updateDeployment.model_validate({"litellm_params": { + "complexity_router_config": incoming_config, "complexity_router_default_model": "allowed", + }}) + operation: Final = authorize_member_auto_router_write( + incoming=patch, existing=existing, user_api_key_dict=_actor( + models=["allowed"] if restricted or expected_identity is None else ["allowed", expected_identity], + ), + team=_team(models=["allowed", "laya/english", "laya/multilingual", "typesafe/jev-latest"]), + premium_user=True, prisma_client=_Client(), llm_router=catalog, + ) + violation: Final = _strategy_router_write_violation(patch.litellm_params, existing.litellm_params) + if expected_identity is None: + assert violation is not None + with pytest.raises(HTTPException) as rejected: + await operation + assert rejected.value.status_code == 400 + return + assert violation is None + if restricted: + with pytest.raises(ProxyException, match=expected_identity): + await operation + return + grant: Final = await operation + persisted: Final = update_db_model(existing, patch) + saved: Final = RequestComplexityRouterConfig.model_validate( + json.loads(persisted["litellm_params"])["complexity_router_config"] + ) + assert grant.config == saved + assert saved.jev_classifier_config is not None + assert f"{saved.jev_classifier_config.provider}/{saved.jev_classifier_config.model}" == expected_identity + assert saved.jev_classifier_config.api_key == ( + "stored-classifier-key" if expected_identity.startswith("laya/") else None + ) + assert saved.jev_classifier_config.timeout_ms == ( + 12000 if nested == "omit-config" else 500 if nested == {"model": "english", "timeout_ms": 500} else 3000 + ) + assert existing.litellm_params.complexity_router_config == stored_config + + @pytest.mark.asyncio @pytest.mark.parametrize("target", ["missing", "nested"]) async def test_member_dependencies_require_plain_configured_models(target: str) -> None: @@ -246,13 +338,17 @@ async def test_member_dependencies_require_plain_configured_models(target: str) @pytest.mark.asyncio @pytest.mark.parametrize("restricted", ["key", "team", None]) +@pytest.mark.parametrize(("provider", "model"), [("typesafe", "jev-latest"), ("laya", "english")]) async def test_jev_evaluation_requires_model_access_but_no_completion_deployment( - catalog: Router, restricted: str | None + catalog: Router, restricted: str | None, provider: str, model: str ) -> None: - permitted: Final = ["allowed", "typesafe/jev-latest"] + permitted: Final = ["allowed", f"{provider}/{model}"] operation: Final = authorize_member_auto_router_dependencies( config=validate_member_auto_router_config( - {"tiers": {"SIMPLE": "allowed"}, "classifier_type": "jev", "jev_classifier_config": {}} + { + "tiers": {"SIMPLE": "allowed"}, "classifier_type": "jev", + "jev_classifier_config": {"provider": provider, "model": model}, + } ), default_model=None, user_api_key_dict=_actor(models=["allowed"] if restricted == "key" else permitted), @@ -261,17 +357,20 @@ async def test_jev_evaluation_requires_model_access_but_no_completion_deployment llm_router=catalog, ) if restricted is not None: - with pytest.raises(ProxyException, match="jev-latest"): + with pytest.raises(ProxyException, match=model): await operation return await operation - assert not catalog.get_model_list("typesafe/jev-latest") + assert not catalog.get_model_list(f"{provider}/{model}") @pytest.mark.asyncio @pytest.mark.parametrize("restricted", ["member", "project", "organization", None]) -async def test_jev_evaluation_obeys_each_containing_scope(catalog: Router, restricted: str | None) -> None: - allowed: Final = ["allowed", "typesafe/jev-latest"] +@pytest.mark.parametrize(("provider", "model"), [("typesafe", "jev-latest"), ("laya", "english")]) +async def test_jev_evaluation_obeys_each_containing_scope( + catalog: Router, restricted: str | None, provider: str, model: str +) -> None: + allowed: Final = ["allowed", f"{provider}/{model}"] membership: Final = LiteLLM_TeamMembership.model_validate( { "user_id": "owner", @@ -293,7 +392,10 @@ async def test_jev_evaluation_obeys_each_containing_scope(catalog: Router, restr ) operation: Final = authorize_member_auto_router_dependencies( config=validate_member_auto_router_config( - {"tiers": {"SIMPLE": "allowed"}, "classifier_type": "jev", "jev_classifier_config": {}} + { + "tiers": {"SIMPLE": "allowed"}, "classifier_type": "jev", + "jev_classifier_config": {"provider": provider, "model": model}, + } ), default_model=None, user_api_key_dict=_actor(models=allowed, project_id="project-a"), @@ -303,8 +405,8 @@ async def test_jev_evaluation_obeys_each_containing_scope(catalog: Router, restr dependency_objects=MemberAutoRouterDependencyObjects(membership, organization, project), ) if restricted is not None: - with pytest.raises(ProxyException, match="jev-latest"): + with pytest.raises(ProxyException, match=model): await operation return await operation - assert not catalog.get_model_list("typesafe/jev-latest") + assert not catalog.get_model_list(f"{provider}/{model}") diff --git a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_typesafe_passthrough_logging_handler.py b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_typesafe_passthrough_logging_handler.py index e0a5ef063e8..acf05dcdfde 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_typesafe_passthrough_logging_handler.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_typesafe_passthrough_logging_handler.py @@ -1,10 +1,12 @@ from datetime import datetime +from typing import Final from unittest.mock import MagicMock import httpx import pytest import litellm +from litellm.litellm_core_utils.litellm_logging import Logging from litellm.proxy.pass_through_endpoints.llm_provider_handlers.typesafe_passthrough_logging_handler import ( TypeSafePassthroughLoggingHandler, ) @@ -137,6 +139,84 @@ def test_success_handler_dispatches_to_typesafe_handler(): assert normalized["kwargs"]["model"] == "typesafe/jev-1.13.0" +@pytest.mark.asyncio +@pytest.mark.parametrize("guardrail_cost", [0.0, 0.25]) +@pytest.mark.parametrize("metadata_slot", ["metadata", "litellm_metadata"]) +@pytest.mark.parametrize("routing_model", ["multilingual", None]) +async def test_laya_gateway_accounts_for_checkpoint_usage_and_registered_cost( + monkeypatch: pytest.MonkeyPatch, routing_model: str | None, metadata_slot: str, guardrail_cost: float +) -> None: + checkpoint: Final = routing_model or "english" + model: Final = f"laya/{checkpoint}" + input_rate: Final = 0.002 + output_rate: Final = 0.005 + monkeypatch.setitem(litellm.model_cost, model, { + "input_cost_per_token": input_rate, "output_cost_per_token": output_rate, + "litellm_provider": "laya", "mode": "evaluation", + }) + start: Final = datetime.now() + logging_obj: Final = Logging( + model="english", messages=[], stream=False, call_type="pass_through_endpoint", + start_time=start, litellm_call_id="laya-accounting", function_id="laya-accounting", kwargs={}, + ) + from fastapi import Request + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.pass_through_endpoints.pass_through_endpoints import HttpPassThroughEndpointHelpers + + request: Final = Request({ + "type": "http", "method": "POST", "path": "/laya/v1/systemone", + "headers": [], "query_string": b"", + }) + auth: Final = UserAPIKeyAuth( + api_key="laya-budget-key", token="laya-budget-key", + model_max_budget={"laya/english": {"budget_limit": 0.01, "time_period": "1d"}}, + ) + request_body: Final = {"model": "english", metadata_slot: {"model_group": "unbounded-client-choice"}} + logging_kwargs: Final = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint( + request=request, user_api_key_dict=auth, logging_obj=logging_obj, + passthrough_logging_payload={"url": "https://laya.test/v1/systemone"}, _parsed_body=request_body, + ) + logging_kwargs["litellm_params"]["metadata"]["standard_logging_guardrail_information"] = [ + {"guardrail_name": "trusted-hook", "guardrail_cost": guardrail_cost}, + ] + logging_obj.update_environment_variables( + model="english", user="unknown", optional_params={}, + litellm_params=logging_kwargs["litellm_params"], call_type="pass_through_endpoint", + ) + body: Final = { + "model": "laya-rl-agent", "usage": {"input_tokens": 10, "output_tokens": 3}, + **({"routing": {"model": routing_model}} if routing_model else {}), + } + normalized: Final = PassThroughEndpointLogging().normalize_llm_passthrough_logging_payload( + httpx_response=httpx.Response(200, request=httpx.Request("POST", "https://laya.test/v1/systemone"), json=body), + response_body=body, request_body={"model": "english"}, logging_obj=logging_obj, + url_route="https://laya.test/v1/systemone", result="{}", start_time=start, + end_time=datetime.now(), cache_hit=False, custom_llm_provider="laya", **logging_kwargs, + ) + logged: Final = normalized["kwargs"] + expected_cost: Final = 10 * input_rate + 3 * output_rate + assert (logged["model"], logged["custom_llm_provider"]) == (model, "laya") + assert logged["response_cost"] == pytest.approx(expected_cost) + assert logged["combined_usage_object"].model_dump(exclude_none=True) == { + "prompt_tokens": 10, "completion_tokens": 3, "total_tokens": 13, + } + assert logging_obj.model_call_details["model"] == model + assert logging_obj.model_call_details["response_cost"] == pytest.approx(expected_cost) + assert logged["standard_logging_object"]["model"] == model + assert logged["standard_logging_object"]["model_group"] == "laya/english" + assert logged["standard_logging_object"]["response_cost"] == pytest.approx(expected_cost + guardrail_cost) + + from litellm.caching.caching import DualCache + from litellm.exceptions import BudgetExceededError + from litellm.proxy.hooks.model_max_budget_limiter import _PROXY_VirtualKeyModelMaxBudgetLimiter + + budget_limiter: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(DualCache()) + assert await budget_limiter.is_key_within_model_budget(auth, "laya/english") + await budget_limiter.async_log_success_event(logged, None, start, datetime.now()) + with pytest.raises(BudgetExceededError): + await budget_limiter.is_key_within_model_budget(auth, "laya/english") + + def test_openrouter_decisions_response_is_priced_from_request_model_registry_row(): logging_obj = _logging_obj() model_cost = litellm.model_cost["openrouter/typesafe/jev-1.13"] diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 227921d6150..06c43c06c64 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -23,6 +23,8 @@ from starlette.datastructures import FormData import litellm +from litellm.caching.caching import DualCache +from litellm.types.utils import CallTypesLiteral from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing from tests.unit.llms.bedrock.event_loop_probe import EventLoopProbe from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS @@ -7119,6 +7121,152 @@ class TestTypeSafePassthroughRoute: ) +class TestLayaPassthroughRoute: + @pytest.fixture + def client(self, monkeypatch: pytest.MonkeyPatch) -> Iterator[TestClient]: + from litellm.proxy.proxy_server import app + + monkeypatch.setenv("LAYA_API_BASE", "http://laya.test/base") + monkeypatch.setenv("TYPESAFE_API_KEY", "never-send-typesafe-key") + monkeypatch.delenv("LAYA_API_KEY", raising=False) + monkeypatch.delenv("SERVER_ROOT_PATH", raising=False) + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + monkeypatch.setitem(app.dependency_overrides, user_api_key_auth, lambda: UserAPIKeyAuth(api_key="sk-virtual")) + yield TestClient(app) + + @pytest.mark.parametrize("api_key", [None, "laya-provider-key"]) + def test_laya_forwards_native_decisions_without_gateway_or_typesafe_credentials( + self, client: TestClient, monkeypatch: pytest.MonkeyPatch, api_key: str | None + ) -> None: + if api_key is not None: + monkeypatch.setenv("LAYA_API_KEY", api_key) + body: Final = { + "model": "english", + "state": "refund", + "questions": {"department": {"type": "choice", "criteria": {"billing": "refunds"}}}, + } + answer: Final = {"model": "laya-rl-agent", "routing": {"model": "english"}, "answers": {}} + with respx.mock(assert_all_called=True) as upstream: + route: Final = upstream.post("http://laya.test/base/v1/systemone?trace=yes").respond(200, json=answer) + response: Final = client.post( + "/laya/v1/systemone?trace=yes", + json=body, + headers={"Authorization": "Bearer sk-virtual", "x-pass-authorization": "Bearer attacker"}, + ) + + assert (response.status_code, response.json()) == (200, answer) + sent: Final = route.calls.last.request + assert sent.headers.get("authorization") == (f"Bearer {api_key}" if api_key else None) + assert json.loads(sent.content) == body + + def test_laya_missing_server_fails_without_contacting_another_provider( + self, client: TestClient, monkeypatch: pytest.MonkeyPatch + ) -> None: + monkeypatch.delenv("LAYA_API_BASE") + with respx.mock(assert_all_called=False) as upstream: + response: Final = client.post("/laya/v1/systemone", json={"model": "english"}) + assert response.status_code == 503 + assert "LAYA_API_BASE" in response.text + assert len(upstream.calls) == 0 + + def test_laya_does_not_forward_unsupported_endpoints(self, client: TestClient) -> None: + with respx.mock(assert_all_called=False) as upstream: + response: Final = client.post("/laya/v1/evaluate", json={"model": "english"}) + assert response.status_code == 404 + assert len(upstream.calls) == 0 + + @pytest.mark.parametrize("model", [None, "auto", "jev-latest"]) + def test_laya_rejects_implicit_checkpoint_selection(self, client: TestClient, model: str | None) -> None: + with respx.mock(assert_all_called=False) as upstream: + response: Final = client.post("/laya/v1/systemone", json={"model": model}) + assert response.status_code == 400 + assert len(upstream.calls) == 0 + + @pytest.mark.parametrize( + "controls", + [{"custom_body": {"model": "multilingual", "state": "refund"}}, {"stream": True}, {"stream": "true"}], + ) + def test_laya_rejects_controls_that_change_authorized_body_or_usage_accounting( + self, client: TestClient, controls: Mapping[str, object] + ) -> None: + with respx.mock(assert_all_called=False) as upstream: + route: Final = upstream.post("http://laya.test/base/v1/systemone").respond(200, json={"answers": {}}) + response: Final = client.post("/laya/v1/systemone", json={"model": "english", **controls}) + assert response.status_code == 400 + assert not route.called + + + @pytest.mark.parametrize("metadata_slot", ["metadata", "litellm_metadata"]) + def test_laya_hooks_enforce_canonical_model_limits_and_keep_native_wire_body( + self, client: TestClient, monkeypatch: pytest.MonkeyPatch, metadata_slot: str + ) -> None: + from litellm.integrations.custom_logger import CustomLogger + from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 + from litellm.proxy.utils import InternalUsageCache + from litellm.proxy.proxy_server import app + + cache: Final = DualCache() + limiter: Final = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(cache)) + auth: Final = UserAPIKeyAuth( + api_key="laya-native-rpm", metadata={"model_rpm_limit": {"laya/english": 1}}, + ) + def authenticated_key() -> UserAPIKeyAuth: + return auth + + monkeypatch.setitem(app.dependency_overrides, user_api_key_auth, authenticated_key) + + class LimitHook(CustomLogger): + async def async_pre_call_hook( + self, user_api_key_dict: UserAPIKeyAuth, cache: DualCache, + data: dict[str, object], call_type: CallTypesLiteral, + ) -> dict[str, object]: + assert data["model"] == "laya/english" + metadata: Final = data.get(metadata_slot) + assert isinstance(metadata, dict) + assert "standard_logging_guardrail_information" not in metadata + assert metadata["customer_label"] == "retained" + await limiter.async_pre_call_hook(user_api_key_dict, cache, data, call_type) + return data + + monkeypatch.setattr(litellm, "callbacks", [LimitHook()]) + body: Final = { + "model": "english", "state": "refund", + metadata_slot: { + "customer_label": "retained", "model_group": "unbounded-client-choice", + "standard_logging_guardrail_information": [{"guardrail_cost": 25.0}], + }, + } + with respx.mock(assert_all_called=True) as upstream: + route: Final = upstream.post("http://laya.test/base/v1/systemone").respond(200, json={"answers": {}}) + first: Final = client.post("/laya/v1/systemone", json=body) + second: Final = client.post("/laya/v1/systemone", json=body) + assert first.status_code == 200, first.text + assert second.status_code == 429, second.text + assert route.call_count == 1 + assert json.loads(route.calls.last.request.content) == {"model": "english", "state": "refund"} + + def test_laya_preserves_trusted_hook_checkpoint_changes( + self, client: TestClient, monkeypatch: pytest.MonkeyPatch + ) -> None: + from litellm.integrations.custom_logger import CustomLogger + + class CheckpointHook(CustomLogger): + async def async_pre_call_hook( + self, user_api_key_dict: UserAPIKeyAuth, cache: DualCache, + data: dict[str, object], call_type: CallTypesLiteral, + ) -> dict[str, object]: + assert data["model"] == "laya/english" + return {**data, "model": "laya/multilingual"} + + monkeypatch.setattr(litellm, "callbacks", [CheckpointHook()]) + with respx.mock(assert_all_called=True) as upstream: + route: Final = upstream.post("http://laya.test/base/v1/systemone").respond(200, json={"answers": {}}) + response: Final = client.post("/laya/v1/systemone", json={"model": "english", "state": "refund"}) + assert response.status_code == 200, response.text + assert json.loads(route.calls.last.request.content) == {"model": "multilingual", "state": "refund"} + + class TestFalAIPassthroughRoute: @pytest.fixture def client(self, monkeypatch: pytest.MonkeyPatch) -> Iterator[TestClient]: diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 3469df082e0..96457d1e58c 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -1468,7 +1468,7 @@ async def test_pass_through_request_contains_proxy_server_request_in_kwargs(): # Create mock request mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://test-proxy.com/api/endpoint" + mock_request.url = httpx.URL("http://test-proxy.com/api/endpoint") mock_request.body = AsyncMock(return_value=b'{"message": "test request"}') mock_request.headers = Headers({}) mock_request.query_params = QueryParams({}) @@ -1573,7 +1573,7 @@ async def test_pass_through_request_streaming_marks_logging_obj_as_stream(): mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://test-proxy.com/v1/messages" + mock_request.url = httpx.URL("http://test-proxy.com/v1/messages") mock_request.body = AsyncMock(return_value=b'{"model": "claude-3", "stream": true}') mock_request.headers = Headers({}) mock_request.query_params = QueryParams({}) @@ -1635,7 +1635,7 @@ async def test_pass_through_request_sse_response_marks_logging_obj_as_stream(): mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://test-proxy.com/v1/messages" + mock_request.url = httpx.URL("http://test-proxy.com/v1/messages") mock_request.body = AsyncMock(return_value=b'{"model": "claude-3"}') mock_request.headers = Headers({}) mock_request.query_params = QueryParams({}) @@ -2505,7 +2505,7 @@ async def test_pass_through_request_query_params_forwarding(): # Create mock request with query parameters (Azure API version) mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://localhost:4000/azure-assistant/openai/assistants" + mock_request.url = httpx.URL("http://localhost:4000/azure-assistant/openai/assistants") mock_request.body = AsyncMock(return_value=json.dumps(test_body).encode()) mock_request.headers = Headers({"Content-Type": "application/json"}) @@ -3016,7 +3016,7 @@ async def test_bedrock_router_passthrough_metadata_initialization(): # Create mock request with headers mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://localhost:4000/bedrock/model/my-model/invoke" + mock_request.url = httpx.URL("http://localhost:4000/bedrock/model/my-model/invoke") mock_request.headers = Headers( { "content-type": "application/json", @@ -3850,7 +3850,7 @@ def _lit3538_request(): r = MagicMock() r.method = "POST" r.query_params = {} - r.url = "http://testserver/mock/echo" + r.url = httpx.URL("http://testserver/mock/echo") r.state = SimpleNamespace() headers = MagicMock() headers.copy.return_value = {} @@ -3983,7 +3983,7 @@ async def test_pass_through_request_non_streaming_upstream_error_returned_unchan mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://test-proxy.com/mock-upstream/api/denied" + mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/denied") mock_request.body = AsyncMock(return_value=b'{"action": "read"}') mock_request.headers = Headers({"content-type": "application/json"}) mock_request.query_params = QueryParams({}) @@ -4069,7 +4069,7 @@ async def test_pass_through_request_upstream_error_failure_hook_exception_is_swa mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://test-proxy.com/mock-upstream/api/denied" + mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/denied") mock_request.body = AsyncMock(return_value=b'{"action": "read"}') mock_request.headers = Headers({"content-type": "application/json"}) mock_request.query_params = QueryParams({}) @@ -4118,7 +4118,7 @@ async def test_pass_through_request_streaming_upstream_error_returned_unchanged( mock_request = MagicMock(spec=Request) mock_request.method = "GET" - mock_request.url = "http://test-proxy.com/mock-upstream/api/stream-denied" + mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/stream-denied") mock_request.body = AsyncMock(return_value=b"") mock_request.headers = Headers({}) mock_request.query_params = QueryParams({}) @@ -4169,7 +4169,7 @@ class _UpstreamErrorBodyStream(httpx.AsyncByteStream): def _upstream_error_request() -> MagicMock: mock_request: Final = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://test-proxy.com/mock-upstream/v1beta/models/claude-nope-9:generateContent" + mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/v1beta/models/claude-nope-9:generateContent") mock_request.body = AsyncMock(return_value=b'{"contents": []}') mock_request.headers = Headers({"content-type": "application/json"}) mock_request.query_params = QueryParams({}) @@ -4964,7 +4964,7 @@ async def test_pass_through_request_non_streaming_success_unchanged(): mock_request = MagicMock(spec=Request) mock_request.method = "GET" - mock_request.url = "http://test-proxy.com/mock-upstream/api/success" + mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/success") mock_request.body = AsyncMock(return_value=b"") mock_request.headers = Headers({}) mock_request.query_params = QueryParams({}) @@ -5027,7 +5027,7 @@ async def test_pass_through_request_claims_the_budget_reservation_only_when_its_ mock_get_client.return_value = MagicMock(client=async_client) mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://test-proxy.com/mock-upstream/api/generate" + mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/generate") mock_request.body = AsyncMock(return_value=b'{"prompt": "hi"}') mock_request.headers = Headers({"content-type": "application/json"}) mock_request.query_params = QueryParams({}) @@ -5079,7 +5079,7 @@ async def test_pass_through_request_leaves_the_budget_reservation_for_the_reques mock_get_client.return_value = MagicMock(client=async_client) mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://test-proxy.com/mock-upstream/api/generate" + mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/generate") mock_request.body = AsyncMock(return_value=b'{"prompt": "hi"}') mock_request.headers = Headers({"content-type": "application/json"}) mock_request.query_params = QueryParams({}) @@ -5110,7 +5110,7 @@ async def test_pass_through_request_internal_failure_still_raises_proxy_exceptio mock_request = MagicMock(spec=Request) mock_request.method = "GET" - mock_request.url = "http://test-proxy.com/mock-upstream/api/success" + mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/success") mock_request.body = AsyncMock(return_value=b"") mock_request.headers = Headers({}) mock_request.query_params = QueryParams({}) @@ -5211,7 +5211,7 @@ def _enter_relay_logging_mocks(stack, parsed_body): def _relay_client_request(method="GET"): mock_request = MagicMock(spec=Request) mock_request.method = method - mock_request.url = "http://localhost:4000/passthrough-relay/results" + mock_request.url = httpx.URL("http://localhost:4000/passthrough-relay/results") mock_request.body = AsyncMock(return_value=b"") mock_request.headers = Headers({}) mock_request.query_params = QueryParams({}) @@ -6587,7 +6587,7 @@ def _passthrough_kwargs_for_reservation( ) -> dict: mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent" + mock_request.url = httpx.URL("http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent") mock_request.headers = Headers({}) mock_request.scope = {"endpoint": _marked_pass_through_endpoint()} if user_defined_route else {} @@ -6734,7 +6734,7 @@ async def _drive_streaming_pass_through(upstream_content_type, chunk_delay_secon mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://test-proxy.com/v1/messages" + mock_request.url = httpx.URL("http://test-proxy.com/v1/messages") mock_request.body = AsyncMock( return_value=b'{"model": "claude-3", "stream": true}' if client_asked_for_stream @@ -6922,36 +6922,72 @@ def _marked_pass_through_endpoint(): return _endpoint -def test_user_defined_passthrough_is_neither_tracked_nor_enforced(): - """ - `get_model_from_request` returns None for a user-defined pass-through on - purpose: the body is forwarded verbatim, so its `model` names an UPSTREAM - model rather than a LiteLLM-managed one, and enforcing key/team allowlists - against it would reject valid requests. Enforcement is therefore skipped - on those routes. +@pytest.mark.asyncio +@pytest.mark.parametrize("metadata_slot", ["metadata", "litellm_metadata"]) +async def test_user_defined_passthrough_is_neither_tracked_nor_enforced(metadata_slot: str) -> None: + from datetime import datetime - Attaching the budget metadata anyway would charge a counter that nothing on - that route can refuse, and would attribute the spend to a budget the operator - scoped to a LiteLLM model that merely shares the name. Tracking and - enforcement have to agree: both on for the built-in provider routes, both off - here. - """ - kwargs = _passthrough_kwargs_for_reservation( - UserAPIKeyAuth( - token="hash", - user_id="u-1", - model_max_budget={"claude-opus-4-8": {"budget_limit": 1.0, "time_period": "18h"}}, - ), - user_defined_route=True, + from litellm.caching.caching import DualCache + from litellm.proxy.auth.auth_utils import get_model_from_request + from litellm.proxy.hooks.model_max_budget_limiter import _PROXY_VirtualKeyModelMaxBudgetLimiter + + budget: Final = {"managed-model": {"budget_limit": 0.1, "time_period": "1d"}} + limiter: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(DualCache()) + auth: Final = UserAPIKeyAuth( + api_key="custom-key", token="custom-key", team_id="shared-team", team_model_max_budget=budget, ) - - metadata = kwargs["litellm_params"]["metadata"] - for field in ( - "user_api_key_model_max_budget", - "user_api_key_user_model_max_budget", - "user_api_key_end_user_model_max_budget", - ): - assert field not in metadata, f"{field} was attached on a route that never enforces it" + endpoint: Final = create_pass_through_route( + endpoint="/custom-budget-test", target="https://upstream.test/echo", custom_headers={}, cost_per_request=0.25, + ) + request: Final = Request({ + "type": "http", "method": "POST", "path": "/custom-budget-test", "headers": [], + "query_string": b"", "endpoint": endpoint, + }) + body: Final = { + "model": "upstream-only-model", metadata_slot: { + "model_group": "managed-model", "customer_label": "retained", + "user_api_key_team_model_max_budget": budget, + }, + } + assert get_model_from_request(body, "/custom-budget-test", request=request) is None + assert await limiter.is_team_within_model_budget("shared-team", budget, None, "managed-model") + start: Final = datetime.now() + logging_obj: Final = LiteLLMLoggingObj( + model="upstream-only-model", messages=[], stream=False, call_type="pass_through_endpoint", + start_time=start, litellm_call_id="custom-budget", function_id="custom-budget", kwargs={}, + dynamic_async_success_callbacks=[limiter], + ) + payload: Final = { + "url": "https://upstream.test/echo", "request_body": body, "request_method": "POST", "cost_per_request": 0.25, + } + kwargs: Final = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint( + request=request, user_api_key_dict=auth, passthrough_logging_payload=payload, logging_obj=logging_obj, + _parsed_body=body, litellm_call_id="custom-budget", + ) + logging_obj.update_environment_variables( + model="upstream-only-model", user="unknown", optional_params={}, + litellm_params=kwargs["litellm_params"], call_type="pass_through_endpoint", + ) + response: Final = httpx.Response( + 200, request=httpx.Request("POST", "https://upstream.test/echo"), json={"ok": True}, + ) + await PassThroughEndpointLogging().pass_through_async_success_handler( + httpx_response=response, response_body={"ok": True}, request_body=body, logging_obj=logging_obj, + url_route="https://upstream.test/echo", result=response.text, start_time=start, end_time=datetime.now(), + cache_hit=False, **kwargs, + ) + assert logging_obj.model_call_details["response_cost"] == 0.25 + assert await limiter.is_team_within_model_budget("shared-team", budget, None, "managed-model") + metadata: Final = kwargs["litellm_params"]["metadata"] + assert (metadata["model_group"], metadata["customer_label"]) == ("managed-model", "retained") + assert metadata.keys().isdisjoint({ + "user_api_key_model_max_budget", "user_api_key_team_model_max_budget", + "user_api_key_user_model_max_budget", "user_api_key_end_user_model_max_budget", + }) + assert metadata.keys().isdisjoint({ + "user_api_key_model_max_budget", "user_api_key_team_model_max_budget", + "user_api_key_user_model_max_budget", "user_api_key_end_user_model_max_budget", + }) @pytest.mark.parametrize( @@ -7279,7 +7315,7 @@ def test_passthrough_client_cannot_forge_session_id_omission(client_metadata_key mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent" + mock_request.url = httpx.URL("http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent") mock_request.headers = Headers({}) mock_request.scope = {} @@ -7312,7 +7348,7 @@ def test_passthrough_logs_the_resolved_deployment_model_info_over_the_request_bo the call to (LIT-1761: passthrough successes carried model_id="").""" mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://0.0.0.0:4000/vertex_ai/v1/projects/p/locations/global/publishers/google/models/gemini-3.8-flash:generateContent" + mock_request.url = httpx.URL("http://0.0.0.0:4000/vertex_ai/v1/projects/p/locations/global/publishers/google/models/gemini-3.8-flash:generateContent") mock_request.headers = Headers({}) mock_request.scope = {} mock_request.state = SimpleNamespace( @@ -7344,7 +7380,7 @@ _PROXY_SERVER_REQUEST: Final = TypeAdapter(dict[str, object]) def _split_pass_through_body(body: str) -> _PassThroughSplit: mock_request: Final = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent" + mock_request.url = httpx.URL("http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent") mock_request.headers = Headers() mock_request.scope = MappingProxyType({}) @@ -7600,7 +7636,7 @@ def test_passthrough_attributes_a_cli_session_to_its_alias_not_the_login_token() mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://0.0.0.0:4000/anthropic/v1/messages" + mock_request.url = httpx.URL("http://0.0.0.0:4000/anthropic/v1/messages") mock_request.headers = Headers({}) mock_request.scope = {} session = UserAPIKeyAuth( diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_guardrail_block_otel_span.py b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_guardrail_block_otel_span.py index 73927e92c15..09987b2781c 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_guardrail_block_otel_span.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_guardrail_block_otel_span.py @@ -18,7 +18,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest -from fastapi import HTTPException +from fastapi import HTTPException, Request pytest.importorskip("opentelemetry") @@ -81,15 +81,16 @@ def _user_api_key_dict(): return d -def _mock_request(): - r = MagicMock() - r.method = "POST" - r.query_params = {} - r.url = "http://testserver/mock/echo" - headers = MagicMock() - headers.copy.return_value = {} - r.headers = headers - return r +def _mock_request() -> Request: + return Request({ + "type": "http", + "method": "POST", + "scheme": "http", + "server": ("testserver", 80), + "path": "/mock/echo", + "headers": [], + "query_string": b"", + }) def _httpx_response(text: str) -> httpx.Response: diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py b/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py index 29a635e9b27..d6b69c7c010 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py @@ -1,8 +1,8 @@ +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi import Request -from starlette.datastructures import Headers, State from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( @@ -771,12 +771,15 @@ async def test_vertex_passthrough_attributes_the_call_to_the_resolved_deployment """The router deployment that rewrote the upstream URL is the one the logging kwargs must name, so the Prometheus model_id label (and SpendLogs.model_id) on a Vertex passthrough success reads the deployment's id instead of "" (LIT-1761).""" - mock_request = MagicMock(spec=Request) - mock_request.method = "POST" - mock_request.url = "http://0.0.0.0:4000/vertex_ai/v1/projects/p/locations/global/publishers/google/models/gemini-3.8-flash:generateContent" - mock_request.headers = Headers({}) - mock_request.scope = {} - mock_request.state = State() + mock_request: Final = Request({ + "type": "http", + "method": "POST", + "scheme": "http", + "server": ("0.0.0.0", 4000), + "path": "/vertex_ai/v1/projects/p/locations/global/publishers/google/models/gemini-3.8-flash:generateContent", + "headers": [], + "query_string": b"", + }) mock_handler = MagicMock() mock_handler.get_default_base_target_url.return_value = "https://aiplatform.googleapis.com" diff --git a/tests/unit/llms/laya/__init__.py b/tests/unit/llms/laya/__init__.py new file mode 100644 index 00000000000..8b137891791 --- /dev/null +++ b/tests/unit/llms/laya/__init__.py @@ -0,0 +1 @@ + diff --git a/tests/unit/llms/laya/test_common_utils.py b/tests/unit/llms/laya/test_common_utils.py new file mode 100644 index 00000000000..c9ee0062cd2 --- /dev/null +++ b/tests/unit/llms/laya/test_common_utils.py @@ -0,0 +1,60 @@ +from collections.abc import Mapping +from typing import Final + +import pytest + +from litellm.llms.laya.common_utils import laya_connection, laya_response_model + + +@pytest.mark.parametrize( + ("base", "key", "expected_base", "expected_key"), + [ + (None, None, "http://laya.test/root", "laya-env-key"), + ("http://custom.test/", None, "http://custom.test", None), + ("http://custom.test/", "explicit-key", "http://custom.test", "explicit-key"), + ], +) +def test_laya_credentials_stay_with_their_configured_destination( + monkeypatch: pytest.MonkeyPatch, + base: str | None, + key: str | None, + expected_base: str, + expected_key: str | None, +) -> None: + monkeypatch.setenv("LAYA_API_BASE", "http://laya.test/root/") + monkeypatch.setenv("LAYA_API_KEY", "laya-env-key") + monkeypatch.setenv("TYPESAFE_API_KEY", "never-send-this") + connection: Final = laya_connection(base, key) + assert (connection.api_base, connection.api_key) == (expected_base, expected_key) + assert "key" not in repr(connection) + + +@pytest.mark.parametrize( + "base", + ["", "ftp://laya.test", "http://user:password@laya.test", "https://laya.test?key=x", "http://laya.test/#x"], +) +def test_laya_rejects_ambiguous_server_urls(base: str) -> None: + with pytest.raises(ValueError, match="Laya"): + laya_connection(base) + + +def test_laya_missing_server_does_not_fall_back_to_typesafe(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("LAYA_API_BASE", raising=False) + monkeypatch.setenv("TYPESAFE_API_BASE", "https://typesafe.test") + with pytest.raises(ValueError, match="LAYA_API_BASE"): + laya_connection() + + +@pytest.mark.parametrize( + ("routing", "requested", "expected"), + [ + ({"model": "multilingual"}, "english", "multilingual"), + (None, "english", "english"), + ({"model": 42}, "english", "english"), + (None, None, "unknown"), + ], +) +def test_laya_identity_tracks_the_checkpoint_not_the_shared_agent_name( + routing: Mapping[str, object] | None, requested: str | None, expected: str +) -> None: + assert laya_response_model({"model": "laya-rl-agent", "routing": routing}, requested) == expected diff --git a/tests/unit/router_strategy/complexity_router/test_jev_classifier.py b/tests/unit/router_strategy/complexity_router/test_jev_classifier.py index 45070dfd3a7..3fab380bfc6 100644 --- a/tests/unit/router_strategy/complexity_router/test_jev_classifier.py +++ b/tests/unit/router_strategy/complexity_router/test_jev_classifier.py @@ -8,6 +8,7 @@ from unittest.mock import create_autospec import httpx import pytest +import respx import litellm from litellm._logging import verbose_router_logger @@ -30,14 +31,15 @@ from litellm.types.utils import AUTOROUTER_CLASSIFIER_CALL_ORIGIN class _UsageRecorder(CustomLogger): - def __init__(self) -> None: + def __init__(self, model_key: str = "typesafe/jev-accounting") -> None: super().__init__() + self.model_key = model_key self.calls: tuple[Mapping[str, object], ...] = () async def async_log_success_event( self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime ) -> None: - if str(kwargs.get("model", "")).removeprefix("typesafe/") != "jev-accounting": + if str(kwargs.get("model", "")) != self.model_key: return self.calls = (*self.calls, kwargs) @@ -420,6 +422,62 @@ def test_jev_config_requires_classifier_config() -> None: ComplexityRouterConfig.model_validate({"classifier_type": "jev"}) +@pytest.mark.parametrize("config", [{"provider": "laya"}, {"provider": "laya", "model": " "}]) +def test_laya_requires_its_own_checkpoint(config: Mapping[str, object]) -> None: + with pytest.raises(ValueError, match="Laya model must be"): + JevClassifierConfig.model_validate(config) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("custom_base", [False, True]) +async def test_laya_routes_with_its_own_credentials_and_accounts_the_checkpoint( + monkeypatch: pytest.MonkeyPatch, custom_base: bool +) -> None: + monkeypatch.setenv("TYPESAFE_API_KEY", "never-send-typesafe-key") + monkeypatch.setenv("LAYA_API_BASE", "https://laya.test") + monkeypatch.setenv("LAYA_API_KEY", "laya-env-key") + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + monkeypatch.setitem(litellm.model_cost, "laya/english", {"input_cost_per_token": 0.01}) + recorder: Final = _UsageRecorder("laya/english") + monkeypatch.setattr(litellm, "_async_success_callback", [recorder]) + router: Final = ComplexityRouter( + "laya-route", + litellm.Router(model_list=[]), + { + "classifier_type": "jev", + "jev_classifier_config": { + "provider": "laya", + "model": "english", + **({"api_base": "https://laya.test"} if custom_base else {}), + }, + "tiers": {"SIMPLE": "cheap"}, + }, + derive_savings_baseline=False, + ) + with respx.mock(assert_all_called=True) as upstream: + route: Final = upstream.post("https://laya.test/v1/systemone").respond( + 200, + json={ + "model": "laya-rl-agent", + "routing": {"model": "english"}, + "answers": {"tier": _answer().model_dump()}, + "usage": {"input_tokens": 31, "output_tokens": 0}, + }, + ) + outcome: Final = await router.aclassify("choose a tier") + await GLOBAL_LOGGING_WORKER.flush() + + assert outcome.cause == "jev_classifier" + assert outcome.jev_verdict is not None + assert (outcome.jev_verdict.provider, outcome.jev_verdict.model) == ("laya", "english") + assert outcome.classifier_cost == pytest.approx(0.31) + sent: Final = route.calls.last.request + assert sent.headers.get("authorization") == (None if custom_base else "Bearer laya-env-key") + assert json.loads(sent.content)["model"] == "english" + assert len(recorder.calls) == 1 + assert recorder.calls[0]["response_cost"] == pytest.approx(0.31) + + def test_jev_config_is_rejected_for_other_classifier_types() -> None: with pytest.raises(ValueError, match="has no effect"): ComplexityRouterConfig.model_validate( diff --git a/tests/unit/router_utils/test_auto_router_model_naming.py b/tests/unit/router_utils/test_auto_router_model_naming.py index 645f9e5e62a..1ffb0f8194d 100644 --- a/tests/unit/router_utils/test_auto_router_model_naming.py +++ b/tests/unit/router_utils/test_auto_router_model_naming.py @@ -23,21 +23,23 @@ COMPLEXITY_FIELDS = frozenset({"complexity_router_config"}) SEMANTIC_FIELDS = frozenset({"auto_router_config", "auto_router_default_model", "auto_router_embedding_model"}) -@pytest.mark.parametrize("model", ["jev-latest", "jev-preview"]) -def test_jev_enumerates_a_paid_evaluation_without_a_completion_classifier(model: str) -> None: - found = strategy_router_dependencies( +@pytest.mark.parametrize( + ("provider", "model"), [(None, "jev-latest"), ("typesafe", "jev-preview"), ("laya", "english")] +) +def test_jev_enumerates_a_paid_evaluation_without_a_completion_classifier(provider: str | None, model: str) -> None: + found: Final = strategy_router_dependencies( { "model": "auto_router/complexity_router", "complexity_router_config": { "classifier_type": "jev", - "jev_classifier_config": {"model": model}, + "jev_classifier_config": {"model": model, **({"provider": provider} if provider else {})}, "tiers": {"SIMPLE": "cheap"}, }, } ) assert tuple((dep.model_name, dep.role) for dep in found) == ( ("cheap", "tier"), - (f"typesafe/{model}", "evaluation"), + (f"{provider or 'typesafe'}/{model}", "evaluation"), ) diff --git a/tests/unit/test_utils.py b/tests/unit/test_utils.py index 2c612aa350c..a48811ec1b8 100644 --- a/tests/unit/test_utils.py +++ b/tests/unit/test_utils.py @@ -1008,6 +1008,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "/v1/videos", "/vertex_ai/live", "/v1/listen", + "/v1/systemone", "/v1beta/interactions", ], }, diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 16325970ed1..bd49e9151ff 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -8575,6 +8575,23 @@ export interface paths { patch: operations["langfuse_proxy_route_langfuse__endpoint__patch"]; trace?: never; }; + "/laya/v1/systemone": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** Laya Proxy Route */ + post: operations["laya_proxy_route_laya_v1_systemone_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/lazy/warm/{name}": { parameters: { query?: never; @@ -31536,12 +31553,12 @@ export interface components { JevClassifierConfig: { /** * Api Base - * @description TypeSafe API base, falling back to TYPESAFE_API_BASE and then https://api.typesafe.ai + * @description Provider API base; defaults to TYPESAFE_API_BASE or LAYA_API_BASE for the selected provider */ api_base?: string | null; /** * Api Key - * @description TypeSafe API key, falling back to TYPESAFE_API_KEY + * @description Provider API key; optional for self-hosted Laya */ api_key?: string | null; /** @@ -31564,6 +31581,12 @@ export interface components { * @default jev-latest */ model: string; + /** + * Provider + * @default typesafe + * @enum {string} + */ + provider: "typesafe" | "laya"; /** * Timeout Ms * @default 3000 @@ -58274,6 +58297,26 @@ export interface operations { }; }; }; + laya_proxy_route_laya_v1_systemone_post: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + }; + }; warm_lazy_warm__name__post: { parameters: { query?: never;