This commit is contained in:
tin-berri 2026-09-30 10:36:40 -04:00 • committed by GitHub
commit ae68a1f2b7
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
34 changed files with 1292 additions and 153 deletions

68
cookbook/laya/README.md Normal file
View file

@ -0,0 +1,68 @@
# Laya gateway and classifier
[Laya](https://github.com/NandhaKishorM/laya) serves typed decisions using the System One protocol. LiteLLM exposes its native API at `POST /laya/v1/systemone` and offers Laya under the auto-router's Decision Model classifier
## Start Laya
Install the server in a separate environment from LiteLLM, then start it on a port reachable from the gateway
```bash
python -m pip install 'laya[serve]==0.3.21'
LAYA_HOST=127.0.0.1 LAYA_PORT=8000 LAYA_MODELS=english laya-serve
```
Set `LAYA_API_BASE=http://127.0.0.1:8000` in the gateway environment. If the Laya server requires authentication, set the same `LAYA_API_KEY` on both processes. The server URL must use HTTP or HTTPS and contain no embedded credentials, query, or fragment
## Call the native API
Use a LiteLLM virtual key with access to `laya/english`. The request names the bare checkpoint; LiteLLM uses `laya/english` for authorization and usage logs
```bash
curl http://localhost:4000/laya/v1/systemone \
-H "Authorization: Bearer $LITELLM_API_KEY" \
-H 'Content-Type: application/json' \
-d '{
"model": "english",
"state": "My invoice has two identical charges",
"questions": {
"department": {
"type": "choice",
"instructions": "Choose the department that should help",
"criteria": {
"billing": "Invoices, payments, and refunds",
"technical": "Bugs and connectivity problems"
}
}
}
}'
```
The response preserves Laya's `answers`, `usage`, and `routing` fields. Every request must explicitly choose `english`, `multilingual`, or `typed-decisions`; unknown names and automatic selection are rejected so model permissions match the checkpoint being called. This integration does not expose `/v1/evaluate` or chat completions
## Use Laya as a classifier
Add this router alongside your existing `small-solver` and `large-solver` deployments
```yaml
model_list:
- model_name: laya-router
litellm_params:
model: auto_router/complexity_router
complexity_router_config:
classifier_type: jev
jev_classifier_config:
provider: laya
model: english
timeout_ms: 15000
tiers:
SIMPLE: small-solver
MEDIUM: small-solver
COMPLEX: large-solver
REASONING: large-solver
```
In the dashboard, choose Decision Model, then Laya and a checkpoint. `classifier_type: jev` remains the stored type for both Jev and Laya
Omitting `api_base` uses the gateway's `LAYA_API_BASE` and optional `LAYA_API_KEY`. Administrators may instead set `api_base` and optional `api_key` in `jev_classifier_config`; an explicit base without a key is keyless and does not inherit an environment key. Switching providers or changing the saved server URL discards the previous hidden key. Team members can select a centrally configured Laya checkpoint but cannot set connection credentials or URLs
Laya's catalog token prices are zero because self-hosting has no provider API fee. Hosting costs are separate; operators may register their own token rates. Logs use the selected checkpoint, such as `laya/english`, and record the usage returned by Laya. Evaluate classification accuracy on your own prompts before enabling routing

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

View file

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

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

View file

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

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

View file

@ -1883,6 +1883,14 @@ def _extract_model_candidates_from_request(
llm_router: Router | None = None,
team_id: str | None = None,
) -> list[str]:
if route.rstrip("/") == "/laya/v1/systemone":
from litellm.llms.laya.common_utils import validate_laya_model
try:
laya_model: Final = validate_laya_model(request_data.get("model"))
except ValueError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
return [f"laya/{laya_model}"]
if route == "/cost/predict-cache":
prediction_models: Final = _cache_prediction_model_candidates(request_data, llm_router, team_id) # pyright: ignore[reportUnknownArgumentType] # the typed reader validates each deployment ID from this legacy payload
return _dedupe_model_candidates(prediction_models)

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,
@ -114,6 +115,7 @@ from litellm.router_strategy.complexity_router import (
normalize_classification_examples,
normalize_classification_prompt,
)
from litellm.router_strategy.complexity_router.config import resolve_complexity_router_config_write
from litellm.router_utils.auto_router_model_naming import (
GATED_AUTO_ROUTER_CAPABILITIES,
STRATEGY_ROUTER_PARAM_FIELDS,
@ -186,6 +188,25 @@ class _ProxyModelRow(Protocol):
def model_dump_json(self, *, exclude_none: bool = False) -> str: ...
def _model_write_response(
row: _ProxyModelRow, member_write: MemberAutoRouterWrite | None
) -> _ProxyModelRow | Mapping[str, object]:
if member_write is None:
return row
payload: Final = TypeAdapter(dict[str, object]).validate_json(row.model_dump_json())
stored_params: Final = payload.get("litellm_params")
params: Final = (
TypeAdapter(dict[str, object]).validate_json(stored_params)
if isinstance(stored_params, str)
else TypeAdapter(dict[str, object]).validate_python(stored_params)
)
redacted: Final = redact_credentials_in_payload(params)
return {
**payload,
"litellm_params": json.dumps(redacted) if isinstance(stored_params, str) else redacted,
}
class _ProxyModelTable(Protocol):
def find_unique(self, *, where: Mapping[str, object]) -> Awaitable[BaseModel | None]: ...
@ -406,34 +427,10 @@ WHERE model_id <> $1
def _effective_complexity_router_config(
incoming_params: GenericLiteLLMParams | None, existing_params: GenericLiteLLMParams | None
) -> object:
) -> Mapping[str, object] | None:
incoming: Final = None if incoming_params is None else incoming_params.complexity_router_config
existing: Final = None if existing_params is None else existing_params.complexity_router_config
if incoming is None:
return existing
if existing is None or incoming.get("classifier_type") != "jev" or existing.get("classifier_type") != "jev":
return incoming
incoming_jev: Final[object] = incoming.get("jev_classifier_config")
existing_jev: Final[object] = existing.get("jev_classifier_config")
if not isinstance(incoming_jev, Mapping) or not isinstance(existing_jev, Mapping):
return incoming
supplied: Final = TypeAdapter(dict[str, object]).validate_python(incoming_jev)
stored: Final = TypeAdapter(dict[str, object]).validate_python(existing_jev)
same_base: Final = "api_base" not in supplied or supplied["api_base"] == stored.get("api_base")
transport: Final = MappingProxyType(
{
key: value
for key, value in stored.items()
if key in ("api_key", "api_base") and (key != "api_key" or same_base)
}
)
return { # mutable-ok: persisted JSON requires concrete nested dicts
**incoming,
"jev_classifier_config": { # mutable-ok: json.dumps cannot serialize MappingProxyType
**transport,
**supplied,
},
}
return resolve_complexity_router_config_write(incoming, existing).effective
def _effective_model(
@ -1303,7 +1300,7 @@ async def patch_model(
live_after=reload_outcome.live_after,
)
return updated_model
return _model_write_response(updated_model, member_write)
except Exception as e:
verbose_proxy_logger.exception("Error in patch_model: %s", e)
@ -2538,7 +2535,7 @@ async def add_new_model(
live_after=reload_outcome.live_after,
)
return model_response
return _model_write_response(model_response, member_write)
except Exception as e:
verbose_proxy_logger.exception("litellm.proxy.proxy_server.add_new_model(): Exception occured - %s", e)
@ -2764,7 +2761,7 @@ async def update_model(
live_after=reload_outcome.live_after,
)
return model_response
return _model_write_response(model_response, member_write)
except Exception as e:
verbose_proxy_logger.exception("litellm.proxy.proxy_server.update_model(): Exception occured - %s", e)
if isinstance(e, HTTPException):

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
@ -71,6 +75,7 @@ class _MemberJevClassifierConfig(BaseModel):
model_config = ConfigDict(extra="forbid")
provider: Literal["typesafe", "laya"] = "typesafe"
model: str
api_key: None = None
api_base: None = None
@ -123,14 +128,21 @@ def authorize_member_auto_router_team(
def validate_member_auto_router_config(config: Mapping[str, object]) -> RequestComplexityRouterConfig:
return _validate_member_auto_router_config_write(resolve_complexity_router_config_write(config, None))
def _validate_member_auto_router_config_write(write: ComplexityRouterConfigWrite) -> RequestComplexityRouterConfig:
if write.effective is None:
raise HTTPException(status_code=400, detail="A complexity_router_config is required.")
try:
validated: Final = _MemberComplexityRouterConfig.model_validate(config)
for entries in validated.tier_model_configs.values():
for entry in entries:
_MemberRouterGenerationParams.model_validate(entry.litellm_params)
if validated.jev_classifier_config is not None:
_MemberJevClassifierConfig.model_validate(validated.jev_classifier_config.model_dump())
return validated
if write.submitted is not None:
validated: Final = _MemberComplexityRouterConfig.model_validate(write.submitted)
for entries in validated.tier_model_configs.values():
for entry in entries:
_MemberRouterGenerationParams.model_validate(entry.litellm_params)
if validated.jev_classifier_config is not None:
_MemberJevClassifierConfig.model_validate(validated.jev_classifier_config.model_dump())
return RequestComplexityRouterConfig.model_validate(write.effective)
except ValidationError as exc:
location: Final = ".".join(str(part) for part in exc.errors()[0]["loc"])
raise HTTPException(status_code=400, detail=f"Invalid member auto-router configuration at {location}.") from exc
@ -334,16 +346,15 @@ async def authorize_member_auto_router_write(
if existing is not None and incoming.model_name not in (None, public_name, existing.model_name):
raise HTTPException(status_code=403, detail="Team members cannot rename an auto router.")
supplied_config: Final = _RouterConfigSource.model_validate(params.model_dump()).complexity_router_config
raw_config: Final = (
supplied_config
if supplied_config is not None
else _RouterConfigSource.model_validate(existing.litellm_params.model_dump()).complexity_router_config
stored_config: Final = (
_RouterConfigSource.model_validate(existing.litellm_params.model_dump()).complexity_router_config
if existing is not None
else None
)
if raw_config is None:
raise HTTPException(status_code=400, detail="A complexity_router_config is required.")
config: Final = validate_member_auto_router_config(raw_config)
resolved_config: Final = resolve_complexity_router_config_write(supplied_config, stored_config)
if resolved_config.supplied_connection_fields:
raise HTTPException(status_code=403, detail="Team members cannot change classifier connections.")
config: Final = _validate_member_auto_router_config_write(resolved_config)
stored_default: Final = existing.litellm_params.complexity_router_default_model if existing is not None else None
default_model: Final = (
params.complexity_router_default_model

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 TypeAdapter
from starlette.websockets import WebSocketState
from typing_extensions import ReadOnly, TypedDict
@ -54,6 +55,7 @@ from litellm.llms.deepgram.common_utils import (
deepgram_listen_websocket_target,
)
from litellm.llms.fal_ai.cost_calculator import fal_ai_passthrough_cost, fal_ai_queue_base
from litellm.llms.laya.common_utils import laya_connection, validate_laya_request
from litellm.llms.nvidia_nim.passthrough.transformation import nvidia_nim_model_group_in_path
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.passthrough.main import AsyncPassthroughStreamingResponse
@ -633,6 +635,43 @@ async def typesafe_proxy_route(
return await endpoint_func(request, fastapi_response, user_api_key_dict)
@router.post(
"/laya/v1/systemone",
tags=["Laya Pass-through", "pass-through"], # mutable-ok: FastAPI route metadata requires a list
)
async def laya_proxy_route(
request: Request,
fastapi_response: Response,
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
) -> Response:
body: Final = TypeAdapter(dict[str, object]).validate_python(await _read_request_body(request))
try:
_ = validate_laya_request(body)
except ValueError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
try:
connection: Final = laya_connection()
except ValueError as exc:
raise HTTPException(
status_code=503, detail="Laya server is not configured correctly; check LAYA_API_BASE"
) from exc
base_url: Final = httpx.URL(connection.api_base)
updated_url: Final = base_url.copy_with(
path=HttpPassThroughEndpointHelpers.join_base_and_endpoint_path(base_url, "/v1/systemone"),
)
endpoint_func: Final = create_pass_through_route(
endpoint="v1/systemone",
target=str(updated_url),
custom_headers={ # mutable-ok: pass-through request headers require a mutable mapping
**({"Authorization": f"Bearer {connection.api_key}"} if connection.api_key else {}),
"Content-Type": "application/json",
},
custom_llm_provider="laya",
is_streaming_request=False,
)
return cast(Response, await endpoint_func(request, fastapi_response, user_api_key_dict))
@router.api_route(
"/openrouter/{endpoint:path}",
methods=["GET", "POST", "PUT", "DELETE", "PATCH"], # mutable-ok: FastAPI route metadata requires a list

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

@ -74,7 +74,11 @@ from litellm.proxy._types import (
ProxyException,
UserAPIKeyAuth,
)
from litellm.proxy.auth.auth_utils import request_dispatched_to_pass_through_endpoint
from litellm.proxy.auth.auth_utils import (
get_model_from_request,
get_request_route,
request_dispatched_to_pass_through_endpoint,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_request_processing import (
ProxyBaseLLMRequestProcessing,
@ -100,6 +104,8 @@ from litellm.proxy.common_utils.sse_keepalive import (
from litellm.proxy.litellm_pre_call_utils import (
LiteLLMProxyRequestSetup,
_get_dynamic_logging_metadata, # pyright: ignore[reportPrivateUsage] # shared proxy helper, same import style as _read_request_body above
_key_or_team_allows_client_pricing_override, # pyright: ignore[reportPrivateUsage] # reuse the proxy's pricing trust policy
_strip_client_pricing_overrides, # pyright: ignore[reportPrivateUsage] # sanitize before trusted hooks add guardrail costs
)
from litellm.proxy.route_llm_request import ProxyModelNotFoundError
from litellm.proxy.utils import normalize_route_for_root_path
@ -583,7 +589,18 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
"""
Filter out litellm params from the request body
"""
from litellm.proxy.proxy_server import llm_router
_parsed_body = _parsed_body or {}
managed_model: Final = get_model_from_request(
request_data=_parsed_body,
route=get_request_route(request),
request_headers=request.headers,
request_query_params=request.query_params,
llm_router=llm_router,
request=request,
team_id=user_api_key_dict.team_id,
)
litellm_keys_in_body: Final = MappingProxyType(
{k: _parsed_body.pop(k) for k in types_utils.all_litellm_params if k in _parsed_body}
@ -629,10 +646,19 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
# would attribute it to a budget the operator scoped to a LiteLLM model that
# merely shares the name.
if not request_dispatched_to_pass_through_endpoint(request):
_metadata["model_group"] = managed_model if isinstance(managed_model, str) else None
_metadata["user_api_key_model_max_budget"] = user_api_key_dict.model_max_budget
_metadata["user_api_key_team_model_max_budget"] = user_api_key_dict.team_model_max_budget
_metadata["user_api_key_user_model_max_budget"] = user_api_key_dict.user_model_max_budget
_metadata["user_api_key_end_user_model_max_budget"] = user_api_key_dict.end_user_model_max_budget
else:
for field in (
"user_api_key_model_max_budget",
"user_api_key_team_model_max_budget",
"user_api_key_user_model_max_budget",
"user_api_key_end_user_model_max_budget",
):
_metadata.pop(field, None)
_metadata.update(
LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict=user_api_key_dict)
)
@ -1127,6 +1153,14 @@ async def pass_through_request(
_parsed_body,
)
if not _key_or_team_allows_client_pricing_override(user_api_key_dict):
_strip_client_pricing_overrides(_parsed_body)
if custom_llm_provider == "laya":
from litellm.llms.laya.common_utils import validate_laya_request
checkpoint: Final = validate_laya_request(_parsed_body)
_parsed_body["model"] = f"laya/{checkpoint}"
### COLLECT GUARDRAILS FOR PASSTHROUGH ENDPOINT ###
# Passthrough endpoints are opt-in only for guardrails
# When enabled, collect guardrails from org/team/key levels + passthrough-specific
@ -1181,6 +1215,13 @@ async def pass_through_request(
data=_parsed_body,
call_type="pass_through_endpoint",
)
if custom_llm_provider == "laya":
hook_model: Final = _parsed_body.get("model")
laya_body: Final = {
**_parsed_body,
"model": hook_model.removeprefix("laya/") if isinstance(hook_model, str) else hook_model,
}
_parsed_body = {**laya_body, "model": validate_laya_request(laya_body)}
resolved_timeout: Final = resolve_pass_through_request_timeout(timeout)
async_client_obj: Final = get_async_httpx_client(
llm_provider=httpxSpecialProvider.PassThroughEndpoint,
@ -2384,7 +2425,9 @@ async def websocket_passthrough_request(
# with the existing _init_kwargs_for_pass_through_endpoint function
class DummyRequest:
def __init__(self, url: str, method: str = "WEBSOCKET", headers: dict | None = None):
self.url = url
self.url = httpx.URL(url)
self.scope = websocket.scope
self.query_params = websocket.query_params
self.method = method
self.headers = headers or {}

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

@ -1309,6 +1309,16 @@ class ComplexityRouter(CustomLogger):
@staticmethod
def _build_jev_client(config: JevClassifierConfig) -> JevClassifierClient:
if config.provider == "laya":
from litellm.llms.laya.common_utils import laya_connection
connection: Final = laya_connection(config.api_base, config.api_key)
return HttpJevClassifierClient(
api_key=connection.api_key,
api_base=connection.api_base,
http_client=get_async_httpx_client(httpxSpecialProvider.PassThroughEndpoint),
provider="laya",
)
api_key: Final = config.api_key or get_secret_str("TYPESAFE_API_KEY")
if not api_key:
raise ValueError("jev_classifier_config.api_key or TYPESAFE_API_KEY is required for classifier_type 'jev'")
@ -2217,7 +2227,8 @@ class ComplexityRouter(CustomLogger):
probabilities=answer.probabilities,
confidence=answer.confidence,
model=model,
cost=jev_classifier_cost(response, config.model),
cost=jev_classifier_cost(response, config.model, config.provider),
provider=config.provider,
)
if breaker is not None and permit is not None:
breaker.record_success(permit)
@ -2225,8 +2236,8 @@ class ComplexityRouter(CustomLogger):
tier=tier,
score=None,
signals=(
f"jev-classifier:{tier_name}",
f"jev-confidence={answer.confidence:.6f}",
f"{'laya' if config.provider == 'laya' else 'jev'}-classifier:{tier_name}",
f"{'laya' if config.provider == 'laya' else 'jev'}-confidence={answer.confidence:.6f}",
*(
f"tier-probability:{label}={probability:.6f}"
for label, probability in answer.probabilities.items()
@ -4765,7 +4776,7 @@ class ComplexityRouter(CustomLogger):
tier_litellm_params: Final = self._litellm_params_for_model(tier, routed_model)
classifier_model: Final = (
f"typesafe/{outcome.jev_verdict.model}"
f"{outcome.jev_verdict.provider}/{outcome.jev_verdict.model}"
if outcome.cause == "jev_classifier" and outcome.jev_verdict is not None
else self.config.classifier_llm_config.model
if outcome.cause in ("llm_classifier", "capability_classifier", "llm_v2_classifier", "llm_v2_fallback")

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,
@ -681,11 +683,12 @@ class CapabilityClassifierConfig(BaseModel):
class JevClassifierConfig(BaseModel):
model_config = ConfigDict(extra="forbid", frozen=True)
provider: Literal["typesafe", "laya"] = "typesafe"
model: str = "jev-latest"
api_key: str | None = Field(default=None, description="TypeSafe API key, falling back to TYPESAFE_API_KEY")
api_key: str | None = Field(default=None, description="Provider API key; optional for self-hosted Laya")
api_base: str | None = Field(
default=None,
description="TypeSafe API base, falling back to TYPESAFE_API_BASE and then https://api.typesafe.ai",
description="Provider API base; defaults to TYPESAFE_API_BASE or LAYA_API_BASE for the selected provider",
)
timeout_ms: int = Field(default=3000, ge=1)
instructions: str | None = Field(
@ -711,6 +714,13 @@ class JevClassifierConfig(BaseModel):
@model_validator(mode="after")
def _keep_the_environment_key_on_the_environment_base(self) -> "JevClassifierConfig":
if self.provider == "laya":
from litellm.llms.laya.common_utils import validate_laya_api_base, validate_laya_model
_ = validate_laya_model(self.model)
if self.api_base is not None:
_ = validate_laya_api_base(self.api_base)
return self
if self.api_base is not None and self.api_key is None:
raise ValueError(
"jev_classifier_config.api_base requires jev_classifier_config.api_key: TYPESAFE_API_KEY is only sent "
@ -719,6 +729,60 @@ class JevClassifierConfig(BaseModel):
return self
@dataclass(frozen=True, slots=True)
class ComplexityRouterConfigWrite:
submitted: Mapping[str, object] | None
effective: Mapping[str, object] | None
@property
def supplied_connection_fields(self) -> frozenset[str]:
classifier: Final = self.submitted.get("jev_classifier_config") if self.submitted is not None else None
return frozenset(
field for field in ("api_base", "api_key") if isinstance(classifier, Mapping) and field in classifier
)
def resolve_complexity_router_config_write(
incoming: Mapping[str, object] | None, stored: Mapping[str, object] | None
) -> ComplexityRouterConfigWrite:
if incoming is None:
return ComplexityRouterConfigWrite(submitted=None, effective=stored)
if stored is None or incoming.get("classifier_type") != "jev" or stored.get("classifier_type") != "jev":
return ComplexityRouterConfigWrite(submitted=incoming, effective=incoming)
incoming_classifier: Final = incoming.get("jev_classifier_config")
stored_classifier: Final = stored.get("jev_classifier_config")
if not isinstance(incoming_classifier, Mapping) or not isinstance(stored_classifier, Mapping):
return ComplexityRouterConfigWrite(submitted=incoming, effective=incoming)
existing: Final = TypeAdapter(dict[str, object]).validate_python(stored_classifier)
supplied: Final = TypeAdapter(dict[str, object]).validate_python(incoming_classifier)
classifier: Final = (
MappingProxyType({**supplied, "provider": existing["provider"]})
if "provider" not in supplied and "provider" in existing
else supplied
)
same_provider: Final = classifier.get("provider", "typesafe") == existing.get("provider", "typesafe")
same_base: Final = "api_base" not in classifier or (
classifier["api_base"] is not None and classifier["api_base"] == existing.get("api_base")
)
transport: Final = MappingProxyType(
{
key: value
for key, value in existing.items()
if same_provider and key in ("api_key", "api_base") and (key != "api_key" or same_base)
}
)
return ComplexityRouterConfigWrite(
submitted=MappingProxyType({**incoming, "jev_classifier_config": classifier}),
effective={ # mutable-ok: persisted JSON requires concrete nested dicts
**incoming,
"jev_classifier_config": { # mutable-ok: json.dumps cannot serialize MappingProxyType
**transport,
**classifier,
},
},
)
MAX_CUSTOM_PATTERN_REPEAT: Final[int] = 64
MAX_CUSTOM_PATTERN_WORK: Final[int] = 2048
MAX_CUSTOM_DIMENSIONS_WORK: Final[int] = 8192

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,
@ -95,21 +103,25 @@ class HttpJevClassifierClient:
json=request.model_dump(mode="json"),
headers=MappingProxyType(
{
"Authorization": f"Bearer {self._api_key}",
**({"Authorization": f"Bearer {self._api_key}"} if self._api_key else {}),
"Content-Type": "application/json",
}
), # pyright: ignore[reportArgumentType] # HTTP headers are not mutated by AsyncHTTPHandler
timeout=timeout_s,
)
response.raise_for_status()
body: Final = TypeAdapter(dict[str, object]).validate_json(response.content)
normalized_body: Final = (
{**body, "model": laya_response_model(body, request.model)} if self._provider == "laya" else body
)
try:
self._log_response(request, response, request_kwargs, start_time)
except Exception as exc: # noqa: BLE001 # logging integrations must not discard a provider verdict
verbose_router_logger.warning("JEV response logging failed (%s)", type(exc).__name__)
return TypeAdapter(JevSystemOneResponse).validate_python(response.json())
return TypeAdapter(JevSystemOneResponse).validate_python(normalized_body)
@staticmethod
def _log_response(
self,
request: JevSystemOneRequest,
response: httpx.Response,
request_kwargs: Mapping[str, object] | None,
@ -139,7 +151,7 @@ class HttpJevClassifierClient:
"turn_off_message_logging": effective_turn_off_message_logging(request_kwargs),
}
logging_obj: Final = Logging(
model=f"typesafe/{request.model}",
model=f"{self._provider}/{request.model}",
messages=[{"role": "user", "content": request.state}], # mutable-ok: callbacks require JSON message lists
stream=False,
call_type="pass_through_endpoint",
@ -150,7 +162,7 @@ class HttpJevClassifierClient:
kwargs=params,
)
logging_obj.update_environment_variables(
model=f"typesafe/{request.model}",
model=f"{self._provider}/{request.model}",
user=parent_user if isinstance(parent_user := parent.get("user"), str) else None,
optional_params={}, # mutable-ok: Logging's optional_params contract requires a dict
litellm_params=params,
@ -165,7 +177,7 @@ class HttpJevClassifierClient:
end_time=end_time,
cache_hit=False,
request_body=MappingProxyType({"model": request.model}),
custom_llm_provider="typesafe",
custom_llm_provider=self._provider,
litellm_params=params,
)
success_handlers: Final = logging_obj.dispatch_success_handlers(
@ -189,6 +201,7 @@ class JevVerdict(NamedTuple):
confidence: float
model: str
cost: float | None
provider: Literal["typesafe", "laya"] = "typesafe"
class _RegistryPricing(BaseModel):
@ -211,12 +224,14 @@ def build_jev_request(
return JevSystemOneRequest(state=state, model=model, questions=MappingProxyType({"tier": question}))
def jev_classifier_cost(response: JevSystemOneResponse, configured_model: str) -> float | None:
def jev_classifier_cost(
response: JevSystemOneResponse, configured_model: str, provider: Literal["typesafe", "laya"] = "typesafe"
) -> float | None:
usage: Final = response.usage
if usage is None:
return None
model: Final = response.model or configured_model
model_key: Final = f"typesafe/{model}"
model_key: Final = f"{provider}/{model}"
if model_key not in litellm.model_cost: # pyright: ignore[reportUnknownMemberType] # registry is dynamically typed
return None
try:

View file

@ -153,6 +153,7 @@ def strategy_router_dependencies(
)
complexity: Final = _mapping(litellm_params.get("complexity_router_config"))
classifier: Final = _mapping(complexity.get("classifier_llm_config"))
decision_classifier: Final = _mapping(complexity.get("jev_classifier_config"))
return tuple(
dict.fromkeys(
tuple(dep for tier in _mapping(complexity.get("tiers")).values() for dep in _pool(tier, "tier"))
@ -165,7 +166,7 @@ def strategy_router_dependencies(
)
+ (
_named(
f"typesafe/{_mapping(complexity.get('jev_classifier_config')).get('model', 'jev-latest')}",
f"{decision_classifier.get('provider', 'typesafe')}/{decision_classifier.get('model', 'jev-latest')}",
"evaluation",
)
if complexity.get("classifier_type") == "jev"

View file

@ -72497,6 +72497,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",
@ -1458,6 +1459,13 @@
"rerank": false
}
},
"laya": {
"display_name": "Laya (`laya`)",
"url": "https://github.com/BerriAI/litellm/tree/main/cookbook/laya",
"endpoints": {
"systemone": true
}
},
"lambda_ai": {
"display_name": "Lambda AI (`lambda_ai`)",
"url": "https://docs.litellm.ai/docs/providers/lambda_ai",
@ -3336,6 +3344,13 @@
"provider_json_field": "skills",
"url": "https://docs.litellm.ai/docs/skills"
},
"systemone": {
"docs_label": "systemone",
"display_name": "System One Decision API",
"leftnav_label": "/laya/v1/systemone",
"provider_json_field": "systemone",
"url": "https://github.com/BerriAI/litellm/tree/main/cookbook/laya"
},
"text_completion": {
"docs_label": "text_completion",
"display_name": "OpenAI Completions API",

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
@ -7603,6 +7605,148 @@ class TestTeamMemberAutoRouterWrites:
assert row.litellm_params["complexity_router_config"] == stored_config
assert request.litellm_params.complexity_router_config == config
@pytest.mark.asyncio
@pytest.mark.parametrize("endpoint", ["patch", "legacy"])
@pytest.mark.parametrize(
"stored_provider,stored_base,supplied,expected_transport",
[
("laya", "https://decision.test", {"provider": "laya", "model": "english"}, {"api_base": "https://decision.test", "api_key": "stored-secret"}),
("laya", "https://decision.test", {"provider": "laya", "model": "english", "api_base": "https://decision.test"}, {"api_base": "https://decision.test", "api_key": "stored-secret"}),
(
"laya",
"https://decision.test",
{"provider": "laya", "model": "english", "api_key": None},
{"api_base": "https://decision.test"},
),
("laya", "https://decision.test", {"provider": "laya", "model": "english", "api_base": "https://new.test"}, {}),
("laya", "https://decision.test", {"provider": "laya", "model": "english", "api_base": None}, {}),
("laya", None, {"provider": "laya", "model": "english", "api_base": None}, {}),
("laya", "https://decision.test", {"provider": "typesafe", "model": "jev-latest"}, {}),
(
"laya", "https://decision.test", {"model": "english", "timeout_ms": 8100},
{"provider": "laya", "api_base": "https://decision.test", "api_key": "stored-secret"},
),
("typesafe", "https://decision.test", {"provider": "laya", "model": "english"}, {}),
],
)
async def test_decision_provider_changes_cannot_reuse_a_stored_key(
self, endpoint: str, stored_provider: str, stored_base: str | None,
supplied: dict[str, object], expected_transport: dict[str, object]
) -> None:
original: Final = self._row()
row: Final = original.model_copy(update={"litellm_params": {
"model": "auto_router/complexity_router",
"complexity_router_config": {
"classifier_type": "jev", "tiers": {"SIMPLE": "allowed"},
"jev_classifier_config": {
"provider": stored_provider, "model": "english" if stored_provider == "laya" else "jev-latest",
"api_base": stored_base, "api_key": "stored-secret",
},
},
}})
database: Final = self._database(self._team(), row)
config: Final = {"classifier_type": "jev", "tiers": {"SIMPLE": "allowed"}, "jev_classifier_config": supplied}
request: Final = updateDeployment(
litellm_params=updateLiteLLMParams(complexity_router_config=config), model_info=ModelInfo(id=row.model_id),
)
actor: Final = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
with self._environment(database, row):
await (patch_model(row.model_id, request, actor) if endpoint == "patch" else update_model(request, actor))
written: Final = database.db.litellm_proxymodeltable.update.await_args.kwargs["data"]
saved: Final = json.loads(written["litellm_params"])["complexity_router_config"]
assert saved == {**config, "jev_classifier_config": {**expected_transport, **supplied}}
@pytest.mark.asyncio
@pytest.mark.parametrize("endpoint", ["patch", "legacy"])
@pytest.mark.parametrize(
"string_params,reset_field,config_shape",
[
(False, None, "full"), (True, None, "full"), (False, "api_key", "full"),
(False, "api_base", "full"), (False, None, "omit-provider"),
(False, None, "omit-config"), (False, None, "null-config"),
],
)
async def test_member_save_protects_stored_classifier_connection(
self, endpoint: str, string_params: bool, reset_field: str | None, config_shape: str
) -> None:
original: Final = self._row()
config: Final = {
"classifier_type": "jev", "tiers": {"SIMPLE": "allowed"},
"jev_classifier_config": {"provider": "laya", "model": "english"},
}
secret_params: Final = {
"model": "auto_router/complexity_router",
"complexity_router_config": {
**config, "jev_classifier_config": {
**config["jev_classifier_config"], "api_key": "retained-laya-secret", "api_base": "https://laya.test",
},
},
}
row: Final = original.model_copy(update={"litellm_params": secret_params})
team: Final = self._team().model_copy(update={"models": ["allowed", "laya/english"]})
database: Final = self._database(team, row)
database.transaction.litellm_proxymodeltable.update.return_value = row.model_copy(
update={"litellm_params": json.dumps(secret_params) if string_params else secret_params}
)
supplied_config: Final = {
**config, "jev_classifier_config": {
**{
key: value for key, value in config["jev_classifier_config"].items()
if key != "provider" or config_shape != "omit-provider"
},
**({reset_field: None} if reset_field is not None else {}),
},
}
patch_params: Final = (
{"complexity_router_default_model": "allowed"}
if config_shape == "omit-config"
else {"complexity_router_config": None, "complexity_router_default_model": "allowed"}
if config_shape == "null-config"
else {"complexity_router_config": supplied_config}
)
request: Final = updateDeployment(
litellm_params=updateLiteLLMParams.model_validate(patch_params),
model_info=ModelInfo(id=row.model_id, team_id="member-team"),
)
actor: Final = UserAPIKeyAuth(
user_id="owner", user_role=LitellmUserRoles.INTERNAL_USER, models=["allowed", "laya/english"], config={"timeout": 60},
)
with self._environment(database, row):
if reset_field is not None:
expected_error: Final = HTTPException if endpoint == "patch" else ProxyException
with pytest.raises(expected_error, match="Team members cannot change classifier connections") as denied:
await (
patch_model(row.model_id, request, actor) if endpoint == "patch" else update_model(request, actor)
)
assert (
denied.value.status_code if isinstance(denied.value, HTTPException) else int(denied.value.code)
) == 403
database.transaction.litellm_proxymodeltable.update.assert_not_awaited()
assert row.litellm_params == secret_params
return
response: Final = await (patch_model(row.model_id, request, actor) if endpoint == "patch" else update_model(request, actor))
written: Final = database.transaction.litellm_proxymodeltable.update.await_args.kwargs["data"]
saved: Final = json.loads(written["litellm_params"])["complexity_router_config"]["jev_classifier_config"]
assert saved == secret_params["complexity_router_config"]["jev_classifier_config"]
if config_shape in ("omit-config", "null-config"):
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
assert decrypt_value_helper(
json.loads(written["litellm_params"])["complexity_router_default_model"],
key="complexity_router_default_model", return_original_value=True,
) == "allowed"
response_payload: Final = jsonable_encoder(response)
assert "retained-laya-secret" not in json.dumps(response_payload)
response_params: Final = json.loads(response_payload["litellm_params"]) if string_params else response_payload["litellm_params"]
assert response_params == {
**secret_params, "complexity_router_config": {
**config, "jev_classifier_config": {
**config["jev_classifier_config"], "api_key": "REDACTED", "api_base": "https://laya.test",
},
},
}
assert "retained-laya-secret" in row.model_dump_json()
@pytest.mark.asyncio
@pytest.mark.parametrize("endpoint", ["patch", "legacy"])
@pytest.mark.parametrize("access", ["owner", "peer", "limited-key"])

View file

@ -1,3 +1,4 @@
import json
from collections.abc import Mapping
from dataclasses import dataclass
from typing import Final
@ -141,6 +142,8 @@ def test_tier_config_is_normalized_and_unknown_router_extras_are_rejected() -> N
({"api_key": "sk-member"}, "api_key"),
({"api_base": "https://collector.invalid", "api_key": "sk-member"}, "api_key"),
({"api_base": "https://collector.invalid", "api_key": ""}, "jev_classifier_config.api_key"),
({"provider": "laya", "model": "english", "api_base": "https://collector.invalid"}, "api_base"),
({"provider": "laya", "model": "english", "api_key": "sk-member"}, "api_key"),
],
)
def test_members_cannot_move_the_jev_classifier_off_the_proxys_typesafe_account(
@ -154,16 +157,21 @@ def test_members_cannot_move_the_jev_classifier_off_the_proxys_typesafe_account(
assert denied.value.detail == f"Invalid member auto-router configuration at {rejected_at}."
def test_members_can_still_tune_the_jev_classifier() -> None:
@pytest.mark.parametrize(("provider", "model"), [("typesafe", "jev-preview"), ("laya", "english")])
def test_members_can_still_tune_the_jev_classifier(provider: str, model: str) -> None:
validated: Final = validate_member_auto_router_config(
{
"tiers": {"SIMPLE": "allowed"},
"classifier_type": "jev",
"jev_classifier_config": {"model": "jev-preview", "timeout_ms": 500},
"jev_classifier_config": {"provider": provider, "model": model, "timeout_ms": 500},
}
)
assert validated.jev_classifier_config is not None
assert (validated.jev_classifier_config.model, validated.jev_classifier_config.timeout_ms) == ("jev-preview", 500)
assert (
validated.jev_classifier_config.provider,
validated.jev_classifier_config.model,
validated.jev_classifier_config.timeout_ms,
) == (provider, model, 500)
assert validate_member_auto_router_config(validated.model_dump()).jev_classifier_config is not None
@ -217,6 +225,90 @@ async def test_member_updates_restrict_fields_and_preserve_an_inherited_default(
assert granted.default_model == "allowed"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"nested,expected_identity,restricted",
[
("omit-config", "laya/english", False),
("omit-config", "laya/english", True),
("omit-block", None, False),
(None, None, False),
({}, None, False),
({"timeout_ms": 500}, None, False),
({"model": "english", "timeout_ms": 500}, "laya/english", False),
({"model": "english", "timeout_ms": 500}, "laya/english", True),
({"model": "multilingual"}, "laya/multilingual", False),
({"provider": "typesafe", "model": "jev-latest"}, "typesafe/jev-latest", False),
({"provider": "typesafe", "model": "jev-latest"}, "typesafe/jev-latest", True),
],
)
async def test_member_authorization_and_persistence_resolve_the_same_classifier(
catalog: Router, monkeypatch: pytest.MonkeyPatch, nested: object, expected_identity: str | None, restricted: bool
) -> None:
from litellm.proxy.management_endpoints.model_management_endpoints import (
_strategy_router_write_violation,
update_db_model,
)
from litellm.types.management_endpoints.auto_router_endpoints import RequestComplexityRouterConfig
monkeypatch.setenv("LITELLM_SALT_KEY", "member-router-test-salt")
stored_config: Final = {
"classifier_type": "jev", "tiers": {"SIMPLE": "allowed"},
"jev_classifier_config": {
"provider": "laya", "model": "english", "timeout_ms": 12000,
"api_base": "https://laya.test", "api_key": "stored-classifier-key",
},
}
existing: Final = Deployment(
model_name="member-router",
litellm_params=LiteLLM_Params(model="auto_router/complexity_router", complexity_router_config=stored_config),
model_info=ModelInfo(id="router-a", team_id="team-a"), created_by="owner",
)
incoming_config: Final = (
None if nested == "omit-config" else {
"classifier_type": "jev", "tiers": {"SIMPLE": "allowed"},
**({} if nested == "omit-block" else {"jev_classifier_config": nested}),
}
)
patch: Final = updateDeployment.model_validate({"litellm_params": {
"complexity_router_config": incoming_config, "complexity_router_default_model": "allowed",
}})
operation: Final = authorize_member_auto_router_write(
incoming=patch, existing=existing, user_api_key_dict=_actor(
models=["allowed"] if restricted or expected_identity is None else ["allowed", expected_identity],
),
team=_team(models=["allowed", "laya/english", "laya/multilingual", "typesafe/jev-latest"]),
premium_user=True, prisma_client=_Client(), llm_router=catalog,
)
violation: Final = _strategy_router_write_violation(patch.litellm_params, existing.litellm_params)
if expected_identity is None:
assert violation is not None
with pytest.raises(HTTPException) as rejected:
await operation
assert rejected.value.status_code == 400
return
assert violation is None
if restricted:
with pytest.raises(ProxyException, match=expected_identity):
await operation
return
grant: Final = await operation
persisted: Final = update_db_model(existing, patch)
saved: Final = RequestComplexityRouterConfig.model_validate(
json.loads(persisted["litellm_params"])["complexity_router_config"]
)
assert grant.config == saved
assert saved.jev_classifier_config is not None
assert f"{saved.jev_classifier_config.provider}/{saved.jev_classifier_config.model}" == expected_identity
assert saved.jev_classifier_config.api_key == (
"stored-classifier-key" if expected_identity.startswith("laya/") else None
)
assert saved.jev_classifier_config.timeout_ms == (
12000 if nested == "omit-config" else 500 if nested == {"model": "english", "timeout_ms": 500} else 3000
)
assert existing.litellm_params.complexity_router_config == stored_config
@pytest.mark.asyncio
@pytest.mark.parametrize("target", ["missing", "nested"])
async def test_member_dependencies_require_plain_configured_models(target: str) -> None:
@ -246,13 +338,17 @@ async def test_member_dependencies_require_plain_configured_models(target: str)
@pytest.mark.asyncio
@pytest.mark.parametrize("restricted", ["key", "team", None])
@pytest.mark.parametrize(("provider", "model"), [("typesafe", "jev-latest"), ("laya", "english")])
async def test_jev_evaluation_requires_model_access_but_no_completion_deployment(
catalog: Router, restricted: str | None
catalog: Router, restricted: str | None, provider: str, model: str
) -> None:
permitted: Final = ["allowed", "typesafe/jev-latest"]
permitted: Final = ["allowed", f"{provider}/{model}"]
operation: Final = authorize_member_auto_router_dependencies(
config=validate_member_auto_router_config(
{"tiers": {"SIMPLE": "allowed"}, "classifier_type": "jev", "jev_classifier_config": {}}
{
"tiers": {"SIMPLE": "allowed"}, "classifier_type": "jev",
"jev_classifier_config": {"provider": provider, "model": model},
}
),
default_model=None,
user_api_key_dict=_actor(models=["allowed"] if restricted == "key" else permitted),
@ -261,17 +357,20 @@ async def test_jev_evaluation_requires_model_access_but_no_completion_deployment
llm_router=catalog,
)
if restricted is not None:
with pytest.raises(ProxyException, match="jev-latest"):
with pytest.raises(ProxyException, match=model):
await operation
return
await operation
assert not catalog.get_model_list("typesafe/jev-latest")
assert not catalog.get_model_list(f"{provider}/{model}")
@pytest.mark.asyncio
@pytest.mark.parametrize("restricted", ["member", "project", "organization", None])
async def test_jev_evaluation_obeys_each_containing_scope(catalog: Router, restricted: str | None) -> None:
allowed: Final = ["allowed", "typesafe/jev-latest"]
@pytest.mark.parametrize(("provider", "model"), [("typesafe", "jev-latest"), ("laya", "english")])
async def test_jev_evaluation_obeys_each_containing_scope(
catalog: Router, restricted: str | None, provider: str, model: str
) -> None:
allowed: Final = ["allowed", f"{provider}/{model}"]
membership: Final = LiteLLM_TeamMembership.model_validate(
{
"user_id": "owner",
@ -293,7 +392,10 @@ async def test_jev_evaluation_obeys_each_containing_scope(catalog: Router, restr
)
operation: Final = authorize_member_auto_router_dependencies(
config=validate_member_auto_router_config(
{"tiers": {"SIMPLE": "allowed"}, "classifier_type": "jev", "jev_classifier_config": {}}
{
"tiers": {"SIMPLE": "allowed"}, "classifier_type": "jev",
"jev_classifier_config": {"provider": provider, "model": model},
}
),
default_model=None,
user_api_key_dict=_actor(models=allowed, project_id="project-a"),
@ -303,8 +405,8 @@ async def test_jev_evaluation_obeys_each_containing_scope(catalog: Router, restr
dependency_objects=MemberAutoRouterDependencyObjects(membership, organization, project),
)
if restricted is not None:
with pytest.raises(ProxyException, match="jev-latest"):
with pytest.raises(ProxyException, match=model):
await operation
return
await operation
assert not catalog.get_model_list("typesafe/jev-latest")
assert not catalog.get_model_list(f"{provider}/{model}")

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
@ -7119,6 +7121,152 @@ class TestTypeSafePassthroughRoute:
)
class TestLayaPassthroughRoute:
@pytest.fixture
def client(self, monkeypatch: pytest.MonkeyPatch) -> Iterator[TestClient]:
from litellm.proxy.proxy_server import app
monkeypatch.setenv("LAYA_API_BASE", "http://laya.test/base")
monkeypatch.setenv("TYPESAFE_API_KEY", "never-send-typesafe-key")
monkeypatch.delenv("LAYA_API_KEY", raising=False)
monkeypatch.delenv("SERVER_ROOT_PATH", raising=False)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
litellm.in_memory_llm_clients_cache.flush_cache()
monkeypatch.setitem(app.dependency_overrides, user_api_key_auth, lambda: UserAPIKeyAuth(api_key="sk-virtual"))
yield TestClient(app)
@pytest.mark.parametrize("api_key", [None, "laya-provider-key"])
def test_laya_forwards_native_decisions_without_gateway_or_typesafe_credentials(
self, client: TestClient, monkeypatch: pytest.MonkeyPatch, api_key: str | None
) -> None:
if api_key is not None:
monkeypatch.setenv("LAYA_API_KEY", api_key)
body: Final = {
"model": "english",
"state": "refund",
"questions": {"department": {"type": "choice", "criteria": {"billing": "refunds"}}},
}
answer: Final = {"model": "laya-rl-agent", "routing": {"model": "english"}, "answers": {}}
with respx.mock(assert_all_called=True) as upstream:
route: Final = upstream.post("http://laya.test/base/v1/systemone?trace=yes").respond(200, json=answer)
response: Final = client.post(
"/laya/v1/systemone?trace=yes",
json=body,
headers={"Authorization": "Bearer sk-virtual", "x-pass-authorization": "Bearer attacker"},
)
assert (response.status_code, response.json()) == (200, answer)
sent: Final = route.calls.last.request
assert sent.headers.get("authorization") == (f"Bearer {api_key}" if api_key else None)
assert json.loads(sent.content) == body
def test_laya_missing_server_fails_without_contacting_another_provider(
self, client: TestClient, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.delenv("LAYA_API_BASE")
with respx.mock(assert_all_called=False) as upstream:
response: Final = client.post("/laya/v1/systemone", json={"model": "english"})
assert response.status_code == 503
assert "LAYA_API_BASE" in response.text
assert len(upstream.calls) == 0
def test_laya_does_not_forward_unsupported_endpoints(self, client: TestClient) -> None:
with respx.mock(assert_all_called=False) as upstream:
response: Final = client.post("/laya/v1/evaluate", json={"model": "english"})
assert response.status_code == 404
assert len(upstream.calls) == 0
@pytest.mark.parametrize("model", [None, "auto", "jev-latest"])
def test_laya_rejects_implicit_checkpoint_selection(self, client: TestClient, model: str | None) -> None:
with respx.mock(assert_all_called=False) as upstream:
response: Final = client.post("/laya/v1/systemone", json={"model": model})
assert response.status_code == 400
assert len(upstream.calls) == 0
@pytest.mark.parametrize(
"controls",
[{"custom_body": {"model": "multilingual", "state": "refund"}}, {"stream": True}, {"stream": "true"}],
)
def test_laya_rejects_controls_that_change_authorized_body_or_usage_accounting(
self, client: TestClient, controls: Mapping[str, object]
) -> None:
with respx.mock(assert_all_called=False) as upstream:
route: Final = upstream.post("http://laya.test/base/v1/systemone").respond(200, json={"answers": {}})
response: Final = client.post("/laya/v1/systemone", json={"model": "english", **controls})
assert response.status_code == 400
assert not route.called
@pytest.mark.parametrize("metadata_slot", ["metadata", "litellm_metadata"])
def test_laya_hooks_enforce_canonical_model_limits_and_keep_native_wire_body(
self, client: TestClient, monkeypatch: pytest.MonkeyPatch, metadata_slot: str
) -> None:
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3
from litellm.proxy.utils import InternalUsageCache
from litellm.proxy.proxy_server import app
cache: Final = DualCache()
limiter: Final = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(cache))
auth: Final = UserAPIKeyAuth(
api_key="laya-native-rpm", metadata={"model_rpm_limit": {"laya/english": 1}},
)
def authenticated_key() -> UserAPIKeyAuth:
return auth
monkeypatch.setitem(app.dependency_overrides, user_api_key_auth, authenticated_key)
class LimitHook(CustomLogger):
async def async_pre_call_hook(
self, user_api_key_dict: UserAPIKeyAuth, cache: DualCache,
data: dict[str, object], call_type: CallTypesLiteral,
) -> dict[str, object]:
assert data["model"] == "laya/english"
metadata: Final = data.get(metadata_slot)
assert isinstance(metadata, dict)
assert "standard_logging_guardrail_information" not in metadata
assert metadata["customer_label"] == "retained"
await limiter.async_pre_call_hook(user_api_key_dict, cache, data, call_type)
return data
monkeypatch.setattr(litellm, "callbacks", [LimitHook()])
body: Final = {
"model": "english", "state": "refund",
metadata_slot: {
"customer_label": "retained", "model_group": "unbounded-client-choice",
"standard_logging_guardrail_information": [{"guardrail_cost": 25.0}],
},
}
with respx.mock(assert_all_called=True) as upstream:
route: Final = upstream.post("http://laya.test/base/v1/systemone").respond(200, json={"answers": {}})
first: Final = client.post("/laya/v1/systemone", json=body)
second: Final = client.post("/laya/v1/systemone", json=body)
assert first.status_code == 200, first.text
assert second.status_code == 429, second.text
assert route.call_count == 1
assert json.loads(route.calls.last.request.content) == {"model": "english", "state": "refund"}
def test_laya_preserves_trusted_hook_checkpoint_changes(
self, client: TestClient, monkeypatch: pytest.MonkeyPatch
) -> None:
from litellm.integrations.custom_logger import CustomLogger
class CheckpointHook(CustomLogger):
async def async_pre_call_hook(
self, user_api_key_dict: UserAPIKeyAuth, cache: DualCache,
data: dict[str, object], call_type: CallTypesLiteral,
) -> dict[str, object]:
assert data["model"] == "laya/english"
return {**data, "model": "laya/multilingual"}
monkeypatch.setattr(litellm, "callbacks", [CheckpointHook()])
with respx.mock(assert_all_called=True) as upstream:
route: Final = upstream.post("http://laya.test/base/v1/systemone").respond(200, json={"answers": {}})
response: Final = client.post("/laya/v1/systemone", json={"model": "english", "state": "refund"})
assert response.status_code == 200, response.text
assert json.loads(route.calls.last.request.content) == {"model": "multilingual", "state": "refund"}
class TestFalAIPassthroughRoute:
@pytest.fixture
def client(self, monkeypatch: pytest.MonkeyPatch) -> Iterator[TestClient]:

View file

@ -1468,7 +1468,7 @@ async def test_pass_through_request_contains_proxy_server_request_in_kwargs():
# Create mock request
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://test-proxy.com/api/endpoint"
mock_request.url = httpx.URL("http://test-proxy.com/api/endpoint")
mock_request.body = AsyncMock(return_value=b'{"message": "test request"}')
mock_request.headers = Headers({})
mock_request.query_params = QueryParams({})
@ -1573,7 +1573,7 @@ async def test_pass_through_request_streaming_marks_logging_obj_as_stream():
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://test-proxy.com/v1/messages"
mock_request.url = httpx.URL("http://test-proxy.com/v1/messages")
mock_request.body = AsyncMock(return_value=b'{"model": "claude-3", "stream": true}')
mock_request.headers = Headers({})
mock_request.query_params = QueryParams({})
@ -1635,7 +1635,7 @@ async def test_pass_through_request_sse_response_marks_logging_obj_as_stream():
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://test-proxy.com/v1/messages"
mock_request.url = httpx.URL("http://test-proxy.com/v1/messages")
mock_request.body = AsyncMock(return_value=b'{"model": "claude-3"}')
mock_request.headers = Headers({})
mock_request.query_params = QueryParams({})
@ -2505,7 +2505,7 @@ async def test_pass_through_request_query_params_forwarding():
# Create mock request with query parameters (Azure API version)
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://localhost:4000/azure-assistant/openai/assistants"
mock_request.url = httpx.URL("http://localhost:4000/azure-assistant/openai/assistants")
mock_request.body = AsyncMock(return_value=json.dumps(test_body).encode())
mock_request.headers = Headers({"Content-Type": "application/json"})
@ -3016,7 +3016,7 @@ async def test_bedrock_router_passthrough_metadata_initialization():
# Create mock request with headers
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://localhost:4000/bedrock/model/my-model/invoke"
mock_request.url = httpx.URL("http://localhost:4000/bedrock/model/my-model/invoke")
mock_request.headers = Headers(
{
"content-type": "application/json",
@ -3850,7 +3850,7 @@ def _lit3538_request():
r = MagicMock()
r.method = "POST"
r.query_params = {}
r.url = "http://testserver/mock/echo"
r.url = httpx.URL("http://testserver/mock/echo")
r.state = SimpleNamespace()
headers = MagicMock()
headers.copy.return_value = {}
@ -3983,7 +3983,7 @@ async def test_pass_through_request_non_streaming_upstream_error_returned_unchan
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://test-proxy.com/mock-upstream/api/denied"
mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/denied")
mock_request.body = AsyncMock(return_value=b'{"action": "read"}')
mock_request.headers = Headers({"content-type": "application/json"})
mock_request.query_params = QueryParams({})
@ -4069,7 +4069,7 @@ async def test_pass_through_request_upstream_error_failure_hook_exception_is_swa
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://test-proxy.com/mock-upstream/api/denied"
mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/denied")
mock_request.body = AsyncMock(return_value=b'{"action": "read"}')
mock_request.headers = Headers({"content-type": "application/json"})
mock_request.query_params = QueryParams({})
@ -4118,7 +4118,7 @@ async def test_pass_through_request_streaming_upstream_error_returned_unchanged(
mock_request = MagicMock(spec=Request)
mock_request.method = "GET"
mock_request.url = "http://test-proxy.com/mock-upstream/api/stream-denied"
mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/stream-denied")
mock_request.body = AsyncMock(return_value=b"")
mock_request.headers = Headers({})
mock_request.query_params = QueryParams({})
@ -4169,7 +4169,7 @@ class _UpstreamErrorBodyStream(httpx.AsyncByteStream):
def _upstream_error_request() -> MagicMock:
mock_request: Final = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://test-proxy.com/mock-upstream/v1beta/models/claude-nope-9:generateContent"
mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/v1beta/models/claude-nope-9:generateContent")
mock_request.body = AsyncMock(return_value=b'{"contents": []}')
mock_request.headers = Headers({"content-type": "application/json"})
mock_request.query_params = QueryParams({})
@ -4964,7 +4964,7 @@ async def test_pass_through_request_non_streaming_success_unchanged():
mock_request = MagicMock(spec=Request)
mock_request.method = "GET"
mock_request.url = "http://test-proxy.com/mock-upstream/api/success"
mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/success")
mock_request.body = AsyncMock(return_value=b"")
mock_request.headers = Headers({})
mock_request.query_params = QueryParams({})
@ -5027,7 +5027,7 @@ async def test_pass_through_request_claims_the_budget_reservation_only_when_its_
mock_get_client.return_value = MagicMock(client=async_client)
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://test-proxy.com/mock-upstream/api/generate"
mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/generate")
mock_request.body = AsyncMock(return_value=b'{"prompt": "hi"}')
mock_request.headers = Headers({"content-type": "application/json"})
mock_request.query_params = QueryParams({})
@ -5079,7 +5079,7 @@ async def test_pass_through_request_leaves_the_budget_reservation_for_the_reques
mock_get_client.return_value = MagicMock(client=async_client)
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://test-proxy.com/mock-upstream/api/generate"
mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/generate")
mock_request.body = AsyncMock(return_value=b'{"prompt": "hi"}')
mock_request.headers = Headers({"content-type": "application/json"})
mock_request.query_params = QueryParams({})
@ -5110,7 +5110,7 @@ async def test_pass_through_request_internal_failure_still_raises_proxy_exceptio
mock_request = MagicMock(spec=Request)
mock_request.method = "GET"
mock_request.url = "http://test-proxy.com/mock-upstream/api/success"
mock_request.url = httpx.URL("http://test-proxy.com/mock-upstream/api/success")
mock_request.body = AsyncMock(return_value=b"")
mock_request.headers = Headers({})
mock_request.query_params = QueryParams({})
@ -5211,7 +5211,7 @@ def _enter_relay_logging_mocks(stack, parsed_body):
def _relay_client_request(method="GET"):
mock_request = MagicMock(spec=Request)
mock_request.method = method
mock_request.url = "http://localhost:4000/passthrough-relay/results"
mock_request.url = httpx.URL("http://localhost:4000/passthrough-relay/results")
mock_request.body = AsyncMock(return_value=b"")
mock_request.headers = Headers({})
mock_request.query_params = QueryParams({})
@ -6587,7 +6587,7 @@ def _passthrough_kwargs_for_reservation(
) -> dict:
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent"
mock_request.url = httpx.URL("http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent")
mock_request.headers = Headers({})
mock_request.scope = {"endpoint": _marked_pass_through_endpoint()} if user_defined_route else {}
@ -6734,7 +6734,7 @@ async def _drive_streaming_pass_through(upstream_content_type, chunk_delay_secon
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://test-proxy.com/v1/messages"
mock_request.url = httpx.URL("http://test-proxy.com/v1/messages")
mock_request.body = AsyncMock(
return_value=b'{"model": "claude-3", "stream": true}'
if client_asked_for_stream
@ -6922,36 +6922,72 @@ def _marked_pass_through_endpoint():
return _endpoint
def test_user_defined_passthrough_is_neither_tracked_nor_enforced():
"""
`get_model_from_request` returns None for a user-defined pass-through on
purpose: the body is forwarded verbatim, so its `model` names an UPSTREAM
model rather than a LiteLLM-managed one, and enforcing key/team allowlists
against it would reject valid requests. Enforcement is therefore skipped
on those routes.
@pytest.mark.asyncio
@pytest.mark.parametrize("metadata_slot", ["metadata", "litellm_metadata"])
async def test_user_defined_passthrough_is_neither_tracked_nor_enforced(metadata_slot: str) -> None:
from datetime import datetime
Attaching the budget metadata anyway would charge a counter that nothing on
that route can refuse, and would attribute the spend to a budget the operator
scoped to a LiteLLM model that merely shares the name. Tracking and
enforcement have to agree: both on for the built-in provider routes, both off
here.
"""
kwargs = _passthrough_kwargs_for_reservation(
UserAPIKeyAuth(
token="hash",
user_id="u-1",
model_max_budget={"claude-opus-4-8": {"budget_limit": 1.0, "time_period": "18h"}},
),
user_defined_route=True,
from litellm.caching.caching import DualCache
from litellm.proxy.auth.auth_utils import get_model_from_request
from litellm.proxy.hooks.model_max_budget_limiter import _PROXY_VirtualKeyModelMaxBudgetLimiter
budget: Final = {"managed-model": {"budget_limit": 0.1, "time_period": "1d"}}
limiter: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(DualCache())
auth: Final = UserAPIKeyAuth(
api_key="custom-key", token="custom-key", team_id="shared-team", team_model_max_budget=budget,
)
metadata = kwargs["litellm_params"]["metadata"]
for field in (
"user_api_key_model_max_budget",
"user_api_key_user_model_max_budget",
"user_api_key_end_user_model_max_budget",
):
assert field not in metadata, f"{field} was attached on a route that never enforces it"
endpoint: Final = create_pass_through_route(
endpoint="/custom-budget-test", target="https://upstream.test/echo", custom_headers={}, cost_per_request=0.25,
)
request: Final = Request({
"type": "http", "method": "POST", "path": "/custom-budget-test", "headers": [],
"query_string": b"", "endpoint": endpoint,
})
body: Final = {
"model": "upstream-only-model", metadata_slot: {
"model_group": "managed-model", "customer_label": "retained",
"user_api_key_team_model_max_budget": budget,
},
}
assert get_model_from_request(body, "/custom-budget-test", request=request) is None
assert await limiter.is_team_within_model_budget("shared-team", budget, None, "managed-model")
start: Final = datetime.now()
logging_obj: Final = LiteLLMLoggingObj(
model="upstream-only-model", messages=[], stream=False, call_type="pass_through_endpoint",
start_time=start, litellm_call_id="custom-budget", function_id="custom-budget", kwargs={},
dynamic_async_success_callbacks=[limiter],
)
payload: Final = {
"url": "https://upstream.test/echo", "request_body": body, "request_method": "POST", "cost_per_request": 0.25,
}
kwargs: Final = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint(
request=request, user_api_key_dict=auth, passthrough_logging_payload=payload, logging_obj=logging_obj,
_parsed_body=body, litellm_call_id="custom-budget",
)
logging_obj.update_environment_variables(
model="upstream-only-model", user="unknown", optional_params={},
litellm_params=kwargs["litellm_params"], call_type="pass_through_endpoint",
)
response: Final = httpx.Response(
200, request=httpx.Request("POST", "https://upstream.test/echo"), json={"ok": True},
)
await PassThroughEndpointLogging().pass_through_async_success_handler(
httpx_response=response, response_body={"ok": True}, request_body=body, logging_obj=logging_obj,
url_route="https://upstream.test/echo", result=response.text, start_time=start, end_time=datetime.now(),
cache_hit=False, **kwargs,
)
assert logging_obj.model_call_details["response_cost"] == 0.25
assert await limiter.is_team_within_model_budget("shared-team", budget, None, "managed-model")
metadata: Final = kwargs["litellm_params"]["metadata"]
assert (metadata["model_group"], metadata["customer_label"]) == ("managed-model", "retained")
assert metadata.keys().isdisjoint({
"user_api_key_model_max_budget", "user_api_key_team_model_max_budget",
"user_api_key_user_model_max_budget", "user_api_key_end_user_model_max_budget",
})
assert metadata.keys().isdisjoint({
"user_api_key_model_max_budget", "user_api_key_team_model_max_budget",
"user_api_key_user_model_max_budget", "user_api_key_end_user_model_max_budget",
})
@pytest.mark.parametrize(
@ -7279,7 +7315,7 @@ def test_passthrough_client_cannot_forge_session_id_omission(client_metadata_key
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent"
mock_request.url = httpx.URL("http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent")
mock_request.headers = Headers({})
mock_request.scope = {}
@ -7312,7 +7348,7 @@ def test_passthrough_logs_the_resolved_deployment_model_info_over_the_request_bo
the call to (LIT-1761: passthrough successes carried model_id="")."""
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://0.0.0.0:4000/vertex_ai/v1/projects/p/locations/global/publishers/google/models/gemini-3.8-flash:generateContent"
mock_request.url = httpx.URL("http://0.0.0.0:4000/vertex_ai/v1/projects/p/locations/global/publishers/google/models/gemini-3.8-flash:generateContent")
mock_request.headers = Headers({})
mock_request.scope = {}
mock_request.state = SimpleNamespace(
@ -7344,7 +7380,7 @@ _PROXY_SERVER_REQUEST: Final = TypeAdapter(dict[str, object])
def _split_pass_through_body(body: str) -> _PassThroughSplit:
mock_request: Final = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent"
mock_request.url = httpx.URL("http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent")
mock_request.headers = Headers()
mock_request.scope = MappingProxyType({})
@ -7600,7 +7636,7 @@ def test_passthrough_attributes_a_cli_session_to_its_alias_not_the_login_token()
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://0.0.0.0:4000/anthropic/v1/messages"
mock_request.url = httpx.URL("http://0.0.0.0:4000/anthropic/v1/messages")
mock_request.headers = Headers({})
mock_request.scope = {}
session = UserAPIKeyAuth(

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

@ -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,6 +8,7 @@ from unittest.mock import create_autospec
import httpx
import pytest
import respx
import litellm
from litellm._logging import verbose_router_logger
@ -30,14 +31,15 @@ from litellm.types.utils import AUTOROUTER_CLASSIFIER_CALL_ORIGIN
class _UsageRecorder(CustomLogger):
def __init__(self) -> None:
def __init__(self, model_key: str = "typesafe/jev-accounting") -> None:
super().__init__()
self.model_key = model_key
self.calls: tuple[Mapping[str, object], ...] = ()
async def async_log_success_event(
self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime
) -> None:
if str(kwargs.get("model", "")).removeprefix("typesafe/") != "jev-accounting":
if str(kwargs.get("model", "")) != self.model_key:
return
self.calls = (*self.calls, kwargs)
@ -420,6 +422,62 @@ def test_jev_config_requires_classifier_config() -> None:
ComplexityRouterConfig.model_validate({"classifier_type": "jev"})
@pytest.mark.parametrize("config", [{"provider": "laya"}, {"provider": "laya", "model": " "}])
def test_laya_requires_its_own_checkpoint(config: Mapping[str, object]) -> None:
with pytest.raises(ValueError, match="Laya model must be"):
JevClassifierConfig.model_validate(config)
@pytest.mark.asyncio
@pytest.mark.parametrize("custom_base", [False, True])
async def test_laya_routes_with_its_own_credentials_and_accounts_the_checkpoint(
monkeypatch: pytest.MonkeyPatch, custom_base: bool
) -> None:
monkeypatch.setenv("TYPESAFE_API_KEY", "never-send-typesafe-key")
monkeypatch.setenv("LAYA_API_BASE", "https://laya.test")
monkeypatch.setenv("LAYA_API_KEY", "laya-env-key")
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
monkeypatch.setitem(litellm.model_cost, "laya/english", {"input_cost_per_token": 0.01})
recorder: Final = _UsageRecorder("laya/english")
monkeypatch.setattr(litellm, "_async_success_callback", [recorder])
router: Final = ComplexityRouter(
"laya-route",
litellm.Router(model_list=[]),
{
"classifier_type": "jev",
"jev_classifier_config": {
"provider": "laya",
"model": "english",
**({"api_base": "https://laya.test"} if custom_base else {}),
},
"tiers": {"SIMPLE": "cheap"},
},
derive_savings_baseline=False,
)
with respx.mock(assert_all_called=True) as upstream:
route: Final = upstream.post("https://laya.test/v1/systemone").respond(
200,
json={
"model": "laya-rl-agent",
"routing": {"model": "english"},
"answers": {"tier": _answer().model_dump()},
"usage": {"input_tokens": 31, "output_tokens": 0},
},
)
outcome: Final = await router.aclassify("choose a tier")
await GLOBAL_LOGGING_WORKER.flush()
assert outcome.cause == "jev_classifier"
assert outcome.jev_verdict is not None
assert (outcome.jev_verdict.provider, outcome.jev_verdict.model) == ("laya", "english")
assert outcome.classifier_cost == pytest.approx(0.31)
sent: Final = route.calls.last.request
assert sent.headers.get("authorization") == (None if custom_base else "Bearer laya-env-key")
assert json.loads(sent.content)["model"] == "english"
assert len(recorder.calls) == 1
assert recorder.calls[0]["response_cost"] == pytest.approx(0.31)
def test_jev_config_is_rejected_for_other_classifier_types() -> None:
with pytest.raises(ValueError, match="has no effect"):
ComplexityRouterConfig.model_validate(

View file

@ -23,21 +23,23 @@ COMPLEXITY_FIELDS = frozenset({"complexity_router_config"})
SEMANTIC_FIELDS = frozenset({"auto_router_config", "auto_router_default_model", "auto_router_embedding_model"})
@pytest.mark.parametrize("model", ["jev-latest", "jev-preview"])
def test_jev_enumerates_a_paid_evaluation_without_a_completion_classifier(model: str) -> None:
found = strategy_router_dependencies(
@pytest.mark.parametrize(
("provider", "model"), [(None, "jev-latest"), ("typesafe", "jev-preview"), ("laya", "english")]
)
def test_jev_enumerates_a_paid_evaluation_without_a_completion_classifier(provider: str | None, model: str) -> None:
found: Final = strategy_router_dependencies(
{
"model": "auto_router/complexity_router",
"complexity_router_config": {
"classifier_type": "jev",
"jev_classifier_config": {"model": model},
"jev_classifier_config": {"model": model, **({"provider": provider} if provider else {})},
"tiers": {"SIMPLE": "cheap"},
},
}
)
assert tuple((dep.model_name, dep.role) for dep in found) == (
("cheap", "tier"),
(f"typesafe/{model}", "evaluation"),
(f"{provider or 'typesafe'}/{model}", "evaluation"),
)

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

@ -8575,6 +8575,23 @@ export interface paths {
patch: operations["langfuse_proxy_route_langfuse__endpoint__patch"];
trace?: never;
};
"/laya/v1/systemone": {
parameters: {
query?: never;
header?: never;
path?: never;
cookie?: never;
};
get?: never;
put?: never;
/** Laya Proxy Route */
post: operations["laya_proxy_route_laya_v1_systemone_post"];
delete?: never;
options?: never;
head?: never;
patch?: never;
trace?: never;
};
"/lazy/warm/{name}": {
parameters: {
query?: never;
@ -31611,12 +31628,12 @@ export interface components {
JevClassifierConfig: {
/**
* Api Base
* @description TypeSafe API base, falling back to TYPESAFE_API_BASE and then https://api.typesafe.ai
* @description Provider API base; defaults to TYPESAFE_API_BASE or LAYA_API_BASE for the selected provider
*/
api_base?: string | null;
/**
* Api Key
* @description TypeSafe API key, falling back to TYPESAFE_API_KEY
* @description Provider API key; optional for self-hosted Laya
*/
api_key?: string | null;
/**
@ -31639,6 +31656,12 @@ export interface components {
* @default jev-latest
*/
model: string;
/**
* Provider
* @default typesafe
* @enum {string}
*/
provider: "typesafe" | "laya";
/**
* Timeout Ms
* @default 3000
@ -58474,6 +58497,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;