feat: add Laya gateway and OSS classifier providers (#43626)

* feat: add Laya gateway and classifier backend

* test: cover the pass-through model_group pin and repair the shard fakes

MockRequest in tests/pass_through_unit_tests gains an httpx.URL and an ASGI
scope, which get_request_route now reads inside
_init_kwargs_for_pass_through_endpoint, and the POST-only /laya/v1/systemone
route joins the protocol-constrained exemptions. A built-in pass-through pins
metadata.model_group to the resolved model so a client cannot choose its own
per-model budget key; test_pass_through_endpoints now proves that on a
non-Laya route and drops a duplicated assertion.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>

---------

Co-authored-by: Claude Fable 5.1 <noreply@anthropic.com>
This commit is contained in:
tin-berri 2026-10-02 11:29:36 -07:00 • committed by GitHub
parent 4758fce91a
commit 0dc23406eb
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
35 changed files with 1672 additions and 270 deletions

View file

@ -0,0 +1 @@

View file

@ -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 _LayaRouting())
except ValidationError:
return requested_model or "unknown"
return routing.model or requested_model or "unknown"

View file

@ -72622,6 +72622,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",

View file

@ -229,6 +229,7 @@ LAZY_FEATURES: Final[tuple[LazyFeature, ...]] = (
"/tinyfish/",
"/transcribe",
"/typesafe/",
"/laya/",
"/openrouter/",
"/vertex-ai/",
"/vertex_ai/",

View file

@ -27761,6 +27761,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.",

View file

@ -507,6 +507,7 @@ class LiteLLMRoutes(enum.Enum):
"/vllm",
"/mistral",
"/typesafe",
"/laya",
"/openrouter",
"/milvus",
"/gigachat",

View file

@ -1883,6 +1883,15 @@ 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_request: Final = TypeAdapter(Mapping[str, object]).validate_python(request_data)
laya_model: Final = validate_laya_model(laya_request.get("model"))
except ValueError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
return _dedupe_model_candidates((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)

View file

@ -316,7 +316,7 @@ async def _authorize_models_this_test_can_call(
its calls through the proxy. Team and member budgets are already enforced on every route.
"""
models: Final = _models_this_test_can_call(config)
if not models and config.classifier_type != "jev":
if not models and config.classifier_type != "oss_classifier":
return
from litellm.proxy.proxy_server import proxy_logging_obj
@ -342,9 +342,9 @@ async def _authorize_models_this_test_can_call(
code=status.HTTP_400_BAD_REQUEST,
) from e
if config.classifier_type == "jev" and user_api_key_dict.budget_throttle_pct is not None:
if config.classifier_type == "oss_classifier" and user_api_key_dict.budget_throttle_pct is not None:
raise ProxyException(
message="Budget has been exceeded! JEV Test Routing requires available budget.",
message="Budget has been exceeded! OSS Classifier Test Routing requires available budget.",
type=ProxyErrorTypes.budget_exceeded,
param=None,
code=status.HTTP_400_BAD_REQUEST,

View file

@ -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,
@ -115,6 +116,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,
@ -187,6 +189,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]: ...
@ -407,34 +428,13 @@ 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 {
**incoming,
"jev_classifier_config": {
**transport,
**supplied,
},
}
config_adapter: Final[TypeAdapter[Mapping[str, object] | None]] = TypeAdapter(Mapping[str, object] | None)
return resolve_complexity_router_config_write(
config_adapter.validate_python(incoming), config_adapter.validate_python(existing)
).effective
def _effective_model(
@ -1304,7 +1304,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)
@ -1501,10 +1501,18 @@ async def _add_model_to_db(
slot: AbstractAsyncContextManager[_ProxyModelTable] | None = None,
) -> "_ProxyModelRow | LiteLLM_ProxyModelTable":
# encrypt litellm params #
_litellm_params_dict: Final = model_params.litellm_params.dict(exclude_none=True)
_litellm_params_dict: Final = TypeAdapter(dict[str, object]).validate_python(
model_params.litellm_params.model_dump(exclude_none=True)
)
if "complexity_router_config" in _litellm_params_dict:
_litellm_params_dict["complexity_router_config"] = _effective_complexity_router_config(
model_params.litellm_params, None
)
_original_litellm_model_name: Final = model_params.litellm_params.model
for k, v in _litellm_params_dict.items():
encrypted_value = encrypt_value_helper(value=v, new_encryption_key=new_encryption_key)
encrypted_value = (
encrypt_value_helper(value=v, new_encryption_key=new_encryption_key) if isinstance(v, str) else v
)
model_params.litellm_params[k] = encrypted_value
_data: Final[dict] = {
"model_id": model_params.model_info.id,
@ -2536,7 +2544,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)
@ -2760,7 +2768,7 @@ async def update_model(
live_after=reload_outcome.live_after,
)
return model_response
return None if model_response is None else _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):

View file

@ -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
@ -65,12 +69,12 @@ class _MemberRouterGenerationParams(BaseModel):
stop: str | tuple[str, ...] | None = None
class _MemberJevClassifierConfig(BaseModel):
"""The Jev classifier settings a team member may set. Credentials stay the proxy's own: a member-chosen
api_base would receive the proxy's TYPESAFE_API_KEY, and a member-chosen api_key would be sent from the proxy."""
class _MemberOpenSourceClassifierConfig(BaseModel):
"""Classifier settings a team member may set while the gateway owns the connection."""
model_config = ConfigDict(extra="forbid")
provider: Literal["jev", "laya"] = "jev"
model: str
api_key: None = None
api_base: None = None
@ -123,14 +127,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.opensource_classifier_config is not None:
_MemberOpenSourceClassifierConfig.model_validate(validated.opensource_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
@ -332,16 +343,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

View file

@ -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 ConfigDict, TypeAdapter
from starlette.websockets import WebSocketState
from typing_extensions import ReadOnly, TypedDict
@ -57,6 +58,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
@ -636,6 +638,47 @@ 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"],
)
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"),
)
authorization: Final[Mapping[str, str]] = (
MappingProxyType({"Authorization": f"Bearer {connection.api_key}"})
if connection.api_key
else MappingProxyType({})
)
endpoint_func: Final = create_pass_through_route(
endpoint="v1/systemone",
target=str(updated_url),
custom_headers=MappingProxyType({**authorization, "Content-Type": "application/json"}),
custom_llm_provider="laya",
is_streaming_request=False,
)
return TypeAdapter(Response, config=ConfigDict(arbitrary_types_allowed=True)).validate_python(
await endpoint_func(request, fastapi_response, user_api_key_dict)
)
@router.api_route(
"/openrouter/{endpoint:path}",
methods=["GET", "POST", "PUT", "DELETE", "PATCH"],

View file

@ -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()

View file

@ -26,6 +26,7 @@ from fastapi import (
status,
)
from fastapi.responses import StreamingResponse
from pydantic import TypeAdapter
from starlette.datastructures import UploadFile as StarletteUploadFile
from starlette.websockets import WebSocketState
from websockets.asyncio.client import connect
@ -64,6 +65,7 @@ from litellm.llms.base_llm.managed_resources.utils import (
resolve_passthrough_managed_id_provider,
)
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.llms.laya.common_utils import validate_laya_request
from litellm.passthrough import BasePassthroughUtils
from litellm.proxy._types import (
ConfigFieldInfo,
@ -74,7 +76,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 +106,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
@ -585,7 +593,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}
@ -631,10 +650,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)
)
@ -1131,6 +1159,15 @@ async def pass_through_request(
_parsed_body,
)
if not _key_or_team_allows_client_pricing_override(user_api_key_dict):
pricing_body: Final = TypeAdapter(dict[str, object]).validate_python(_parsed_body)
_strip_client_pricing_overrides(pricing_body)
_parsed_body = pricing_body
if custom_llm_provider == "laya":
laya_request: Final = TypeAdapter(Mapping[str, object]).validate_python(_parsed_body)
checkpoint: Final = validate_laya_request(laya_request)
_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
@ -1186,6 +1223,17 @@ async def pass_through_request(
call_type="pass_through_endpoint",
endpoint_type=endpoint_type,
)
if custom_llm_provider == "laya":
hook_body: Final = TypeAdapter(dict[str, object]).validate_python(_parsed_body)
hook_model: Final = hook_body.get("model")
laya_body: Final = MappingProxyType(
{
**hook_body,
"model": hook_model.removeprefix("laya/") if isinstance(hook_model, str) else hook_model,
}
)
_ = validate_laya_request(laya_body)
_parsed_body = TypeAdapter(dict[str, object]).validate_python(laya_body)
resolved_timeout: Final = resolve_pass_through_request_timeout(timeout)
async_client_obj: Final = get_async_httpx_client(
llm_provider=httpxSpecialProvider.PassThroughEndpoint,
@ -2389,7 +2437,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 {}

View file

@ -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,

View file

@ -108,7 +108,7 @@ from .config import (
ComplexityRouterConfig,
ComplexityTier,
CustomDimension,
JevClassifierConfig,
OpenSourceClassifierConfig,
TierDefinition,
)
from .jev_classifier import (
@ -1308,10 +1308,22 @@ class ComplexityRouter(CustomLogger):
"""
@staticmethod
def _build_jev_client(config: JevClassifierConfig) -> JevClassifierClient:
def _build_jev_client(config: OpenSourceClassifierConfig) -> 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'")
raise ValueError(
"opensource_classifier_config.api_key or TYPESAFE_API_KEY is required for classifier_type 'oss_classifier'"
)
api_base: Final = config.api_base or get_secret_str("TYPESAFE_API_BASE") or "https://api.typesafe.ai"
return HttpJevClassifierClient(
api_key=api_key,
@ -1354,12 +1366,12 @@ class ComplexityRouter(CustomLogger):
if default_model:
self.config.default_model = default_model
jev_config: Final = self.config.jev_classifier_config
jev_config: Final = self.config.opensource_classifier_config
self._jev_client: JevClassifierClient | None = (
jev_client
if jev_client is not None
else self._build_jev_client(jev_config)
if self.config.classifier_type == "jev" and jev_config is not None
if self.config.classifier_type == "oss_classifier" and jev_config is not None
else None
)
@ -1459,7 +1471,11 @@ class ComplexityRouter(CustomLogger):
and self.config.classifier_llm_config.circuit_breaker_enabled
)
else jev_config.circuit_breaker_cooldown_seconds
if (self.config.classifier_type == "jev" and jev_config is not None and jev_config.circuit_breaker_enabled)
if (
self.config.classifier_type == "oss_classifier"
and jev_config is not None
and jev_config.circuit_breaker_enabled
)
else None
)
self._classifier_circuit_breaker: _ClassifierCircuitBreaker | None = (
@ -1909,7 +1925,7 @@ class ComplexityRouter(CustomLogger):
return self._classify_with_heuristic_v2(prompt)
if self.config.classifier_type == "custom":
return await self._classify_with_plugin(prompt, system_prompt, request_kwargs, raw_messages)
if self.config.classifier_type == "jev":
if self.config.classifier_type == "oss_classifier":
return await self._jev_classifier_outcome(prompt, system_prompt, request_kwargs, messages)
if self.config.classifier_type in ("heuristic_first", "hybrid") and _encrypted_classifier_task(
request_kwargs, self._reminder_markers_for_request(request_kwargs or EMPTY_MAPPING)
@ -2161,7 +2177,7 @@ class ComplexityRouter(CustomLogger):
request_kwargs: Mapping[str, object] | None,
messages: Sequence[Mapping[str, object]] | None,
) -> ClassificationOutcome:
config: Final = self.config.jev_classifier_config
config: Final = self.config.opensource_classifier_config
client: Final = self._jev_client
if config is None or client is None:
return self._classifier_failure_outcome("jev classifier is not configured", prompt, system_prompt)
@ -2212,12 +2228,14 @@ class ComplexityRouter(CustomLogger):
if not self._tier_pools().get(tier_name):
raise ValueError(f"Jev classifier returned tier {tier_name!r}, which has no models configured")
model: Final = response.model or config.model
accounting_provider: Final = "laya" if config.provider == "laya" else "typesafe"
verdict: Final = JevVerdict(
label=answer.choice,
probabilities=answer.probabilities,
confidence=answer.confidence,
model=model,
cost=jev_classifier_cost(response, config.model),
cost=jev_classifier_cost(response, config.model, accounting_provider),
provider=accounting_provider,
)
if breaker is not None and permit is not None:
breaker.record_success(permit)
@ -2225,8 +2243,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()
@ -4757,7 +4775,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")

View file

@ -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,
@ -674,14 +676,34 @@ class CapabilityClassifierConfig(BaseModel):
return self
class JevClassifierConfig(BaseModel):
def normalize_classifier_config_aliases(config: Mapping[str, object]) -> Mapping[str, object]:
if "jev_classifier_config" in config and "opensource_classifier_config" in config:
return config
normalized: Final = dict(config)
if "jev_classifier_config" in normalized:
normalized["opensource_classifier_config"] = normalized.pop("jev_classifier_config")
if normalized.get("classifier_type") == "jev":
normalized["classifier_type"] = "oss_classifier"
classifier: Final = normalized.get("opensource_classifier_config")
if isinstance(classifier, Mapping):
classifier_fields: Final = TypeAdapter(Mapping[str, object]).validate_python(classifier)
if classifier_fields.get("provider") == "typesafe":
normalized["opensource_classifier_config"] = {
**classifier_fields,
"provider": "jev",
}
return normalized
class OpenSourceClassifierConfig(BaseModel):
model_config = ConfigDict(extra="forbid", frozen=True)
provider: Literal["jev", "laya"] = "jev"
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(
@ -691,30 +713,112 @@ class JevClassifierConfig(BaseModel):
circuit_breaker_enabled: bool = True
circuit_breaker_cooldown_seconds: float = Field(default=30.0, gt=0.0)
@field_validator("provider", mode="before")
@classmethod
def _normalize_provider_alias(cls, value: object) -> object:
return "jev" if value == "typesafe" else value
@field_validator("instructions")
@classmethod
def _reject_blank_instructions(cls, value: str | None) -> str | None:
if value is not None and not value.strip():
raise ValueError("jev_classifier_config.instructions must be non-empty; omit it to use the default")
raise ValueError("opensource_classifier_config.instructions must be non-empty; omit it to use the default")
return value
@field_validator("api_key")
@classmethod
def _reject_blank_api_key(cls, value: str | None) -> str | None:
if value is not None and not value.strip():
raise ValueError("jev_classifier_config.api_key must be non-empty; omit it to use TYPESAFE_API_KEY")
raise ValueError("opensource_classifier_config.api_key must be non-empty; omit it to use TYPESAFE_API_KEY")
return value
@model_validator(mode="after")
def _keep_the_environment_key_on_the_environment_base(self) -> "JevClassifierConfig":
def _keep_the_environment_key_on_the_environment_base(self) -> "OpenSourceClassifierConfig":
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 "
"opensource_classifier_config.api_base requires opensource_classifier_config.api_key: TYPESAFE_API_KEY is only sent "
"to TYPESAFE_API_BASE or https://api.typesafe.ai"
)
return self
JevClassifierConfig = OpenSourceClassifierConfig
@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("opensource_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)
return _resolve_normalized_complexity_router_config_write(
normalize_classifier_config_aliases(incoming),
normalize_classifier_config_aliases(stored) if stored is not None else None,
)
def _resolve_normalized_complexity_router_config_write(
incoming: Mapping[str, object], stored: Mapping[str, object] | None
) -> ComplexityRouterConfigWrite:
if (
stored is None
or incoming.get("classifier_type") != "oss_classifier"
or stored.get("classifier_type") != "oss_classifier"
):
return ComplexityRouterConfigWrite(submitted=incoming, effective=incoming)
incoming_classifier: Final = incoming.get("opensource_classifier_config")
stored_classifier: Final = stored.get("opensource_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", "jev") == existing.get("provider", "jev")
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, "opensource_classifier_config": classifier}),
effective={
**incoming,
"opensource_classifier_config": {
**transport,
**classifier,
},
},
)
MAX_CUSTOM_PATTERN_REPEAT: Final[int] = 64
MAX_CUSTOM_PATTERN_WORK: Final[int] = 2048
MAX_CUSTOM_DIMENSIONS_WORK: Final[int] = 8192
@ -846,6 +950,20 @@ class ContextCompactionConfig(BaseModel):
class ComplexityRouterConfig(BaseModel):
"""Configuration for the ComplexityRouter."""
@model_validator(mode="before")
@classmethod
def _normalize_classifier_aliases(cls, value: object) -> object:
if not isinstance(value, Mapping):
return value
config: Final = TypeAdapter(dict[str, object]).validate_python(value)
if "jev_classifier_config" in config and "opensource_classifier_config" in config:
raise ValueError("Use only opensource_classifier_config; do not also supply jev_classifier_config")
return normalize_classifier_config_aliases(config)
@property
def jev_classifier_config(self) -> OpenSourceClassifierConfig | None:
return self.opensource_classifier_config
# string = pin; list = random pick when adaptive=False, soft-floor home pool when adaptive=True
tiers: dict[str, str | list[str]] = Field(
default_factory=lambda: DEFAULT_TIER_MODELS.copy(),
@ -880,7 +998,7 @@ class ComplexityRouterConfig(BaseModel):
"becomes that tier's rubric bullet; entries named after a built-in tier may omit the "
"description and inherit the built-in criteria. List order is ascending severity and "
"decides which tier wins when several keyword_tier_rules match. Requires classifier_type "
"'llm', 'jev' or 'custom', a fallback_tier, and `tiers` keys matching the defined names exactly. Escalation, "
"'llm', 'oss_classifier' or 'custom', a fallback_tier, and `tiers` keys matching the defined names exactly. Escalation, "
"adaptive selection, session affinity, plugins, tier_labels, and the calibration-example "
"rubric presets are unavailable with a custom tier set: the first four are built on the "
"built-in tier ladder, and the last two rename or exemplify tiers the set replaces."
@ -1024,7 +1142,7 @@ class ComplexityRouterConfig(BaseModel):
"custom",
"heuristic_first",
"hybrid",
"jev",
"oss_classifier",
] = Field(
default="heuristic",
description=(
@ -1032,7 +1150,7 @@ class ComplexityRouterConfig(BaseModel):
"an LLM tier-selection call, a Switchyard-compatible capability forecast, a joint Fuse V2 forecast, "
"a custom classifier plugin, 'heuristic_first', which scores locally and only pays for the LLM classifier when the "
"local scorer does not confidently land a cheap tier, or 'hybrid', which trusts the local scorer "
"everywhere except when its score lands near a tier boundary, or 'jev', a TypeSafe AI Jev structured choice call"
"everywhere except when its score lands near a tier boundary, or 'oss_classifier', a structured choice call using Jev or Laya"
),
)
llm_v2_config: LLMV2Config | None = Field(
@ -1073,7 +1191,7 @@ class ComplexityRouterConfig(BaseModel):
"and otherwise routes to capable_tier"
),
)
jev_classifier_config: JevClassifierConfig | None = None
opensource_classifier_config: OpenSourceClassifierConfig | None = None
heuristic_first_max_tier: str | None = Field(
default=None,
description=(
@ -1639,14 +1757,16 @@ class ComplexityRouterConfig(BaseModel):
return self
@model_validator(mode="after")
def _validate_jev_classifier_config(self) -> "ComplexityRouterConfig":
jev: Final = self.jev_classifier_config
if self.classifier_type != "jev":
def _validate_opensource_classifier_config(self) -> "ComplexityRouterConfig":
jev: Final = self.opensource_classifier_config
if self.classifier_type != "oss_classifier":
if jev is not None:
raise ValueError("jev_classifier_config requires classifier_type 'jev'; otherwise it has no effect")
raise ValueError(
"opensource_classifier_config requires classifier_type 'oss_classifier'; otherwise it has no effect"
)
return self
if jev is None:
raise ValueError("jev_classifier_config is required when classifier_type is 'jev'")
raise ValueError("opensource_classifier_config is required when classifier_type is 'oss_classifier'")
return self
@model_validator(mode="after")
@ -1962,9 +2082,9 @@ class ComplexityRouterConfig(BaseModel):
"enable_non_reasoning_tier cannot be combined with tier_definitions: a custom tier set "
f"replaces the built-in ladder, so name a tier {non_reasoning_key} in tier_definitions instead"
)
if self.classifier_type not in ("llm", "custom", "jev"):
if self.classifier_type not in ("llm", "custom", "oss_classifier"):
raise ValueError(
f"enable_non_reasoning_tier requires classifier_type 'llm', 'jev' or 'custom', got "
f"enable_non_reasoning_tier requires classifier_type 'llm', 'oss_classifier' or 'custom', got "
f"{self.classifier_type!r}: the heuristic scorers only produce the four tiers from SIMPLE up, "
f"so nothing would ever classify as {non_reasoning_key}"
)
@ -1997,7 +2117,7 @@ class ComplexityRouterConfig(BaseModel):
raise ValueError(f"tier_definitions names must be unique (case-insensitive): {', '.join(duplicated)}")
if self.classifier_type in ("heuristic", "heuristic_v2", "capability", "heuristic_first", "hybrid"):
raise ValueError(
"tier_definitions requires classifier_type 'llm', 'jev' or 'custom': the heuristic scorer only "
"tier_definitions requires classifier_type 'llm', 'oss_classifier' or 'custom': the heuristic scorer only "
"produces the built-in tiers from SIMPLE up, as does heuristic_v2"
)
conflicts: Final = self._tier_definition_conflicts()
@ -2164,7 +2284,9 @@ class ComplexityRouterConfig(BaseModel):
)
COMPLEXITY_ROUTER_CONFIG_KEYS: Final[frozenset[str]] = frozenset(ComplexityRouterConfig.model_fields)
COMPLEXITY_ROUTER_CONFIG_KEYS: Final[frozenset[str]] = frozenset(ComplexityRouterConfig.model_fields) | frozenset(
("jev_classifier_config",)
)
"""Every setting name this config owns, derived from the model so a field added later is covered.
These names are disjoint from the OpenAI request params, from ``all_litellm_params``, and from the

View file

@ -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,
@ -90,26 +98,30 @@ class HttpJevClassifierClient:
request_kwargs: Mapping[str, object] | None = None,
) -> JevSystemOneResponse:
start_time: Final = datetime.now(timezone.utc)
authorization: Final[Mapping[str, str]] = (
MappingProxyType({"Authorization": f"Bearer {self._api_key}"}) if self._api_key else MappingProxyType({})
)
response: Final = await self._http_client.post( # pyright: ignore[reportUnknownMemberType] # AsyncHTTPHandler has a dynamic post signature
f"{self._api_base}/v1/systemone",
json=request.model_dump(mode="json"),
headers=MappingProxyType(
{
"Authorization": f"Bearer {self._api_key}",
"Content-Type": "application/json",
}
), # pyright: ignore[reportArgumentType] # HTTP headers are not mutated by AsyncHTTPHandler
headers=MappingProxyType({**authorization, "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 = (
MappingProxyType({**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}],
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={},
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:

View file

@ -19,6 +19,7 @@ from litellm.router_strategy.complexity_router.config import (
COMPLEXITY_ROUTER_CONFIG_KEYS,
DEFAULT_JEV_INSTRUCTIONS,
LLM_CLASSIFIER_TYPES,
normalize_classifier_config_aliases,
)
AUTO_ROUTER_MODEL_PREFIX: Final = "auto_router/"
@ -151,8 +152,11 @@ def strategy_router_dependencies(
)
)
)
complexity: Final = _mapping(litellm_params.get("complexity_router_config"))
complexity: Final = normalize_classifier_config_aliases(_mapping(litellm_params.get("complexity_router_config")))
classifier: Final = _mapping(complexity.get("classifier_llm_config"))
decision_classifier: Final = _mapping(complexity.get("opensource_classifier_config"))
decision_provider: Final = decision_classifier.get("provider", "jev")
accounting_provider: Final = "typesafe" if decision_provider == "jev" else decision_provider
return tuple(
dict.fromkeys(
tuple(dep for tier in _mapping(complexity.get("tiers")).values() for dep in _pool(tier, "tier"))
@ -165,10 +169,10 @@ def strategy_router_dependencies(
)
+ (
_named(
f"typesafe/{_mapping(complexity.get('jev_classifier_config')).get('model', 'jev-latest')}",
f"{accounting_provider}/{decision_classifier.get('model', 'jev-latest')}",
"evaluation",
)
if complexity.get("classifier_type") == "jev"
if complexity.get("classifier_type") == "oss_classifier"
else ()
)
+ (
@ -206,9 +210,9 @@ def defines_custom_classifier_prompt(complexity_router_config: object) -> bool:
Scoped to the classifier types that actually call an LLM, which is also where the config validator
accepts these fields: the heuristic scorers never read them.
"""
config: Final = _mapping(complexity_router_config)
if config.get("classifier_type") == "jev":
instructions: Final = _mapping(config.get("jev_classifier_config")).get("instructions")
config: Final = normalize_classifier_config_aliases(_mapping(complexity_router_config))
if config.get("classifier_type") == "oss_classifier":
instructions: Final = _mapping(config.get("opensource_classifier_config")).get("instructions")
return isinstance(instructions, str) and instructions != DEFAULT_JEV_INSTRUCTIONS
if config.get("classifier_type") not in LLM_CLASSIFIER_TYPES:
return False
@ -272,6 +276,9 @@ _OPERATOR_PROMPT_FIELDS_SQL: Final = " OR ".join(
f"{{config}} ->> '{field}' IS NOT NULL" for field in OPERATOR_CLASSIFIER_PROMPT_FIELDS
)
_DEFAULT_JEV_INSTRUCTIONS_SQL: Final = DEFAULT_JEV_INSTRUCTIONS.replace("'", "''")
_OPENSOURCE_CLASSIFIER_CONFIG_SQL: Final = (
"COALESCE({config} -> 'opensource_classifier_config', {config} -> 'jev_classifier_config')"
)
CUSTOMIZATION_CAPABILITY: Final = GatedAutoRouterCapability(
key="tier_or_classifier_prompt",
@ -286,9 +293,9 @@ CUSTOMIZATION_CAPABILITY: Final = GatedAutoRouterCapability(
f"({{config}} ->> 'classifier_type' IN ({_LLM_CLASSIFIER_TYPES_SQL}) AND ("
"{config} -> 'classifier_llm_config' ->> 'system_prompt' IS NOT NULL OR "
f"{_OPERATOR_PROMPT_FIELDS_SQL})) OR "
"({config} ->> 'classifier_type' = 'jev' AND "
"jsonb_typeof({config} -> 'jev_classifier_config' -> 'instructions') = 'string' AND "
f"{{config}} -> 'jev_classifier_config' ->> 'instructions' <> '{_DEFAULT_JEV_INSTRUCTIONS_SQL}')"
"({config} ->> 'classifier_type' IN ('oss_classifier', 'jev') AND "
f"jsonb_typeof({_OPENSOURCE_CLASSIFIER_CONFIG_SQL} -> 'instructions') = 'string' AND "
f"{_OPENSOURCE_CLASSIFIER_CONFIG_SQL} ->> 'instructions' <> '{_DEFAULT_JEV_INSTRUCTIONS_SQL}')"
),
)

View file

@ -72622,6 +72622,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",

View file

@ -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",
@ -1476,6 +1477,13 @@
"rerank": false
}
},
"laya": {
"display_name": "Laya (`laya`)",
"url": "https://docs.litellm.ai/docs/auto_router/decision_classifiers",
"endpoints": {
"systemone": true
}
},
"lambda_ai": {
"display_name": "Lambda AI (`lambda_ai`)",
"url": "https://docs.litellm.ai/docs/providers/lambda_ai",
@ -3354,6 +3362,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://docs.litellm.ai/docs/auto_router/decision_classifiers"
},
"text_completion": {
"docs_label": "text_completion",
"display_name": "OpenAI Completions API",

View file

@ -63,7 +63,8 @@ def mock_request():
self.method = method
self.request_body = request_body or {}
# Add url attribute that the actual code expects
self.url = "http://localhost:8000/test"
self.url = httpx.URL("http://localhost:8000/test")
self.scope = {"type": "http", "method": method, "path": "/test"}
# Add state attribute that FastAPI requests have
self.state = type("State", (), {})()
@ -414,6 +415,7 @@ PROTOCOL_CONSTRAINED_PASS_THROUGH_ROUTES = {
"/transcribe": {"POST"},
"/transcribe/{operation}": {"POST"},
"/tinyfish/{endpoint:path}": {"GET", "POST"},
"/laya/v1/systemone": {"POST"},
}

View file

@ -0,0 +1 @@

View file

@ -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

View file

@ -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

View file

@ -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
@ -7515,6 +7517,96 @@ class TestTeamMemberAutoRouterWrites:
"model_info": {"id": "allowed-id"},
}])
@staticmethod
def _classifier_config(classifier: Mapping[str, object], legacy: bool) -> Mapping[str, object]:
return {
"classifier_type": "jev" if legacy else "oss_classifier",
"tiers": {"SIMPLE": "allowed"},
"jev_classifier_config" if legacy else "opensource_classifier_config": classifier,
}
@pytest.mark.asyncio
@pytest.mark.parametrize("team_id", [None, "member-team"])
@pytest.mark.parametrize(
"legacy,provider,model",
[(True, "typesafe", "jev-latest"), (False, "jev", "jev-latest"), (True, "laya", "english"), (False, "laya", "english")],
)
async def test_classifier_create_stores_only_canonical_configuration(
self, team_id: str | None, legacy: bool, provider: str, model: str
) -> None:
from litellm.proxy.management_endpoints.model_management_endpoints import add_new_model
row: Final = self._row()
database: Final = self._database(self._team(), row)
classifier: Final = {
"provider": provider, "model": model,
"api_base": "https://decision.test", "api_key": "stored-secret",
}
deployment: Final = Deployment(
model_name="new-classifier-router",
litellm_params=LiteLLM_Params(
model="auto_router/complexity_router",
complexity_router_config=self._classifier_config(classifier, legacy),
),
model_info=ModelInfo(id=row.model_id, team_id=team_id),
)
actor: Final = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
with (
self._environment(database, row),
patch("litellm.proxy.proxy_server.proxy_config.add_deployment", new=AsyncMock(return_value=ReconcileOutcome( # test-quality-ok: [TQ008] model reload I/O boundary
still_desired=frozenset((row.model_id,)), live_after=frozenset((row.model_id,))
))),
patch("litellm.proxy.management_endpoints.model_management_endpoints.append_team_models", new=AsyncMock()), # test-quality-ok: [TQ008] team allowlist persistence boundary
):
await add_new_model(deployment, actor)
written: Final = database.db.litellm_proxymodeltable.create.await_args.kwargs["data"]
saved: Final = json.loads(written["litellm_params"])["complexity_router_config"]
assert saved == {
"classifier_type": "oss_classifier",
"tiers": {"SIMPLE": "allowed"},
"opensource_classifier_config": {**classifier, "provider": "laya" if provider == "laya" else "jev"},
}
@pytest.mark.asyncio
@pytest.mark.parametrize("endpoint", ["create", "patch", "legacy"])
@pytest.mark.parametrize("legacy_config", [None, {"provider": "laya", "model": "english"}])
async def test_ambiguous_classifier_blocks_are_rejected_before_persistence(
self, endpoint: str, legacy_config: Mapping[str, object] | None
) -> None:
from litellm.proxy.management_endpoints.model_management_endpoints import add_new_model
row: Final = self._row()
database: Final = self._database(self._team(), row)
config: Final = {
**self._classifier_config({"provider": "laya", "model": "english"}, False),
"jev_classifier_config": legacy_config,
}
actor: Final = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
request: Final = updateDeployment(
litellm_params=updateLiteLLMParams(complexity_router_config=config), model_info=ModelInfo(id=row.model_id),
)
operation: Final = (
add_new_model(
Deployment(
model_name="ambiguous-classifier-router",
litellm_params=LiteLLM_Params(model="auto_router/complexity_router", complexity_router_config=config),
model_info=ModelInfo(id=row.model_id),
),
actor,
)
if endpoint == "create"
else patch_model(row.model_id, request, actor)
if endpoint == "patch"
else update_model(request, actor)
)
with self._environment(database, row), pytest.raises(ProxyException) as denied:
await operation
assert denied.value.code == "400"
assert "opensource_classifier_config" in denied.value.message
assert "jev_classifier_config" in denied.value.message
database.db.litellm_proxymodeltable.create.assert_not_awaited()
database.db.litellm_proxymodeltable.update.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("endpoint,change", [("patch", "config"), ("legacy", "strategy"), ("patch", "unrelated")])
async def test_admin_router_changes_release_member_scope(self, endpoint: str, change: str) -> None:
@ -7546,15 +7638,16 @@ class TestTeamMemberAutoRouterWrites:
@pytest.mark.asyncio
@pytest.mark.parametrize("endpoint", ["patch", "legacy"])
@pytest.mark.parametrize("stored_legacy,supplied_legacy", [(True, True), (True, False), (False, True), (False, False)])
@pytest.mark.parametrize("change", ["save", "rotate", "move", "move-without-key", "reset", "heuristic"])
async def test_jev_dashboard_save_preserves_server_transport(self, endpoint: str, change: str) -> None:
async def test_jev_dashboard_save_preserves_server_transport(
self, endpoint: str, change: str, stored_legacy: bool, supplied_legacy: bool
) -> None:
original: Final = self._row()
transport: Final = {"api_key": "synthetic-original-jev-key", "api_base": "https://jev.example.com"}
stored_config: Final = {
"classifier_type": "jev",
"tiers": {"SIMPLE": "allowed"},
"jev_classifier_config": {**transport, "instructions": "Old instructions", "timeout_ms": 6100},
}
stored_config: Final = self._classifier_config(
{**transport, "instructions": "Old instructions", "timeout_ms": 6100}, stored_legacy
)
row: Final = original.model_copy(
update={
"litellm_params": {
@ -7572,11 +7665,11 @@ class TestTeamMemberAutoRouterWrites:
"reset": {"api_key": None, "api_base": None},
"heuristic": {},
}[change]
config: Final = {
"tiers": {"SIMPLE": "allowed"},
"classifier_type": "heuristic" if change == "heuristic" else "jev",
**({} if change == "heuristic" else {"jev_classifier_config": {"timeout_ms": 8100, **overrides}}),
}
config: Final = (
{"tiers": {"SIMPLE": "allowed"}, "classifier_type": "heuristic"}
if change == "heuristic"
else self._classifier_config({"timeout_ms": 8100, **overrides}, supplied_legacy)
)
request: Final = updateDeployment(
litellm_params=updateLiteLLMParams(complexity_router_config=config),
model_info=ModelInfo(id=row.model_id),
@ -7597,12 +7690,179 @@ class TestTeamMemberAutoRouterWrites:
expected: Final = (
config
if change == "heuristic"
else {**config, "jev_classifier_config": {**transport, "timeout_ms": 8100, **overrides}}
else {
"classifier_type": "oss_classifier",
"tiers": {"SIMPLE": "allowed"},
"opensource_classifier_config": {**transport, "timeout_ms": 8100, **overrides},
}
)
assert saved == expected
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_legacy,supplied_legacy", [(True, True), (True, False), (False, True), (False, False)])
@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"}, {}),
(
"typesafe", "https://decision.test", {"provider": "jev", "model": "jev-latest"},
{"api_base": "https://decision.test", "api_key": "stored-secret"},
),
(
"jev", "https://decision.test", {"provider": "typesafe", "model": "jev-latest"},
{"api_base": "https://decision.test", "api_key": "stored-secret"},
),
],
)
async def test_decision_provider_changes_cannot_reuse_a_stored_key(
self, endpoint: str, stored_provider: str, stored_base: str | None,
supplied: Mapping[str, object], expected_transport: Mapping[str, object],
stored_legacy: bool, supplied_legacy: bool,
) -> None:
original: Final = self._row()
row: Final = original.model_copy(update={"litellm_params": {
"model": "auto_router/complexity_router",
"complexity_router_config": self._classifier_config(
{
"provider": stored_provider, "model": "english" if stored_provider == "laya" else "jev-latest",
"api_base": stored_base, "api_key": "stored-secret",
},
stored_legacy,
),
}})
database: Final = self._database(self._team(), row)
config: Final = self._classifier_config(supplied, supplied_legacy)
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"]
expected_provider: Final = supplied.get("provider", stored_provider)
assert saved == {
"classifier_type": "oss_classifier",
"tiers": {"SIMPLE": "allowed"},
"opensource_classifier_config": {
**expected_transport, **supplied,
"provider": "jev" if expected_provider == "typesafe" else expected_provider,
},
}
@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_config: Final = json.loads(written["litellm_params"])["complexity_router_config"]
untouched: Final = config_shape in ("omit-config", "null-config")
saved: Final = saved_config["jev_classifier_config" if untouched else "opensource_classifier_config"]
assert saved == secret_params["complexity_router_config"]["jev_classifier_config"]
assert saved_config["classifier_type"] == ("jev" if untouched else "oss_classifier")
if untouched:
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"])

View file

@ -1,3 +1,4 @@
import json
from collections.abc import Mapping
from dataclasses import dataclass
from typing import Final
@ -137,33 +138,48 @@ def test_tier_config_is_normalized_and_unknown_router_extras_are_rejected() -> N
@pytest.mark.parametrize(
("jev_override", "rejected_at"),
[
({"api_base": "https://collector.invalid"}, "jev_classifier_config"),
({"api_base": "https://collector.invalid"}, "opensource_classifier_config"),
({"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"),
({"api_base": "https://collector.invalid", "api_key": ""}, "opensource_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"),
],
)
@pytest.mark.parametrize("legacy", [False, True])
def test_members_cannot_move_the_jev_classifier_off_the_proxys_typesafe_account(
jev_override: Mapping[str, str], rejected_at: str
jev_override: Mapping[str, str], rejected_at: str, legacy: bool
) -> None:
with pytest.raises(HTTPException) as denied:
validate_member_auto_router_config(
{"tiers": {"SIMPLE": "allowed"}, "classifier_type": "jev", "jev_classifier_config": jev_override}
{
"tiers": {"SIMPLE": "allowed"},
"classifier_type": "jev" if legacy else "oss_classifier",
"jev_classifier_config" if legacy else "opensource_classifier_config": jev_override,
}
)
assert denied.value.status_code == 400
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")])
@pytest.mark.parametrize("legacy", [False, True])
def test_members_can_still_tune_the_jev_classifier(provider: str, model: str, legacy: bool) -> None:
validated: Final = validate_member_auto_router_config(
{
"tiers": {"SIMPLE": "allowed"},
"classifier_type": "jev",
"jev_classifier_config": {"model": "jev-preview", "timeout_ms": 500},
"classifier_type": "jev" if legacy else "oss_classifier",
"jev_classifier_config" if legacy else "opensource_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,
) == ("jev" if provider == "typesafe" else provider, model, 500)
assert validate_member_auto_router_config(validated.model_dump()).jev_classifier_config is not None
@ -217,6 +233,92 @@ 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 (
"typesafe" if saved.jev_classifier_config.provider == "jev" else saved.jev_classifier_config.provider
) + f"/{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 +348,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 +367,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 +402,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 +415,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}")

View file

@ -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"]

View file

@ -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
@ -7407,6 +7409,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]:

View file

@ -1470,7 +1470,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({})
@ -1575,7 +1575,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({})
@ -1637,7 +1637,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({})
@ -2507,7 +2507,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({})
@ -4966,7 +4966,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({})
@ -5029,7 +5029,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({})
@ -5081,7 +5081,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({})
@ -5112,7 +5112,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({})
@ -5213,7 +5213,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({})
@ -6650,7 +6650,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 {}
@ -6797,7 +6797,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
@ -6985,36 +6985,82 @@ 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,
)
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",
})
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"
@pytest.mark.parametrize("metadata_slot", ["metadata", "litellm_metadata"])
def test_builtin_passthrough_pins_model_group_to_the_resolved_model(metadata_slot: str) -> None:
request: Final = Request({
"type": "http", "method": "POST", "path": "/gemini/v1beta/models/gemini-2.5-flash:generateContent",
"headers": [], "query_string": b"",
})
kwargs: Final = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint(
request=request, user_api_key_dict=UserAPIKeyAuth(token="hash", user_id="u-1"),
passthrough_logging_payload=MagicMock(), logging_obj=MagicMock(),
_parsed_body={"contents": [], metadata_slot: {"model_group": "unbounded-client-choice"}},
)
assert kwargs["litellm_params"]["metadata"]["model_group"] == "gemini-2.5-flash"
@pytest.mark.parametrize(
@ -7344,7 +7390,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 = {}
@ -7377,7 +7423,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(
@ -7409,7 +7455,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({})
@ -7665,7 +7711,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(

View file

@ -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:

View file

@ -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"

View file

@ -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)
@ -167,8 +169,9 @@ async def test_jev_invalid_usage_never_reaches_spend_callbacks(
@pytest.mark.asyncio
@pytest.mark.parametrize("answer", ["SIMPLE", "UNAVAILABLE", "malformed"])
@pytest.mark.parametrize("private", [False, True])
@pytest.mark.parametrize("legacy", [False, True])
async def test_jev_accounts_once_with_parent_identity_even_when_the_verdict_fails(
monkeypatch: pytest.MonkeyPatch, answer: str, private: bool
monkeypatch: pytest.MonkeyPatch, answer: str, private: bool, legacy: bool
) -> None:
recorder: Final = _UsageRecorder()
monkeypatch.setattr(litellm, "_async_success_callback", [recorder])
@ -196,7 +199,15 @@ async def test_jev_accounts_once_with_parent_identity_even_when_the_verdict_fail
router: Final = ComplexityRouter(
"jev-router",
litellm.Router(model_list=[]),
{"classifier_type": "jev", "jev_classifier_config": {}, "tiers": {"SIMPLE": "cheap"}},
{
"classifier_type": "jev" if legacy else "oss_classifier",
"jev_classifier_config" if legacy else "opensource_classifier_config": {
"provider": "typesafe" if legacy else "jev",
},
"tiers": {"SIMPLE": "cheap"},
"session_affinity": False,
"deployment_affinity": False,
},
jev_client=provider,
derive_savings_baseline=False,
)
@ -209,8 +220,9 @@ async def test_jev_accounts_once_with_parent_identity_even_when_the_verdict_fail
"user_api_key_budget_reservation": {"reservation_id": "parent-reservation"},
"user_api_key_auth": {"budget_reservation": {"reservation_id": "parent-reservation"}},
}
outcome: Final = await router.aclassify(
"private current ask",
result: Final = await router.async_pre_routing_hook(
model="jev-router",
messages=[{"role": "user", "content": "private current ask"}],
request_kwargs={
"metadata": metadata,
"litellm_session_id": "session-a",
@ -221,7 +233,15 @@ async def test_jev_accounts_once_with_parent_identity_even_when_the_verdict_fail
await GLOBAL_LOGGING_WORKER.flush()
await handler.client.aclose()
assert (outcome.cause == "jev_classifier") is (answer == "SIMPLE")
assert result is not None and result.model == "cheap"
assert result.routing_decision is not None
decision: Final = result.routing_decision
assert (decision["cause"] == "jev_classifier") is (answer == "SIMPLE")
if answer == "SIMPLE":
assert decision["classifier_model"] == "typesafe/jev-accounting"
assert decision["classifier_cost"] == pytest.approx(0.007)
assert "jev-classifier:SIMPLE" in decision["signals"]
assert "jev-confidence=1.000000" in decision["signals"]
assert len(recorder.calls) == 1
event: Final = recorder.calls[0]
assert event["response_cost"] == pytest.approx(0.007)
@ -416,10 +436,101 @@ def _answer(choice: str = "SIMPLE") -> JevChoiceAnswer:
def test_jev_config_requires_classifier_config() -> None:
with pytest.raises(ValueError, match="jev_classifier_config is required"):
with pytest.raises(ValueError, match="opensource_classifier_config is required"):
ComplexityRouterConfig.model_validate({"classifier_type": "jev"})
@pytest.mark.parametrize(
("classifier_type", "config_key"),
[
("oss_classifier", "opensource_classifier_config"),
("jev", "jev_classifier_config"),
("oss_classifier", "jev_classifier_config"),
("jev", "opensource_classifier_config"),
],
)
@pytest.mark.parametrize(
("provider", "model", "canonical_provider"),
[(None, "jev-latest", "jev"), ("typesafe", "jev-latest", "jev"), ("jev", "jev-latest", "jev"), ("laya", "english", "laya")],
)
def test_classifier_aliases_load_and_serialize_one_canonical_config(
classifier_type: str, config_key: str, provider: str | None, model: str, canonical_provider: str
) -> None:
incoming: Final = {
"classifier_type": classifier_type,
config_key: {"model": model, "api_key": None, **({"provider": provider} if provider is not None else {})},
}
original: Final = deepcopy(incoming)
config: Final = ComplexityRouterConfig.model_validate(incoming)
assert config.classifier_type == "oss_classifier"
assert config.opensource_classifier_config is not None
assert config.opensource_classifier_config.provider == canonical_provider
assert config.opensource_classifier_config.model == model
assert config.opensource_classifier_config.api_key is None
assert "api_key" in config.opensource_classifier_config.model_fields_set
assert "api_base" not in config.opensource_classifier_config.model_fields_set
assert "jev_classifier_config" not in config.model_dump()
assert config.jev_classifier_config is config.opensource_classifier_config
assert incoming == original
@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])
@pytest.mark.parametrize("legacy", [False, True])
async def test_laya_routes_with_its_own_credentials_and_accounts_the_checkpoint(
monkeypatch: pytest.MonkeyPatch, custom_base: bool, legacy: 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" if legacy else "oss_classifier",
"jev_classifier_config" if legacy else "opensource_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(
@ -437,7 +548,7 @@ def test_jev_instructions_reject_blank_values() -> None:
@pytest.mark.parametrize(
("missing_key", "rejection"),
[
({}, r"api_base requires jev_classifier_config\.api_key"),
({}, r"api_base requires opensource_classifier_config\.api_key"),
({"api_key": ""}, r"api_key must be non-empty"),
({"api_key": " "}, r"api_key must be non-empty"),
],

View file

@ -6,11 +6,11 @@ import pytest
from litellm.router_strategy.complexity_router.fuse_presets import get_fuse_presets
from litellm.router_strategy.complexity_router.jev_classifier import DEFAULT_JEV_INSTRUCTIONS
from litellm.router_utils.auto_router_model_naming import (
carries_complexity_router_settings,
classify_strategy_router_model,
GATED_AUTO_ROUTER_CAPABILITIES,
capability_limit_violation,
carries_complexity_router_settings,
claimed_capability,
classify_strategy_router_model,
count_capability_routers,
gated_capability_of,
strategy_router_dependencies,
@ -23,27 +23,59 @@ 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(
("classifier_type", "config_key"),
[
("jev", "jev_classifier_config"),
("oss_classifier", "opensource_classifier_config"),
("jev", "opensource_classifier_config"),
("oss_classifier", "jev_classifier_config"),
],
)
@pytest.mark.parametrize(
("provider", "model", "accounting_provider"),
[
(None, "jev-latest", "typesafe"),
("typesafe", "jev-preview", "typesafe"),
("jev", "jev-preview", "typesafe"),
("laya", "english", "laya"),
],
)
def test_open_source_classifier_enumerates_its_accounting_model(
classifier_type: str, config_key: str, provider: str | None, model: str, accounting_provider: str
) -> None:
found: Final = strategy_router_dependencies(
{
"model": "auto_router/complexity_router",
"complexity_router_config": {
"classifier_type": "jev",
"jev_classifier_config": {"model": model},
"classifier_type": classifier_type,
config_key: {"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"{accounting_provider}/{model}", "evaluation"),
)
@pytest.mark.parametrize("instructions", [None, DEFAULT_JEV_INSTRUCTIONS, "Route conservatively"])
def test_only_non_default_jev_instructions_claim_the_shared_customization_slot(instructions: str | None) -> None:
capability = claimed_capability({"classifier_type": "jev", "jev_classifier_config": {"instructions": instructions}})
@pytest.mark.parametrize(
("classifier_type", "config_key"),
[
("jev", "jev_classifier_config"),
("oss_classifier", "opensource_classifier_config"),
("jev", "opensource_classifier_config"),
("oss_classifier", "jev_classifier_config"),
],
)
def test_only_non_default_open_source_instructions_claim_the_shared_customization_slot(
instructions: str | None, classifier_type: str, config_key: str
) -> None:
capability: Final = claimed_capability(
{"classifier_type": classifier_type, config_key: {"instructions": instructions}}
)
assert (capability.key if capability else None) == (
"tier_or_classifier_prompt" if instructions == "Route conservatively" else None
)
@ -123,6 +155,21 @@ VALID_TIERS = {
}
@pytest.mark.parametrize("legacy_config", [None, {}, {"provider": "laya", "model": "english"}])
def test_dual_classifier_blocks_return_a_write_validation_error(legacy_config: Mapping[str, object] | None) -> None:
violation: Final = validate_complexity_router_config_write(
{
"tiers": VALID_TIERS,
"classifier_type": "oss_classifier",
"opensource_classifier_config": {"provider": "laya", "model": "english"},
"jev_classifier_config": legacy_config,
}
)
assert violation is not None
assert "opensource_classifier_config" in violation
assert "jev_classifier_config" in violation
@pytest.mark.parametrize(
"keyword_tier_rules,expected_fragment",
[
@ -408,6 +455,8 @@ def test_complexity_embedding_model_is_a_dependency_only_when_semantic_matching_
("token_thresholds", "dimension_weights"),
("reasoning_override_min_score",),
("tiers",),
("jev_classifier_config",),
("opensource_classifier_config",),
],
)
def test_placement_rejects_settings_written_beside_the_config(misplaced):
@ -447,7 +496,7 @@ def test_placement_guards_every_setting_the_config_owns():
ComplexityRouterConfig,
)
assert COMPLEXITY_ROUTER_CONFIG_KEYS == frozenset(ComplexityRouterConfig.model_fields)
assert COMPLEXITY_ROUTER_CONFIG_KEYS == frozenset(ComplexityRouterConfig.model_fields) | {"jev_classifier_config"}
assert {"tier_boundaries", "token_thresholds", "dimension_weights"} <= COMPLEXITY_ROUTER_CONFIG_KEYS

View file

@ -1020,6 +1020,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
"/v1/videos",
"/vertex_ai/live",
"/v1/listen",
"/v1/systemone",
"/v1beta/interactions",
],
},

View file

@ -8830,6 +8830,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;
@ -33152,44 +33169,6 @@ export interface components {
/** Updated By */
updated_by?: string | null;
};
/** JevClassifierConfig */
JevClassifierConfig: {
/**
* Api Base
* @description TypeSafe API base, falling back to TYPESAFE_API_BASE and then https://api.typesafe.ai
*/
api_base?: string | null;
/**
* Api Key
* @description TypeSafe API key, falling back to TYPESAFE_API_KEY
*/
api_key?: string | null;
/**
* Circuit Breaker Cooldown Seconds
* @default 30
*/
circuit_breaker_cooldown_seconds: number;
/**
* Circuit Breaker Enabled
* @default true
*/
circuit_breaker_enabled: boolean;
/**
* Instructions
* @description Replaces the built-in Jev question instructions
*/
instructions?: string | null;
/**
* Model
* @default jev-latest
*/
model: string;
/**
* Timeout Ms
* @default 3000
*/
timeout_ms: number;
};
/** Job */
Job: {
/**
@ -39134,6 +39113,50 @@ export interface components {
*/
type: "openIdConnect";
};
/** OpenSourceClassifierConfig */
OpenSourceClassifierConfig: {
/**
* Api Base
* @description Provider API base; defaults to TYPESAFE_API_BASE or LAYA_API_BASE for the selected provider
*/
api_base?: string | null;
/**
* Api Key
* @description Provider API key; optional for self-hosted Laya
*/
api_key?: string | null;
/**
* Circuit Breaker Cooldown Seconds
* @default 30
*/
circuit_breaker_cooldown_seconds: number;
/**
* Circuit Breaker Enabled
* @default true
*/
circuit_breaker_enabled: boolean;
/**
* Instructions
* @description Replaces the built-in Jev question instructions
*/
instructions?: string | null;
/**
* Model
* @default jev-latest
*/
model: string;
/**
* Provider
* @default jev
* @enum {string}
*/
provider: "jev" | "laya";
/**
* Timeout Ms
* @default 3000
*/
timeout_ms: number;
};
/**
* OperationCreateFile
* @description Instruction describing how to create a file via the apply_patch tool.
@ -41934,11 +41957,11 @@ export interface components {
classifier_plugin_timeout_ms: number;
/**
* Classifier Type
* @description Classification strategy: local regex/keyword scoring, the bundled trained four-tier heuristic, an LLM tier-selection call, a Switchyard-compatible capability forecast, a joint Fuse V2 forecast, a custom classifier plugin, 'heuristic_first', which scores locally and only pays for the LLM classifier when the local scorer does not confidently land a cheap tier, or 'hybrid', which trusts the local scorer everywhere except when its score lands near a tier boundary, or 'jev', a TypeSafe AI Jev structured choice call
* @description Classification strategy: local regex/keyword scoring, the bundled trained four-tier heuristic, an LLM tier-selection call, a Switchyard-compatible capability forecast, a joint Fuse V2 forecast, a custom classifier plugin, 'heuristic_first', which scores locally and only pays for the LLM classifier when the local scorer does not confidently land a cheap tier, or 'hybrid', which trusts the local scorer everywhere except when its score lands near a tier boundary, or 'oss_classifier', a structured choice call using Jev or Laya
* @default heuristic
* @enum {string}
*/
classifier_type: "heuristic" | "heuristic_v2" | "llm" | "capability" | "llm_v2" | "custom" | "heuristic_first" | "hybrid" | "jev";
classifier_type: "heuristic" | "heuristic_v2" | "llm" | "capability" | "llm_v2" | "custom" | "heuristic_first" | "hybrid" | "oss_classifier";
/**
* Code Keywords
* @description Keywords indicating code-related content
@ -42037,7 +42060,6 @@ export interface components {
* @description How close to a tier boundary a heuristic score has to land before the LLM classifier breaks the tie; required when classifier_type is 'hybrid' and rejected otherwise. Everything further than this from every active boundary routes on the scorer's own tier with no classifier call, at any tier, which is what separates 'hybrid' from 'heuristic_first' and its cheap-tier ceiling. A prompt where no dimension fired still goes to the classifier, since the scorer has no opinion to be near a boundary with. 0 escalates only scores sitting exactly on a boundary.
*/
hybrid_boundary_margin?: number | null;
jev_classifier_config?: components["schemas"]["JevClassifierConfig"] | null;
/**
* Keyword Tier Rules
* @description Rules that force a specific tier when their keywords match the prompt
@ -42069,6 +42091,7 @@ export interface components {
* @default false
*/
modality_routing: boolean;
opensource_classifier_config?: components["schemas"]["OpenSourceClassifierConfig"] | null;
/**
* Plan Mode Min Tier
* @description When set, requests carrying a coding-agent plan-mode sentinel (Claude Code plan mode, VS Code Copilot Plan mode, Copilot CLI's exit_plan_mode tool) are routed to at least this tier: the classified tier still wins when it is higher, and the floor also overrides a session-affinity pin to a lower tier for exactly the turns carrying the sentinel, without rewriting the pin -- the first turn after plan mode exits routes as if plan mode had never happened. Names a built-in tier, or with tier_definitions set, one of the defined tier names (list order is ascending severity, same as keyword_tier_rules). Unset disables detection entirely. The sentinels ride in client-injected prompt text, so a caller who pastes one can spend up to this tier's models -- never down, and never outside the configured pools.
@ -42166,7 +42189,7 @@ export interface components {
};
/**
* Tier Definitions
* @description Operator-defined tier set replacing the built-in SIMPLE/MEDIUM/COMPLEX/REASONING. Each entry's name becomes a value the LLM classifier can return and its description becomes that tier's rubric bullet; entries named after a built-in tier may omit the description and inherit the built-in criteria. List order is ascending severity and decides which tier wins when several keyword_tier_rules match. Requires classifier_type 'llm', 'jev' or 'custom', a fallback_tier, and `tiers` keys matching the defined names exactly. Escalation, adaptive selection, session affinity, plugins, tier_labels, and the calibration-example rubric presets are unavailable with a custom tier set: the first four are built on the built-in tier ladder, and the last two rename or exemplify tiers the set replaces.
* @description Operator-defined tier set replacing the built-in SIMPLE/MEDIUM/COMPLEX/REASONING. Each entry's name becomes a value the LLM classifier can return and its description becomes that tier's rubric bullet; entries named after a built-in tier may omit the description and inherit the built-in criteria. List order is ascending severity and decides which tier wins when several keyword_tier_rules match. Requires classifier_type 'llm', 'oss_classifier' or 'custom', a fallback_tier, and `tiers` keys matching the defined names exactly. Escalation, adaptive selection, session affinity, plugins, tier_labels, and the calibration-example rubric presets are unavailable with a custom tier set: the first four are built on the built-in tier ladder, and the last two rename or exemplify tiers the set replaces.
*/
tier_definitions?: components["schemas"]["TierDefinition"][] | null;
/**
@ -61639,6 +61662,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;