mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
feat: add Bespoke Nimble gateway and OSS classifier support (#44246)
* feat: add Bespoke Nimble gateway and OSS classifier support * feat: accept Ollama's nimble model name for the Bespoke provider * test: exempt the POST-only bespoke decisions route from the all-methods check test_pass_through_routes_support_all_methods requires every built-in pass-through route to accept every HTTP method unless it is listed in PROTOCOL_CONSTRAINED_PASS_THROUGH_ROUTES. /bespoke/v1/systemone is POST-only like /laya/v1/systemone, so the test failed at this branch and passed at the merge base. List it alongside Laya. Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com> --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Co-authored-by: Claude Fable 5.1 <noreply@anthropic.com>
This commit is contained in:
parent
c5c0a48ef1
commit
f63d989ff9
33 changed files with 577 additions and 270 deletions
|
|
@ -1,51 +1,7 @@
|
|||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Final, Literal, TypeAlias
|
||||
from typing import Final
|
||||
|
||||
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)
|
||||
from pydantic import BaseModel, TypeAdapter, ValidationError
|
||||
|
||||
|
||||
class _LayaRouting(BaseModel):
|
||||
|
|
|
|||
56
litellm/llms/oss_decision.py
Normal file
56
litellm/llms/oss_decision.py
Normal file
|
|
@ -0,0 +1,56 @@
|
|||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass, field
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal, TypeAlias
|
||||
|
||||
from pydantic import AnyHttpUrl, TypeAdapter, ValidationError
|
||||
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
OssDecisionProvider: TypeAlias = Literal["laya", "bespoke"]
|
||||
OSS_DECISION_MODELS: Final = MappingProxyType(
|
||||
{
|
||||
"laya": ("english", "multilingual", "typed-decisions"),
|
||||
"bespoke": ("nimble-latest", "nimble", "bespokelabs/Bespoke-Nimble-9B"),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def validate_oss_model(provider: OssDecisionProvider, value: object) -> str:
|
||||
if not isinstance(value, str) or value not in OSS_DECISION_MODELS[provider]:
|
||||
raise ValueError(f"{provider} model must be one of {', '.join(OSS_DECISION_MODELS[provider])}")
|
||||
return value
|
||||
|
||||
|
||||
def validate_oss_request(provider: OssDecisionProvider, body: Mapping[str, object]) -> str:
|
||||
if "custom_body" in body:
|
||||
raise ValueError(f"custom_body is not supported for {provider} requests")
|
||||
if body.get("stream"):
|
||||
raise ValueError(f"Streaming is not supported for {provider} requests")
|
||||
return validate_oss_model(provider, body.get("model"))
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class OssDecisionConnection:
|
||||
api_base: str
|
||||
api_key: str | None = field(repr=False)
|
||||
|
||||
|
||||
def validate_oss_api_base(provider: OssDecisionProvider, value: str) -> str:
|
||||
try:
|
||||
url: Final = TypeAdapter(AnyHttpUrl).validate_python(value)
|
||||
except ValidationError as exc:
|
||||
raise ValueError(f"{provider} 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(f"{provider} api_base must not contain credentials, a query, or a fragment")
|
||||
return str(url).rstrip("/")
|
||||
|
||||
|
||||
def oss_connection(
|
||||
provider: OssDecisionProvider, api_base: str | None = None, api_key: str | None = None
|
||||
) -> OssDecisionConnection:
|
||||
base: Final = api_base if api_base is not None else get_secret_str(f"{provider.upper()}_API_BASE")
|
||||
if not base:
|
||||
raise ValueError(f"{provider} requires api_base or {provider.upper()}_API_BASE pointing to its server")
|
||||
key: Final = api_key if api_base is not None else api_key or get_secret_str(f"{provider.upper()}_API_KEY")
|
||||
return OssDecisionConnection(api_base=validate_oss_api_base(provider, base), api_key=key)
|
||||
|
|
@ -72632,6 +72632,48 @@
|
|||
"supports_audio_input": true,
|
||||
"supports_video_input": true
|
||||
},
|
||||
"bespoke/nimble-latest": {
|
||||
"input_cost_per_token": 0.0,
|
||||
"litellm_provider": "bespoke",
|
||||
"max_input_tokens": 8192,
|
||||
"mode": "evaluation",
|
||||
"output_cost_per_token": 0.0,
|
||||
"source": "https://github.com/bespokelabsai/nimble",
|
||||
"supported_endpoints": [
|
||||
"/v1/systemone"
|
||||
],
|
||||
"metadata": {
|
||||
"notes": "Self-hosted decision model; infrastructure costs are paid separately"
|
||||
}
|
||||
},
|
||||
"bespoke/nimble": {
|
||||
"input_cost_per_token": 0.0,
|
||||
"litellm_provider": "bespoke",
|
||||
"max_input_tokens": 8192,
|
||||
"mode": "evaluation",
|
||||
"output_cost_per_token": 0.0,
|
||||
"source": "https://ollama.com/library/nimble",
|
||||
"supported_endpoints": [
|
||||
"/v1/systemone"
|
||||
],
|
||||
"metadata": {
|
||||
"notes": "Self-hosted decision model under the name Ollama serves it as; infrastructure costs are paid separately"
|
||||
}
|
||||
},
|
||||
"bespoke/bespokelabs/Bespoke-Nimble-9B": {
|
||||
"input_cost_per_token": 0.0,
|
||||
"litellm_provider": "bespoke",
|
||||
"max_input_tokens": 8192,
|
||||
"mode": "evaluation",
|
||||
"output_cost_per_token": 0.0,
|
||||
"source": "https://github.com/bespokelabsai/nimble",
|
||||
"supported_endpoints": [
|
||||
"/v1/systemone"
|
||||
],
|
||||
"metadata": {
|
||||
"notes": "Self-hosted decision model; infrastructure costs are paid separately"
|
||||
}
|
||||
},
|
||||
"laya/english": {
|
||||
"input_cost_per_token": 0.0,
|
||||
"litellm_provider": "laya",
|
||||
|
|
|
|||
|
|
@ -230,6 +230,7 @@ LAZY_FEATURES: Final[tuple[LazyFeature, ...]] = (
|
|||
"/transcribe",
|
||||
"/typesafe/",
|
||||
"/laya/",
|
||||
"/bespoke/",
|
||||
"/openrouter/",
|
||||
"/vertex-ai/",
|
||||
"/vertex_ai/",
|
||||
|
|
|
|||
|
|
@ -26318,6 +26318,30 @@
|
|||
]
|
||||
}
|
||||
},
|
||||
"/bespoke/v1/systemone": {
|
||||
"post": {
|
||||
"operationId": "bespoke_proxy_route_bespoke_v1_systemone_post",
|
||||
"responses": {
|
||||
"200": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {}
|
||||
}
|
||||
},
|
||||
"description": "Successful Response"
|
||||
}
|
||||
},
|
||||
"security": [
|
||||
{
|
||||
"APIKeyHeader": []
|
||||
}
|
||||
],
|
||||
"summary": "Bespoke Proxy Route",
|
||||
"tags": [
|
||||
"llm_passthrough"
|
||||
]
|
||||
}
|
||||
},
|
||||
"/cohere/{endpoint}": {
|
||||
"delete": {
|
||||
"description": "[Docs](https://docs.litellm.ai/docs/pass_through/cohere)",
|
||||
|
|
|
|||
|
|
@ -511,6 +511,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
"/mistral",
|
||||
"/typesafe",
|
||||
"/laya",
|
||||
"/bespoke",
|
||||
"/openrouter",
|
||||
"/milvus",
|
||||
"/gigachat",
|
||||
|
|
|
|||
|
|
@ -1883,15 +1883,16 @@ 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
|
||||
if route.rstrip("/") in ("/laya/v1/systemone", "/bespoke/v1/systemone"):
|
||||
from litellm.llms.oss_decision import validate_oss_model
|
||||
|
||||
provider: Final = "bespoke" if route.startswith("/bespoke/") else "laya"
|
||||
try:
|
||||
laya_request: Final = TypeAdapter(Mapping[str, object]).validate_python(request_data)
|
||||
laya_model: Final = validate_laya_model(laya_request.get("model"))
|
||||
decision_request: Final = TypeAdapter(Mapping[str, object]).validate_python(request_data)
|
||||
decision_model: Final = validate_oss_model(provider, decision_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}",))
|
||||
return _dedupe_model_candidates((f"{provider}/{decision_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)
|
||||
|
|
|
|||
|
|
@ -74,7 +74,7 @@ class _MemberOpenSourceClassifierConfig(BaseModel):
|
|||
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
provider: Literal["jev", "laya"] = "jev"
|
||||
provider: Literal["jev", "laya", "bespoke"] = "jev"
|
||||
model: str
|
||||
api_key: None = None
|
||||
api_base: None = None
|
||||
|
|
|
|||
|
|
@ -58,8 +58,8 @@ 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.oss_decision import OssDecisionProvider, oss_connection, validate_oss_request
|
||||
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
|
||||
from litellm.passthrough.main import AsyncPassthroughStreamingResponse
|
||||
from litellm.proxy._types import *
|
||||
|
|
@ -646,17 +646,32 @@ async def laya_proxy_route(
|
|||
request: Request,
|
||||
fastapi_response: Response,
|
||||
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
|
||||
) -> Response:
|
||||
return await _oss_decision_proxy_route("laya", request, fastapi_response, user_api_key_dict)
|
||||
|
||||
|
||||
@router.post("/bespoke/v1/systemone", tags=["Bespoke Nimble Pass-through", "pass-through"])
|
||||
async def bespoke_proxy_route(
|
||||
request: Request,
|
||||
fastapi_response: Response,
|
||||
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
|
||||
) -> Response:
|
||||
return await _oss_decision_proxy_route("bespoke", request, fastapi_response, user_api_key_dict)
|
||||
|
||||
|
||||
async def _oss_decision_proxy_route(
|
||||
provider: OssDecisionProvider, request: Request, fastapi_response: Response, user_api_key_dict: UserAPIKeyAuth
|
||||
) -> Response:
|
||||
body: Final = TypeAdapter(dict[str, object]).validate_python(await _read_request_body(request))
|
||||
try:
|
||||
_ = validate_laya_request(body)
|
||||
_ = validate_oss_request(provider, body)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
try:
|
||||
connection: Final = laya_connection()
|
||||
connection: Final = oss_connection(provider)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(
|
||||
status_code=503, detail="Laya server is not configured correctly; check LAYA_API_BASE"
|
||||
status_code=503, detail=f"{provider} server is not configured correctly; check {provider.upper()}_API_BASE"
|
||||
) from exc
|
||||
base_url: Final = httpx.URL(connection.api_base)
|
||||
updated_url: Final = base_url.copy_with(
|
||||
|
|
@ -671,7 +686,7 @@ async def laya_proxy_route(
|
|||
endpoint="v1/systemone",
|
||||
target=str(updated_url),
|
||||
custom_headers=MappingProxyType({**authorization, "Content-Type": "application/json"}),
|
||||
custom_llm_provider="laya",
|
||||
custom_llm_provider=provider,
|
||||
is_streaming_request=False,
|
||||
)
|
||||
return TypeAdapter(Response, config=ConfigDict(arbitrary_types_allowed=True)).validate_python(
|
||||
|
|
|
|||
|
|
@ -65,7 +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.llms.oss_decision import validate_oss_request
|
||||
from litellm.passthrough import BasePassthroughUtils
|
||||
from litellm.proxy._types import (
|
||||
ConfigFieldInfo,
|
||||
|
|
@ -387,7 +387,9 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
|
|||
@staticmethod
|
||||
def get_endpoint_type(url: str, custom_llm_provider: str | None = None) -> EndpointType:
|
||||
parsed_url: Final = urlparse(url)
|
||||
if custom_llm_provider == "typesafe" and parsed_url.path.removesuffix("/").endswith("/v1/systemone"):
|
||||
if custom_llm_provider in ("typesafe", "laya", "bespoke") and parsed_url.path.removesuffix("/").endswith(
|
||||
"/v1/systemone"
|
||||
):
|
||||
return EndpointType.DECISIONS
|
||||
if (
|
||||
("generateContent") in url
|
||||
|
|
@ -1163,10 +1165,10 @@ async def pass_through_request(
|
|||
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}"
|
||||
if custom_llm_provider in ("laya", "bespoke"):
|
||||
decision_request: Final = TypeAdapter(Mapping[str, object]).validate_python(_parsed_body)
|
||||
checkpoint: Final = validate_oss_request(custom_llm_provider, decision_request)
|
||||
_parsed_body["model"] = f"{custom_llm_provider}/{checkpoint}"
|
||||
|
||||
### COLLECT GUARDRAILS FOR PASSTHROUGH ENDPOINT ###
|
||||
# Passthrough endpoints are opt-in only for guardrails
|
||||
|
|
@ -1223,17 +1225,19 @@ async def pass_through_request(
|
|||
call_type="pass_through_endpoint",
|
||||
endpoint_type=endpoint_type,
|
||||
)
|
||||
if custom_llm_provider == "laya":
|
||||
if custom_llm_provider in ("laya", "bespoke"):
|
||||
hook_body: Final = TypeAdapter(dict[str, object]).validate_python(_parsed_body)
|
||||
hook_model: Final = hook_body.get("model")
|
||||
laya_body: Final = MappingProxyType(
|
||||
decision_body: Final = MappingProxyType(
|
||||
{
|
||||
**hook_body,
|
||||
"model": hook_model.removeprefix("laya/") if isinstance(hook_model, str) else hook_model,
|
||||
"model": hook_model.removeprefix(f"{custom_llm_provider}/")
|
||||
if isinstance(hook_model, str)
|
||||
else hook_model,
|
||||
}
|
||||
)
|
||||
_ = validate_laya_request(laya_body)
|
||||
_parsed_body = TypeAdapter(dict[str, object]).validate_python(laya_body)
|
||||
_ = validate_oss_request(custom_llm_provider, decision_body)
|
||||
_parsed_body = TypeAdapter(dict[str, object]).validate_python(decision_body)
|
||||
resolved_timeout: Final = resolve_pass_through_request_timeout(timeout)
|
||||
async_client_obj: Final = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.PassThroughEndpoint,
|
||||
|
|
|
|||
|
|
@ -336,7 +336,7 @@ class PassThroughEndpointLogging:
|
|||
kwargs = transcribe_handler_result["kwargs"] # rebind-ok: elif-chain contract
|
||||
elif (
|
||||
self.is_typesafe_route(custom_llm_provider)
|
||||
or custom_llm_provider == "laya"
|
||||
or custom_llm_provider in ("laya", "bespoke")
|
||||
or self.is_openrouter_decisions_route(url_route, custom_llm_provider)
|
||||
):
|
||||
from .llm_provider_handlers.typesafe_passthrough_logging_handler import (
|
||||
|
|
|
|||
|
|
@ -1309,15 +1309,15 @@ class ComplexityRouter(CustomLogger):
|
|||
|
||||
@staticmethod
|
||||
def _build_jev_client(config: OpenSourceClassifierConfig) -> JevClassifierClient:
|
||||
if config.provider == "laya":
|
||||
from litellm.llms.laya.common_utils import laya_connection
|
||||
if config.provider in ("laya", "bespoke"):
|
||||
from litellm.llms.oss_decision import oss_connection
|
||||
|
||||
connection: Final = laya_connection(config.api_base, config.api_key)
|
||||
connection: Final = oss_connection(config.provider, 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",
|
||||
provider=config.provider,
|
||||
)
|
||||
api_key: Final = config.api_key or get_secret_str("TYPESAFE_API_KEY")
|
||||
if not api_key:
|
||||
|
|
@ -2228,7 +2228,7 @@ 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"
|
||||
accounting_provider: Final = "typesafe" if config.provider == "jev" else config.provider
|
||||
verdict: Final = JevVerdict(
|
||||
label=answer.choice,
|
||||
probabilities=answer.probabilities,
|
||||
|
|
@ -2243,8 +2243,8 @@ class ComplexityRouter(CustomLogger):
|
|||
tier=tier,
|
||||
score=None,
|
||||
signals=(
|
||||
f"{'laya' if config.provider == 'laya' else 'jev'}-classifier:{tier_name}",
|
||||
f"{'laya' if config.provider == 'laya' else 'jev'}-confidence={answer.confidence:.6f}",
|
||||
f"{config.provider}-classifier:{tier_name}",
|
||||
f"{config.provider}-confidence={answer.confidence:.6f}",
|
||||
*(
|
||||
f"tier-probability:{label}={probability:.6f}"
|
||||
for label, probability in answer.probabilities.items()
|
||||
|
|
|
|||
|
|
@ -698,17 +698,17 @@ def normalize_classifier_config_aliases(config: Mapping[str, object]) -> Mapping
|
|||
class OpenSourceClassifierConfig(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid", frozen=True)
|
||||
|
||||
provider: Literal["jev", "laya"] = "jev"
|
||||
provider: Literal["jev", "laya", "bespoke"] = "jev"
|
||||
model: str = "jev-latest"
|
||||
api_key: str | None = Field(default=None, description="Provider API key; optional for self-hosted Laya")
|
||||
api_key: str | None = Field(default=None, description="Provider API key; optional for self-hosted providers")
|
||||
api_base: str | None = Field(
|
||||
default=None,
|
||||
description="Provider API base; defaults to TYPESAFE_API_BASE or LAYA_API_BASE for the selected provider",
|
||||
description="Provider API base; defaults to the selected provider API_BASE environment variable",
|
||||
)
|
||||
timeout_ms: int = Field(default=3000, ge=1)
|
||||
instructions: str | None = Field(
|
||||
default=None,
|
||||
description="Replaces the built-in Jev question instructions",
|
||||
description="Replaces the built-in classification instructions",
|
||||
)
|
||||
circuit_breaker_enabled: bool = True
|
||||
circuit_breaker_cooldown_seconds: float = Field(default=30.0, gt=0.0)
|
||||
|
|
@ -729,17 +729,19 @@ class OpenSourceClassifierConfig(BaseModel):
|
|||
@classmethod
|
||||
def _reject_blank_api_key(cls, value: str | None) -> str | None:
|
||||
if value is not None and not value.strip():
|
||||
raise ValueError("opensource_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 the provider environment key"
|
||||
)
|
||||
return value
|
||||
|
||||
@model_validator(mode="after")
|
||||
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
|
||||
if self.provider in ("laya", "bespoke"):
|
||||
from litellm.llms.oss_decision import validate_oss_api_base, validate_oss_model
|
||||
|
||||
_ = validate_laya_model(self.model)
|
||||
_ = validate_oss_model(self.provider, self.model)
|
||||
if self.api_base is not None:
|
||||
_ = validate_laya_api_base(self.api_base)
|
||||
_ = validate_oss_api_base(self.provider, self.api_base)
|
||||
return self
|
||||
if self.api_base is not None and self.api_key is None:
|
||||
raise ValueError(
|
||||
|
|
@ -1150,7 +1152,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 'oss_classifier', a structured choice call using Jev or Laya"
|
||||
"everywhere except when its score lands near a tier boundary, or 'oss_classifier', a structured choice call using Jev, Laya or Bespoke Nimble"
|
||||
),
|
||||
)
|
||||
llm_v2_config: LLMV2Config | None = Field(
|
||||
|
|
|
|||
|
|
@ -84,7 +84,7 @@ class HttpJevClassifierClient:
|
|||
api_key: str | None,
|
||||
api_base: str,
|
||||
http_client: AsyncHTTPHandler,
|
||||
provider: Literal["typesafe", "laya"] = "typesafe",
|
||||
provider: Literal["typesafe", "laya", "bespoke"] = "typesafe",
|
||||
) -> None:
|
||||
self._api_key = api_key
|
||||
self._api_base = api_base.rstrip("/")
|
||||
|
|
@ -201,7 +201,7 @@ class JevVerdict(NamedTuple):
|
|||
confidence: float
|
||||
model: str
|
||||
cost: float | None
|
||||
provider: Literal["typesafe", "laya"] = "typesafe"
|
||||
provider: Literal["typesafe", "laya", "bespoke"] = "typesafe"
|
||||
|
||||
|
||||
class _RegistryPricing(BaseModel):
|
||||
|
|
@ -225,7 +225,7 @@ def build_jev_request(
|
|||
|
||||
|
||||
def jev_classifier_cost(
|
||||
response: JevSystemOneResponse, configured_model: str, provider: Literal["typesafe", "laya"] = "typesafe"
|
||||
response: JevSystemOneResponse, configured_model: str, provider: Literal["typesafe", "laya", "bespoke"] = "typesafe"
|
||||
) -> float | None:
|
||||
usage: Final = response.usage
|
||||
if usage is None:
|
||||
|
|
|
|||
|
|
@ -72632,6 +72632,48 @@
|
|||
"supports_audio_input": true,
|
||||
"supports_video_input": true
|
||||
},
|
||||
"bespoke/nimble-latest": {
|
||||
"input_cost_per_token": 0.0,
|
||||
"litellm_provider": "bespoke",
|
||||
"max_input_tokens": 8192,
|
||||
"mode": "evaluation",
|
||||
"output_cost_per_token": 0.0,
|
||||
"source": "https://github.com/bespokelabsai/nimble",
|
||||
"supported_endpoints": [
|
||||
"/v1/systemone"
|
||||
],
|
||||
"metadata": {
|
||||
"notes": "Self-hosted decision model; infrastructure costs are paid separately"
|
||||
}
|
||||
},
|
||||
"bespoke/nimble": {
|
||||
"input_cost_per_token": 0.0,
|
||||
"litellm_provider": "bespoke",
|
||||
"max_input_tokens": 8192,
|
||||
"mode": "evaluation",
|
||||
"output_cost_per_token": 0.0,
|
||||
"source": "https://ollama.com/library/nimble",
|
||||
"supported_endpoints": [
|
||||
"/v1/systemone"
|
||||
],
|
||||
"metadata": {
|
||||
"notes": "Self-hosted decision model under the name Ollama serves it as; infrastructure costs are paid separately"
|
||||
}
|
||||
},
|
||||
"bespoke/bespokelabs/Bespoke-Nimble-9B": {
|
||||
"input_cost_per_token": 0.0,
|
||||
"litellm_provider": "bespoke",
|
||||
"max_input_tokens": 8192,
|
||||
"mode": "evaluation",
|
||||
"output_cost_per_token": 0.0,
|
||||
"source": "https://github.com/bespokelabsai/nimble",
|
||||
"supported_endpoints": [
|
||||
"/v1/systemone"
|
||||
],
|
||||
"metadata": {
|
||||
"notes": "Self-hosted decision model; infrastructure costs are paid separately"
|
||||
}
|
||||
},
|
||||
"laya/english": {
|
||||
"input_cost_per_token": 0.0,
|
||||
"litellm_provider": "laya",
|
||||
|
|
|
|||
|
|
@ -1477,6 +1477,13 @@
|
|||
"rerank": false
|
||||
}
|
||||
},
|
||||
"bespoke": {
|
||||
"display_name": "Bespoke Nimble (`bespoke`)",
|
||||
"url": "https://docs.litellm.ai/docs/auto_router/decision_classifiers",
|
||||
"endpoints": {
|
||||
"systemone": true
|
||||
}
|
||||
},
|
||||
"laya": {
|
||||
"display_name": "Laya (`laya`)",
|
||||
"url": "https://docs.litellm.ai/docs/auto_router/decision_classifiers",
|
||||
|
|
|
|||
|
|
@ -416,6 +416,7 @@ PROTOCOL_CONSTRAINED_PASS_THROUGH_ROUTES = {
|
|||
"/transcribe/{operation}": {"POST"},
|
||||
"/tinyfish/{endpoint:path}": {"GET", "POST"},
|
||||
"/laya/v1/systemone": {"POST"},
|
||||
"/bespoke/v1/systemone": {"POST"},
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1,48 +1,8 @@
|
|||
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()
|
||||
from litellm.llms.laya.common_utils import laya_response_model
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
|
|
|
|||
60
tests/unit/llms/test_oss_decision.py
Normal file
60
tests/unit/llms/test_oss_decision.py
Normal file
|
|
@ -0,0 +1,60 @@
|
|||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.llms.oss_decision import OssDecisionProvider, oss_connection, validate_oss_request
|
||||
|
||||
pytestmark: Final = pytest.mark.parametrize("provider", ["laya", "bespoke"])
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("base", "key", "expected_base", "expected_key"),
|
||||
[
|
||||
(None, None, "http://decision.test/root", "oss-env-key"),
|
||||
("http://custom.test/", None, "http://custom.test", None),
|
||||
("http://custom.test/", "explicit-key", "http://custom.test", "explicit-key"),
|
||||
],
|
||||
)
|
||||
def test_oss_credentials_stay_with_their_configured_destination(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
provider: OssDecisionProvider,
|
||||
base: str | None,
|
||||
key: str | None,
|
||||
expected_base: str,
|
||||
expected_key: str | None,
|
||||
) -> None:
|
||||
monkeypatch.setenv(f"{provider.upper()}_API_BASE", "http://decision.test/root/")
|
||||
monkeypatch.setenv(f"{provider.upper()}_API_KEY", "oss-env-key")
|
||||
monkeypatch.setenv("TYPESAFE_API_KEY", "never-send-this")
|
||||
monkeypatch.setenv("NIMBLE_API_KEY", "never-send-nimble-search-key")
|
||||
connection: Final = oss_connection(provider, 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_oss_rejects_ambiguous_server_urls(provider: OssDecisionProvider, base: str) -> None:
|
||||
with pytest.raises(ValueError, match=provider):
|
||||
oss_connection(provider, base)
|
||||
|
||||
|
||||
def test_oss_missing_server_does_not_fall_back_to_typesafe(
|
||||
monkeypatch: pytest.MonkeyPatch, provider: OssDecisionProvider
|
||||
) -> None:
|
||||
monkeypatch.delenv(f"{provider.upper()}_API_BASE", raising=False)
|
||||
monkeypatch.setenv("TYPESAFE_API_BASE", "https://typesafe.test")
|
||||
monkeypatch.setenv("NIMBLE_API_BASE", "https://nimble-search.test")
|
||||
with pytest.raises(ValueError, match=f"{provider.upper()}_API_BASE"):
|
||||
oss_connection(provider)
|
||||
|
||||
|
||||
def test_oss_request_accepts_the_name_ollama_serves_nimble_under_only_for_bespoke(provider: OssDecisionProvider) -> None:
|
||||
body: Final = {"model": "nimble"}
|
||||
if provider == "bespoke":
|
||||
assert validate_oss_request(provider, body) == "nimble"
|
||||
return
|
||||
with pytest.raises(ValueError, match=f"{provider} model must be one of"):
|
||||
validate_oss_request(provider, body)
|
||||
|
|
@ -463,16 +463,22 @@ 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("provider,model", [
|
||||
("laya", "english"), ("laya", "multilingual"), ("laya", "typed-decisions"),
|
||||
("bespoke", "nimble-latest"), ("bespoke", "bespokelabs/Bespoke-Nimble-9B"),
|
||||
])
|
||||
@pytest.mark.parametrize("suffix", ["", "/"])
|
||||
def test_oss_native_model_uses_the_classifier_permission_identity(provider: str, model: str, suffix: str) -> None:
|
||||
assert get_model_from_request(
|
||||
request_data={"model": model}, route=f"/{provider}/v1/systemone{suffix}"
|
||||
) == f"{provider}/{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:
|
||||
@pytest.mark.parametrize("provider", ["laya", "bespoke"])
|
||||
@pytest.mark.parametrize("model", [None, "", "auto", "laya/english", "bespoke/nimble-latest", "unknown", ["english"], 7])
|
||||
def test_oss_native_model_cannot_implicitly_select_an_unauthorized_checkpoint(provider: str, model: object) -> None:
|
||||
with pytest.raises(HTTPException) as denied:
|
||||
get_model_from_request(request_data={"model": model}, route="/laya/v1/systemone")
|
||||
get_model_from_request(request_data={"model": model}, route=f"/{provider}/v1/systemone")
|
||||
assert denied.value.status_code == 400
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -7706,6 +7706,10 @@ class TestTeamMemberAutoRouterWrites:
|
|||
@pytest.mark.parametrize(
|
||||
"stored_provider,stored_base,supplied,expected_transport",
|
||||
[
|
||||
("bespoke", "https://decision.test", {"provider": "bespoke", "model": "nimble-latest"}, {"api_base": "https://decision.test", "api_key": "stored-secret"}),
|
||||
("bespoke", "https://decision.test", {"provider": "bespoke", "model": "nimble-latest", "api_base": "https://new.test"}, {}),
|
||||
("bespoke", "https://decision.test", {"provider": "laya", "model": "english"}, {}),
|
||||
("laya", "https://decision.test", {"provider": "bespoke", "model": "nimble-latest"}, {}),
|
||||
("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"}),
|
||||
(
|
||||
|
|
@ -7743,7 +7747,7 @@ class TestTeamMemberAutoRouterWrites:
|
|||
"model": "auto_router/complexity_router",
|
||||
"complexity_router_config": self._classifier_config(
|
||||
{
|
||||
"provider": stored_provider, "model": "english" if stored_provider == "laya" else "jev-latest",
|
||||
"provider": stored_provider, "model": {"laya": "english", "bespoke": "nimble-latest"}.get(stored_provider, "jev-latest"),
|
||||
"api_base": stored_base, "api_key": "stored-secret",
|
||||
},
|
||||
stored_legacy,
|
||||
|
|
|
|||
|
|
@ -144,6 +144,8 @@ def test_tier_config_is_normalized_and_unknown_router_extras_are_rejected() -> N
|
|||
({"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"),
|
||||
({"provider": "bespoke", "model": "nimble-latest", "api_base": "https://collector.invalid"}, "api_base"),
|
||||
({"provider": "bespoke", "model": "nimble-latest", "api_key": "sk-member"}, "api_key"),
|
||||
],
|
||||
)
|
||||
@pytest.mark.parametrize("legacy", [False, True])
|
||||
|
|
@ -162,7 +164,7 @@ 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}."
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("provider", "model"), [("typesafe", "jev-preview"), ("laya", "english")])
|
||||
@pytest.mark.parametrize(("provider", "model"), [("typesafe", "jev-preview"), ("laya", "english"), ("bespoke", "nimble-latest")])
|
||||
@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(
|
||||
|
|
@ -348,7 +350,7 @@ 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")])
|
||||
@pytest.mark.parametrize(("provider", "model"), [("typesafe", "jev-latest"), ("laya", "english"), ("bespoke", "nimble-latest")])
|
||||
async def test_jev_evaluation_requires_model_access_but_no_completion_deployment(
|
||||
catalog: Router, restricted: str | None, provider: str, model: str
|
||||
) -> None:
|
||||
|
|
@ -376,7 +378,7 @@ async def test_jev_evaluation_requires_model_access_but_no_completion_deployment
|
|||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("restricted", ["member", "project", "organization", None])
|
||||
@pytest.mark.parametrize(("provider", "model"), [("typesafe", "jev-latest"), ("laya", "english")])
|
||||
@pytest.mark.parametrize(("provider", "model"), [("typesafe", "jev-latest"), ("laya", "english"), ("bespoke", "nimble-latest")])
|
||||
async def test_jev_evaluation_obeys_each_containing_scope(
|
||||
catalog: Router, restricted: str | None, provider: str, model: str
|
||||
) -> None:
|
||||
|
|
|
|||
|
|
@ -142,60 +142,65 @@ def test_success_handler_dispatches_to_typesafe_handler():
|
|||
@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
|
||||
@pytest.mark.parametrize("provider,requested,routing_model", [
|
||||
("laya", "english", "multilingual"), ("laya", "english", None),
|
||||
("bespoke", "nimble-latest", None),
|
||||
("bespoke", "bespokelabs/Bespoke-Nimble-9B", None),
|
||||
])
|
||||
async def test_oss_gateway_accounts_for_checkpoint_usage_and_registered_cost(
|
||||
monkeypatch: pytest.MonkeyPatch, routing_model: str | None, metadata_slot: str, guardrail_cost: float,
|
||||
provider: str, requested: str
|
||||
) -> None:
|
||||
checkpoint: Final = routing_model or "english"
|
||||
model: Final = f"laya/{checkpoint}"
|
||||
checkpoint: Final = routing_model or requested
|
||||
model: Final = f"{provider}/{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",
|
||||
"litellm_provider": provider, "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={},
|
||||
model=requested, messages=[], stream=False, call_type="pass_through_endpoint",
|
||||
start_time=start, litellm_call_id="oss-accounting", function_id="oss-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",
|
||||
"type": "http", "method": "POST", "path": f"/{provider}/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"}},
|
||||
api_key="oss-budget-key", token="oss-budget-key",
|
||||
model_max_budget={f"{provider}/{requested}": {"budget_limit": 0.01, "time_period": "1d"}},
|
||||
)
|
||||
request_body: Final = {"model": "english", metadata_slot: {"model_group": "unbounded-client-choice"}}
|
||||
request_body: Final = {"model": requested, 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,
|
||||
passthrough_logging_payload={"url": f"https://{provider}.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={},
|
||||
model=requested, 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},
|
||||
"model": "laya-rl-agent" if provider == "laya" else requested, "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,
|
||||
httpx_response=httpx.Response(200, request=httpx.Request("POST", f"https://{provider}.test/v1/systemone"), json=body),
|
||||
response_body=body, request_body={"model": requested}, logging_obj=logging_obj,
|
||||
url_route=f"https://{provider}.test/v1/systemone", result="{}", start_time=start,
|
||||
end_time=datetime.now(), cache_hit=False, custom_llm_provider=provider, **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["model"], logged["custom_llm_provider"]) == (model, provider)
|
||||
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,
|
||||
|
|
@ -203,7 +208,7 @@ async def test_laya_gateway_accounts_for_checkpoint_usage_and_registered_cost(
|
|||
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"]["model_group"] == f"{provider}/{requested}"
|
||||
assert logged["standard_logging_object"]["response_cost"] == pytest.approx(expected_cost + guardrail_cost)
|
||||
|
||||
from litellm.caching.caching import DualCache
|
||||
|
|
@ -211,10 +216,10 @@ async def test_laya_gateway_accounts_for_checkpoint_usage_and_registered_cost(
|
|||
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")
|
||||
assert await budget_limiter.is_key_within_model_budget(auth, f"{provider}/{requested}")
|
||||
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")
|
||||
await budget_limiter.is_key_within_model_budget(auth, f"{provider}/{requested}")
|
||||
|
||||
|
||||
def test_openrouter_decisions_response_is_priced_from_request_model_registry_row():
|
||||
|
|
|
|||
|
|
@ -7298,6 +7298,8 @@ class TestTypeSafePassthroughRoute:
|
|||
"provider, endpoint, is_decision_request",
|
||||
(
|
||||
("typesafe", "systemone", True),
|
||||
("laya", "systemone", True),
|
||||
("bespoke", "systemone", True),
|
||||
("typesafe", "systemone/", True),
|
||||
("typesafe", "systemone?trace=1", True),
|
||||
("typesafe", "systemone/?trace=1", True),
|
||||
|
|
@ -7316,7 +7318,7 @@ class TestTypeSafePassthroughRoute:
|
|||
self,
|
||||
client: TestClient,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
provider: Literal["typesafe", "openrouter"],
|
||||
provider: Literal["typesafe", "openrouter", "laya", "bespoke"],
|
||||
endpoint: str,
|
||||
is_decision_request: bool,
|
||||
quota_scope: Literal["key", "project_output"],
|
||||
|
|
@ -7337,12 +7339,15 @@ class TestTypeSafePassthroughRoute:
|
|||
monkeypatch.setattr(proxy_server, "proxy_logging_obj", ProxyLogging(user_api_key_cache=cache))
|
||||
monkeypatch.setenv("OPENROUTER_API_KEY", "openrouter-test-key")
|
||||
monkeypatch.setenv("OPENROUTER_API_BASE", "https://typesafe.example/base")
|
||||
model: Final = "jev-latest" if provider == "typesafe" else "test-generative-model"
|
||||
monkeypatch.setenv("LAYA_API_BASE", "https://typesafe.example/base")
|
||||
monkeypatch.setenv("BESPOKE_API_BASE", "https://typesafe.example/base")
|
||||
model: Final = {"typesafe": "jev-latest", "laya": "english", "bespoke": "nimble-latest"}.get(provider, "test-generative-model")
|
||||
permission_model: Final = f"{provider}/{model}" if provider in ("laya", "bespoke") else model
|
||||
auth: Final = UserAPIKeyAuth(
|
||||
api_key="sk-limited",
|
||||
tpm_limit=token_limit if quota_scope == "key" else None,
|
||||
project_id="test-project" if quota_scope == "project_output" else None,
|
||||
project_metadata={"model_otpm_limit": {model: token_limit}} if quota_scope == "project_output" else {},
|
||||
project_metadata={"model_otpm_limit": {permission_model: token_limit}} if quota_scope == "project_output" else {},
|
||||
)
|
||||
monkeypatch.setitem(proxy_server.app.dependency_overrides, user_api_key_auth, lambda: auth)
|
||||
body: Final = (
|
||||
|
|
@ -7409,36 +7414,44 @@ class TestTypeSafePassthroughRoute:
|
|||
)
|
||||
|
||||
|
||||
class TestLayaPassthroughRoute:
|
||||
@pytest.mark.parametrize("provider", ["laya", "bespoke"])
|
||||
class TestOssDecisionPassthroughRoute:
|
||||
@pytest.fixture
|
||||
def client(self, monkeypatch: pytest.MonkeyPatch) -> Iterator[TestClient]:
|
||||
def checkpoint(self, provider: str) -> str:
|
||||
return "english" if provider == "laya" else "nimble-latest"
|
||||
|
||||
@pytest.fixture
|
||||
def client(self, monkeypatch: pytest.MonkeyPatch, provider: str) -> Iterator[TestClient]:
|
||||
from litellm.proxy.proxy_server import app
|
||||
|
||||
monkeypatch.setenv("LAYA_API_BASE", "http://laya.test/base")
|
||||
monkeypatch.setenv(f"{provider.upper()}_API_BASE", f"http://{provider}.test/base")
|
||||
monkeypatch.setenv("TYPESAFE_API_KEY", "never-send-typesafe-key")
|
||||
monkeypatch.delenv("LAYA_API_KEY", raising=False)
|
||||
monkeypatch.delenv(f"{provider.upper()}_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
|
||||
@pytest.mark.parametrize("api_key", [None, "oss-provider-key"])
|
||||
def test_oss_forwards_native_decisions_without_gateway_or_typesafe_credentials(
|
||||
self, client: TestClient, monkeypatch: pytest.MonkeyPatch, api_key: str | None, provider: str, checkpoint: str
|
||||
) -> None:
|
||||
if api_key is not None:
|
||||
monkeypatch.setenv("LAYA_API_KEY", api_key)
|
||||
monkeypatch.setenv(f"{provider.upper()}_API_KEY", api_key)
|
||||
body: Final = {
|
||||
"model": "english",
|
||||
"model": checkpoint,
|
||||
"state": "refund",
|
||||
"questions": {"department": {"type": "choice", "criteria": {"billing": "refunds"}}},
|
||||
}
|
||||
answer: Final = {"model": "laya-rl-agent", "routing": {"model": "english"}, "answers": {}}
|
||||
answer: Final = {
|
||||
"model": "laya-rl-agent" if provider == "laya" else checkpoint, "answers": {},
|
||||
**({"routing": {"model": checkpoint}} if provider == "laya" else {}),
|
||||
}
|
||||
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)
|
||||
route: Final = upstream.post(f"http://{provider}.test/base/v1/systemone?trace=yes").respond(200, json=answer)
|
||||
response: Final = client.post(
|
||||
"/laya/v1/systemone?trace=yes",
|
||||
f"/{provider}/v1/systemone?trace=yes",
|
||||
json=body,
|
||||
headers={"Authorization": "Bearer sk-virtual", "x-pass-authorization": "Bearer attacker"},
|
||||
)
|
||||
|
|
@ -7448,26 +7461,26 @@ class TestLayaPassthroughRoute:
|
|||
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
|
||||
def test_oss_missing_server_fails_without_contacting_another_provider(
|
||||
self, client: TestClient, monkeypatch: pytest.MonkeyPatch, provider: str, checkpoint: str
|
||||
) -> None:
|
||||
monkeypatch.delenv("LAYA_API_BASE")
|
||||
monkeypatch.delenv(f"{provider.upper()}_API_BASE")
|
||||
with respx.mock(assert_all_called=False) as upstream:
|
||||
response: Final = client.post("/laya/v1/systemone", json={"model": "english"})
|
||||
response: Final = client.post(f"/{provider}/v1/systemone", json={"model": checkpoint})
|
||||
assert response.status_code == 503
|
||||
assert "LAYA_API_BASE" in response.text
|
||||
assert f"{provider.upper()}_API_BASE" in response.text
|
||||
assert len(upstream.calls) == 0
|
||||
|
||||
def test_laya_does_not_forward_unsupported_endpoints(self, client: TestClient) -> None:
|
||||
def test_oss_does_not_forward_unsupported_endpoints(self, client: TestClient, provider: str, checkpoint: str) -> None:
|
||||
with respx.mock(assert_all_called=False) as upstream:
|
||||
response: Final = client.post("/laya/v1/evaluate", json={"model": "english"})
|
||||
response: Final = client.post(f"/{provider}/v1/evaluate", json={"model": checkpoint})
|
||||
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:
|
||||
def test_oss_rejects_implicit_checkpoint_selection(self, client: TestClient, model: str | None, provider: str) -> None:
|
||||
with respx.mock(assert_all_called=False) as upstream:
|
||||
response: Final = client.post("/laya/v1/systemone", json={"model": model})
|
||||
response: Final = client.post(f"/{provider}/v1/systemone", json={"model": model})
|
||||
assert response.status_code == 400
|
||||
assert len(upstream.calls) == 0
|
||||
|
||||
|
|
@ -7475,19 +7488,19 @@ class TestLayaPassthroughRoute:
|
|||
"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]
|
||||
def test_oss_rejects_controls_that_change_authorized_body_or_usage_accounting(
|
||||
self, client: TestClient, controls: Mapping[str, object], provider: str, checkpoint: str
|
||||
) -> 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})
|
||||
route: Final = upstream.post(f"http://{provider}.test/base/v1/systemone").respond(200, json={"answers": {}})
|
||||
response: Final = client.post(f"/{provider}/v1/systemone", json={"model": checkpoint, **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
|
||||
def test_oss_hooks_enforce_canonical_model_limits_and_keep_native_wire_body(
|
||||
self, client: TestClient, monkeypatch: pytest.MonkeyPatch, metadata_slot: str, provider: str, checkpoint: str
|
||||
) -> None:
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3
|
||||
|
|
@ -7497,7 +7510,7 @@ class TestLayaPassthroughRoute:
|
|||
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}},
|
||||
api_key="oss-native-rpm", metadata={"model_rpm_limit": {f"{provider}/{checkpoint}": 1}},
|
||||
)
|
||||
def authenticated_key() -> UserAPIKeyAuth:
|
||||
return auth
|
||||
|
|
@ -7509,7 +7522,7 @@ class TestLayaPassthroughRoute:
|
|||
self, user_api_key_dict: UserAPIKeyAuth, cache: DualCache,
|
||||
data: dict[str, object], call_type: CallTypesLiteral,
|
||||
) -> dict[str, object]:
|
||||
assert data["model"] == "laya/english"
|
||||
assert data["model"] == f"{provider}/{checkpoint}"
|
||||
metadata: Final = data.get(metadata_slot)
|
||||
assert isinstance(metadata, dict)
|
||||
assert "standard_logging_guardrail_information" not in metadata
|
||||
|
|
@ -7519,40 +7532,42 @@ class TestLayaPassthroughRoute:
|
|||
|
||||
monkeypatch.setattr(litellm, "callbacks", [LimitHook()])
|
||||
body: Final = {
|
||||
"model": "english", "state": "refund",
|
||||
"model": checkpoint, "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)
|
||||
route: Final = upstream.post(f"http://{provider}.test/base/v1/systemone").respond(200, json={"answers": {}})
|
||||
first: Final = client.post(f"/{provider}/v1/systemone", json=body)
|
||||
second: Final = client.post(f"/{provider}/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"}
|
||||
assert json.loads(route.calls.last.request.content) == {"model": checkpoint, "state": "refund"}
|
||||
|
||||
def test_laya_preserves_trusted_hook_checkpoint_changes(
|
||||
self, client: TestClient, monkeypatch: pytest.MonkeyPatch
|
||||
def test_oss_preserves_trusted_hook_checkpoint_changes(
|
||||
self, client: TestClient, monkeypatch: pytest.MonkeyPatch, provider: str, checkpoint: str
|
||||
) -> None:
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
changed_checkpoint: Final = "multilingual" if provider == "laya" else "bespokelabs/Bespoke-Nimble-9B"
|
||||
|
||||
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"}
|
||||
assert data["model"] == f"{provider}/{checkpoint}"
|
||||
return {**data, "model": f"{provider}/{changed_checkpoint}"}
|
||||
|
||||
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"})
|
||||
route: Final = upstream.post(f"http://{provider}.test/base/v1/systemone").respond(200, json={"answers": {}})
|
||||
response: Final = client.post(f"/{provider}/v1/systemone", json={"model": checkpoint, "state": "refund"})
|
||||
assert response.status_code == 200, response.text
|
||||
assert json.loads(route.calls.last.request.content) == {"model": "multilingual", "state": "refund"}
|
||||
assert json.loads(route.calls.last.request.content) == {"model": changed_checkpoint, "state": "refund"}
|
||||
|
||||
|
||||
class TestFalAIPassthroughRoute:
|
||||
|
|
|
|||
|
|
@ -451,7 +451,7 @@ def test_jev_config_requires_classifier_config() -> None:
|
|||
)
|
||||
@pytest.mark.parametrize(
|
||||
("provider", "model", "canonical_provider"),
|
||||
[(None, "jev-latest", "jev"), ("typesafe", "jev-latest", "jev"), ("jev", "jev-latest", "jev"), ("laya", "english", "laya")],
|
||||
[(None, "jev-latest", "jev"), ("typesafe", "jev-latest", "jev"), ("jev", "jev-latest", "jev"), ("laya", "english", "laya"), ("bespoke", "nimble-latest", "bespoke")],
|
||||
)
|
||||
def test_classifier_aliases_load_and_serialize_one_canonical_config(
|
||||
classifier_type: str, config_key: str, provider: str | None, model: str, canonical_provider: str
|
||||
|
|
@ -474,45 +474,47 @@ def test_classifier_aliases_load_and_serialize_one_canonical_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.parametrize("provider", ["laya", "bespoke"])
|
||||
@pytest.mark.parametrize("model", [None, " "])
|
||||
def test_oss_requires_its_own_checkpoint(provider: str, model: str | None) -> None:
|
||||
with pytest.raises(ValueError, match=f"{provider} model must be"):
|
||||
JevClassifierConfig.model_validate({"provider": provider, **({"model": model} if model is not None else {})})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("provider,model", [("laya", "english"), ("bespoke", "nimble-latest")])
|
||||
@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
|
||||
async def test_oss_routes_with_its_own_credentials_and_accounts_the_checkpoint(
|
||||
monkeypatch: pytest.MonkeyPatch, custom_base: bool, legacy: bool, provider: str, model: str
|
||||
) -> 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.setenv(f"{provider.upper()}_API_BASE", f"https://{provider}.test")
|
||||
monkeypatch.setenv(f"{provider.upper()}_API_KEY", "oss-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.setitem(litellm.model_cost, f"{provider}/{model}", {"input_cost_per_token": 0.01})
|
||||
recorder: Final = _UsageRecorder(f"{provider}/{model}")
|
||||
monkeypatch.setattr(litellm, "_async_success_callback", [recorder])
|
||||
router: Final = ComplexityRouter(
|
||||
"laya-route",
|
||||
f"{provider}-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 {}),
|
||||
"provider": provider,
|
||||
"model": model,
|
||||
**({"api_base": f"https://{provider}.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(
|
||||
route: Final = upstream.post(f"https://{provider}.test/v1/systemone").respond(
|
||||
200,
|
||||
json={
|
||||
"model": "laya-rl-agent",
|
||||
"routing": {"model": "english"},
|
||||
"model": "laya-rl-agent" if provider == "laya" else model,
|
||||
**({"routing": {"model": model}} if provider == "laya" else {}),
|
||||
"answers": {"tier": _answer().model_dump()},
|
||||
"usage": {"input_tokens": 31, "output_tokens": 0},
|
||||
},
|
||||
|
|
@ -522,11 +524,11 @@ async def test_laya_routes_with_its_own_credentials_and_accounts_the_checkpoint(
|
|||
|
||||
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.jev_verdict.provider, outcome.jev_verdict.model) == (provider, model)
|
||||
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 sent.headers.get("authorization") == (None if custom_base else "Bearer oss-env-key")
|
||||
assert json.loads(sent.content)["model"] == model
|
||||
assert len(recorder.calls) == 1
|
||||
assert recorder.calls[0]["response_cost"] == pytest.approx(0.31)
|
||||
|
||||
|
|
|
|||
|
|
@ -39,6 +39,7 @@ SEMANTIC_FIELDS = frozenset({"auto_router_config", "auto_router_default_model",
|
|||
("typesafe", "jev-preview", "typesafe"),
|
||||
("jev", "jev-preview", "typesafe"),
|
||||
("laya", "english", "laya"),
|
||||
("bespoke", "nimble-latest", "bespoke"),
|
||||
],
|
||||
)
|
||||
def test_open_source_classifier_enumerates_its_accounting_model(
|
||||
|
|
|
|||
|
|
@ -150,7 +150,7 @@ const AutoRouterClassifierTabs: React.FC<AutoRouterClassifierTabsProps> = ({ val
|
|||
if (next === "jev") changeType("jev");
|
||||
};
|
||||
const changeProvider = (provider: unknown) => {
|
||||
if (provider !== "jev" && provider !== "laya") return;
|
||||
if (provider !== "jev" && provider !== "laya" && provider !== "bespoke") return;
|
||||
const defaults = defaultJevClassifierConfig(provider);
|
||||
onChange({
|
||||
...value,
|
||||
|
|
@ -173,7 +173,7 @@ const AutoRouterClassifierTabs: React.FC<AutoRouterClassifierTabsProps> = ({ val
|
|||
{[
|
||||
{ value: "heuristics", label: "Heuristics", description: "Classify locally, with no API call" },
|
||||
{ value: "llm", label: "LLM", description: "Use a judge model to choose a solver" },
|
||||
{ value: "jev", label: "OSS Classifier", description: "Use Jev or Laya to choose a tier" },
|
||||
{ value: "jev", label: "OSS Classifier", description: "Use Jev, Laya, or Bespoke Nimble to choose a tier" },
|
||||
].map((option) => (
|
||||
<Label
|
||||
key={option.value}
|
||||
|
|
@ -214,6 +214,10 @@ const AutoRouterClassifierTabs: React.FC<AutoRouterClassifierTabsProps> = ({ val
|
|||
<RadioGroupItem value="laya" />
|
||||
Laya
|
||||
</Label>
|
||||
<Label>
|
||||
<RadioGroupItem value="bespoke" />
|
||||
Bespoke Nimble
|
||||
</Label>
|
||||
</RadioGroup>
|
||||
</fieldset>
|
||||
)}
|
||||
|
|
|
|||
|
|
@ -97,9 +97,13 @@ function Form() {
|
|||
|
||||
describe("JEV classifier editor", () => {
|
||||
afterEach(() => vi.mocked(useAuthorized).mockReset());
|
||||
it.each(["jev", "laya"] as const)(
|
||||
it.each([
|
||||
["jev", "Jev", "jev-test"],
|
||||
["laya", "Laya", "multilingual"],
|
||||
["bespoke", "Bespoke Nimble", "bespokelabs/Bespoke-Nimble-9B"],
|
||||
] as const)(
|
||||
"preserves %s, custom tiers and context through save, reload and probe",
|
||||
async (provider) => {
|
||||
async (provider, label, model) => {
|
||||
renderWithProviders(<Form />);
|
||||
expect(screen.getByLabelText("Judge model")).toBeInTheDocument();
|
||||
expect(screen.getByText("Reasoning Effort")).toBeInTheDocument();
|
||||
|
|
@ -117,12 +121,13 @@ describe("JEV classifier editor", () => {
|
|||
expect(screen.getByLabelText("Classifier Model")).toHaveTextContent("english");
|
||||
fireEvent.click(screen.getByRole("radio", { name: "Jev" }));
|
||||
expect(screen.getByLabelText("Classifier Model")).toHaveValue("jev-latest");
|
||||
if (provider === "laya") {
|
||||
fireEvent.click(screen.getByRole("radio", { name: "Laya" }));
|
||||
fireEvent.click(screen.getByRole("radio", { name: label }));
|
||||
if (provider === "bespoke") expect(screen.getByLabelText("Classifier Model")).toHaveTextContent("nimble-latest");
|
||||
if (provider !== "jev") {
|
||||
await userEvent.click(screen.getByLabelText("Classifier Model"));
|
||||
await userEvent.click(screen.getByRole("option", { name: "multilingual" }));
|
||||
await userEvent.click(screen.getByRole("option", { name: model }));
|
||||
} else {
|
||||
fireEvent.change(screen.getByLabelText("Classifier Model"), { target: { value: "jev-test" } });
|
||||
fireEvent.change(screen.getByLabelText("Classifier Model"), { target: { value: model } });
|
||||
}
|
||||
fireEvent.change(screen.getByLabelText("Classifier Timeout (ms)"), { target: { value: "4200" } });
|
||||
fireEvent.change(screen.getByLabelText("Context Window Size"), { target: { value: "6" } });
|
||||
|
|
@ -131,9 +136,9 @@ describe("JEV classifier editor", () => {
|
|||
fireEvent.click(screen.getByRole("button", { name: "Customize tiers" }));
|
||||
fireEvent.click(screen.getByRole("button", { name: "Save and reload" }));
|
||||
expect(screen.getByRole("radio", { name: /^OSS Classifier$/ })).toBeChecked();
|
||||
expect(screen.getByRole("radio", { name: provider === "laya" ? "Laya" : "Jev" })).toBeChecked();
|
||||
if (provider === "laya") expect(screen.getByLabelText("Classifier Model")).toHaveTextContent("multilingual");
|
||||
else expect(screen.getByLabelText("Classifier Model")).toHaveValue("jev-test");
|
||||
expect(screen.getByRole("radio", { name: label })).toBeChecked();
|
||||
if (provider !== "jev") expect(screen.getByLabelText("Classifier Model")).toHaveTextContent(model);
|
||||
else expect(screen.getByLabelText("Classifier Model")).toHaveValue(model);
|
||||
expect(screen.getByLabelText("Classifier Timeout (ms)")).toHaveValue(4200);
|
||||
expect(screen.getByLabelText("Context Window Size")).toHaveValue("6");
|
||||
expect(screen.getByRole("switch", { name: "Classifier circuit breaker" })).not.toBeChecked();
|
||||
|
|
@ -145,7 +150,7 @@ describe("JEV classifier editor", () => {
|
|||
classifier_type: "oss_classifier",
|
||||
opensource_classifier_config: {
|
||||
provider,
|
||||
model: provider === "laya" ? "multilingual" : "jev-test",
|
||||
model,
|
||||
timeout_ms: 4200,
|
||||
circuit_breaker_enabled: false,
|
||||
circuit_breaker_cooldown_seconds: 50,
|
||||
|
|
|
|||
|
|
@ -7,7 +7,14 @@ import { Textarea } from "@/components/ui/textarea";
|
|||
import ClassifierCircuitBreakerConfig from "./ClassifierCircuitBreakerConfig";
|
||||
import type { ComplexityRouterConfigValue } from "./ComplexityRouterConfig";
|
||||
import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select";
|
||||
import { defaultJevClassifierConfig, LAYA_MODELS } from "./jev_classifier_config";
|
||||
import { defaultJevClassifierConfig, OSS_CLASSIFIER_MODELS } from "./jev_classifier_config";
|
||||
|
||||
const providerDescriptions = {
|
||||
jev: "Uses TypeSafe System One Choice evaluation with your configured tiers",
|
||||
laya: "Uses Laya with your configured tiers. Set LAYA_API_BASE on the gateway to connect your Laya server.",
|
||||
bespoke:
|
||||
"Uses Bespoke Nimble with your configured tiers. Set BESPOKE_API_BASE on the gateway to connect your Nimble server.",
|
||||
};
|
||||
|
||||
export default function JevClassifierConfig({
|
||||
value,
|
||||
|
|
@ -18,26 +25,22 @@ export default function JevClassifierConfig({
|
|||
}) {
|
||||
const id = useId();
|
||||
const config = value.jev_classifier_config ?? defaultJevClassifierConfig();
|
||||
const isLaya = config.provider === "laya";
|
||||
const models = config.provider && config.provider !== "jev" ? OSS_CLASSIFIER_MODELS[config.provider] : undefined;
|
||||
const update = (patch: Partial<typeof config>) =>
|
||||
onChange({ ...value, jev_classifier_config: { ...config, ...patch } });
|
||||
|
||||
return (
|
||||
<div className="mt-4 space-y-3">
|
||||
<p className="text-sm text-muted-foreground">
|
||||
{isLaya
|
||||
? "Uses Laya with your configured tiers. Set LAYA_API_BASE on the gateway to connect your Laya server."
|
||||
: "Uses TypeSafe System One Choice evaluation with your configured tiers"}
|
||||
</p>
|
||||
<p className="text-sm text-muted-foreground">{providerDescriptions[config.provider ?? "jev"]}</p>
|
||||
<div>
|
||||
<Label htmlFor={`${id}-model`}>Classifier Model</Label>
|
||||
{isLaya ? (
|
||||
{models ? (
|
||||
<Select value={config.model} onValueChange={(model) => model && update({ model })}>
|
||||
<SelectTrigger id={`${id}-model`} className="w-full">
|
||||
<SelectValue />
|
||||
</SelectTrigger>
|
||||
<SelectContent>
|
||||
{LAYA_MODELS.map((model) => (
|
||||
{models.map((model) => (
|
||||
<SelectItem key={model} value={model}>
|
||||
{model}
|
||||
</SelectItem>
|
||||
|
|
|
|||
|
|
@ -43,10 +43,19 @@ describe("buildAutoRouterRoutingTestRequest", () => {
|
|||
expect(request?.complexity_router_config.opensource_classifier_config).not.toHaveProperty("api_key");
|
||||
expect(request?.complexity_router_config.opensource_classifier_config).not.toHaveProperty("api_base");
|
||||
});
|
||||
it.each(["object", "json"])("probes saved Laya %s configuration with custom tiers and team context", (format) => {
|
||||
it.each([
|
||||
["jev", "jev-latest", "object"],
|
||||
["jev", "jev-latest", "json"],
|
||||
["laya", "english", "object"],
|
||||
["laya", "english", "json"],
|
||||
["bespoke", "nimble-latest", "object"],
|
||||
["bespoke", "nimble-latest", "json"],
|
||||
["bespoke", "bespokelabs/Bespoke-Nimble-9B", "object"],
|
||||
["bespoke", "bespokelabs/Bespoke-Nimble-9B", "json"],
|
||||
])("probes saved %s/%s %s configuration with custom tiers and team context", (provider, model, format) => {
|
||||
const config = {
|
||||
classifier_type: "oss_classifier",
|
||||
opensource_classifier_config: { provider: "laya", model: "english", timeout_ms: 900 },
|
||||
opensource_classifier_config: { provider, model, timeout_ms: 900 },
|
||||
tiers: { QUICK: ["fast"], DEEP: ["strong"] },
|
||||
tier_definitions: { QUICK: "Simple questions", DEEP: "Complex questions" },
|
||||
fallback_tier: "DEEP",
|
||||
|
|
@ -62,6 +71,14 @@ describe("buildAutoRouterRoutingTestRequest", () => {
|
|||
buildSavedJevConnectionTestRequest(format === "json" ? JSON.stringify(config) : config, "saved-id", "team-1"),
|
||||
).toEqual(expectedRequest);
|
||||
});
|
||||
it.each(["laya", "bespoke"])("does not probe unsupported %s models", (provider) => {
|
||||
const config = {
|
||||
classifier_type: "oss_classifier",
|
||||
opensource_classifier_config: { provider, model: "unsupported" },
|
||||
tiers: CONFIG.tiers,
|
||||
};
|
||||
expect(buildSavedJevConnectionTestRequest(config, "saved-id")).toBeUndefined();
|
||||
});
|
||||
it.each([undefined, null, "not json", "[]", {}, { classifier_type: "llm", tiers: {} }, { classifier_type: "jev" }])(
|
||||
"does not build a JEV probe for invalid or other classifier configurations: %j",
|
||||
(config) => {
|
||||
|
|
|
|||
|
|
@ -65,6 +65,7 @@ describe("buildComplexityRouterConfig", () => {
|
|||
{ model: "" },
|
||||
{ model: " " },
|
||||
{ provider: "laya" as const, model: "unsupported" },
|
||||
{ provider: "bespoke" as const, model: "unsupported" },
|
||||
{ timeout_ms: 0 },
|
||||
{ timeout_ms: 1.5 },
|
||||
{ timeout_ms: Number.NaN },
|
||||
|
|
@ -79,13 +80,35 @@ describe("buildComplexityRouterConfig", () => {
|
|||
).toBe("Enter a valid classifier model, a positive whole-number timeout and a positive cooldown");
|
||||
});
|
||||
|
||||
it.each([false, true])("serializes Laya with shared context and no LLM config, custom tiers: %s", (custom) => {
|
||||
it.each([
|
||||
["bespoke", "nimble-latest"],
|
||||
["bespoke", "nimble"],
|
||||
["bespoke", "bespokelabs/Bespoke-Nimble-9B"],
|
||||
["jev", "custom-jev-model"],
|
||||
[undefined, "custom-jev-model"],
|
||||
] as const)("accepts %s model %s before saving or testing", (provider, model) => {
|
||||
expect(
|
||||
getClassifierModelError({
|
||||
classifier_type: "jev",
|
||||
jev_classifier_config: { provider, model, timeout_ms: 3000 },
|
||||
}),
|
||||
).toBeNull();
|
||||
});
|
||||
|
||||
it.each([
|
||||
["jev", "jev-latest", false],
|
||||
["jev", "jev-latest", true],
|
||||
["laya", "english", false],
|
||||
["laya", "english", true],
|
||||
["bespoke", "nimble-latest", false],
|
||||
["bespoke", "nimble-latest", true],
|
||||
] as const)("serializes %s/%s with shared context and no LLM config, custom tiers: %s", (provider, model, custom) => {
|
||||
const params: BuildComplexityRouterConfigParams = {
|
||||
...baseParams,
|
||||
classifierType: "jev",
|
||||
jevClassifierConfig: {
|
||||
provider: "laya",
|
||||
model: "english",
|
||||
provider,
|
||||
model,
|
||||
timeout_ms: 4500,
|
||||
instructions: " Choose the configured tier ",
|
||||
circuit_breaker_enabled: false,
|
||||
|
|
@ -112,8 +135,8 @@ describe("buildComplexityRouterConfig", () => {
|
|||
const config = buildComplexityRouterConfig(params);
|
||||
expect(config.classifier_type).toBe("oss_classifier");
|
||||
const expectedJevConfig = {
|
||||
provider: "laya",
|
||||
model: "english",
|
||||
provider,
|
||||
model,
|
||||
timeout_ms: 4500,
|
||||
instructions: "Choose the configured tier",
|
||||
circuit_breaker_enabled: false,
|
||||
|
|
|
|||
|
|
@ -1,10 +1,16 @@
|
|||
import { z } from "zod";
|
||||
import type { ClassifierType } from "./classifier_types";
|
||||
|
||||
export const LAYA_MODELS = ["english", "multilingual", "typed-decisions"] as const;
|
||||
export const OSS_CLASSIFIER_MODELS = {
|
||||
laya: ["english", "multilingual", "typed-decisions"],
|
||||
bespoke: ["nimble-latest", "nimble", "bespokelabs/Bespoke-Nimble-9B"],
|
||||
} as const;
|
||||
|
||||
const jevClassifierConfigFields = {
|
||||
provider: z.preprocess((value) => (value === "typesafe" ? "jev" : value), z.enum(["jev", "laya"]).optional()),
|
||||
provider: z.preprocess(
|
||||
(value) => (value === "typesafe" ? "jev" : value),
|
||||
z.enum(["jev", "laya", "bespoke"]).optional(),
|
||||
),
|
||||
model: z.string().trim().min(1).optional(),
|
||||
timeout_ms: z.number().int().positive().default(3000),
|
||||
instructions: z
|
||||
|
|
@ -19,16 +25,21 @@ export const jevClassifierConfigSchema = z
|
|||
.object(jevClassifierConfigFields)
|
||||
.transform((config) => ({
|
||||
...config,
|
||||
model: config.model ?? (config.provider === "laya" ? "english" : "jev-latest"),
|
||||
model:
|
||||
config.model ??
|
||||
(config.provider && config.provider !== "jev" ? OSS_CLASSIFIER_MODELS[config.provider][0] : "jev-latest"),
|
||||
}))
|
||||
.refine((config) => config.provider !== "laya" || LAYA_MODELS.some((model) => model === config.model), {
|
||||
message: "Select a supported Laya model",
|
||||
path: ["model"],
|
||||
});
|
||||
.refine(
|
||||
(config) =>
|
||||
!config.provider ||
|
||||
config.provider === "jev" ||
|
||||
OSS_CLASSIFIER_MODELS[config.provider].some((model) => model === config.model),
|
||||
{ message: "Select a supported classifier model", path: ["model"] },
|
||||
);
|
||||
|
||||
export type JevClassifierConfig = z.infer<typeof jevClassifierConfigSchema>;
|
||||
|
||||
export const defaultJevClassifierConfig = (provider: "jev" | "laya" = "jev"): JevClassifierConfig =>
|
||||
export const defaultJevClassifierConfig = (provider: JevClassifierConfig["provider"] = "jev"): JevClassifierConfig =>
|
||||
jevClassifierConfigSchema.parse({ provider });
|
||||
|
||||
export const hydrateOssClassifier = (config: {
|
||||
|
|
|
|||
47
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
47
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -1936,6 +1936,23 @@ export interface paths {
|
|||
patch: operations["bedrock_proxy_route_bedrock__endpoint__patch"];
|
||||
trace?: never;
|
||||
};
|
||||
"/bespoke/v1/systemone": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
get?: never;
|
||||
put?: never;
|
||||
/** Bespoke Proxy Route */
|
||||
post: operations["bespoke_proxy_route_bespoke_v1_systemone_post"];
|
||||
delete?: never;
|
||||
options?: never;
|
||||
head?: never;
|
||||
patch?: never;
|
||||
trace?: never;
|
||||
};
|
||||
"/budget/delete": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
|
|
@ -39149,12 +39166,12 @@ export interface components {
|
|||
OpenSourceClassifierConfig: {
|
||||
/**
|
||||
* Api Base
|
||||
* @description Provider API base; defaults to TYPESAFE_API_BASE or LAYA_API_BASE for the selected provider
|
||||
* @description Provider API base; defaults to the selected provider API_BASE environment variable
|
||||
*/
|
||||
api_base?: string | null;
|
||||
/**
|
||||
* Api Key
|
||||
* @description Provider API key; optional for self-hosted Laya
|
||||
* @description Provider API key; optional for self-hosted providers
|
||||
*/
|
||||
api_key?: string | null;
|
||||
/**
|
||||
|
|
@ -39169,7 +39186,7 @@ export interface components {
|
|||
circuit_breaker_enabled: boolean;
|
||||
/**
|
||||
* Instructions
|
||||
* @description Replaces the built-in Jev question instructions
|
||||
* @description Replaces the built-in classification instructions
|
||||
*/
|
||||
instructions?: string | null;
|
||||
/**
|
||||
|
|
@ -39182,7 +39199,7 @@ export interface components {
|
|||
* @default jev
|
||||
* @enum {string}
|
||||
*/
|
||||
provider: "jev" | "laya";
|
||||
provider: "jev" | "laya" | "bespoke";
|
||||
/**
|
||||
* Timeout Ms
|
||||
* @default 3000
|
||||
|
|
@ -41989,7 +42006,7 @@ 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 'oss_classifier', a structured choice call using Jev or Laya
|
||||
* @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, Laya or Bespoke Nimble
|
||||
* @default heuristic
|
||||
* @enum {string}
|
||||
*/
|
||||
|
|
@ -53033,6 +53050,26 @@ export interface operations {
|
|||
};
|
||||
};
|
||||
};
|
||||
bespoke_proxy_route_bespoke_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;
|
||||
};
|
||||
};
|
||||
};
|
||||
};
|
||||
delete_budget_budget_delete_post: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue