diff --git a/litellm/__init__.py b/litellm/__init__.py index a7c8f0dd5eb..c965968bb39 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -1687,6 +1687,9 @@ if TYPE_CHECKING: from .llms.strands_decider.decisions.transformation import ( StrandsDeciderDecisionsConfig as StrandsDeciderDecisionsConfig, ) + from .llms.databricks.decisions.transformation import ( + DatabricksDecisionsConfig as DatabricksDecisionsConfig, + ) from .llms.hosted_vllm.decisions.transformation import ( HostedVLLMDecisionsConfig as HostedVLLMDecisionsConfig, ) diff --git a/litellm/_lazy_imports_registry.py b/litellm/_lazy_imports_registry.py index 77a45a18db2..f70caf8bb0d 100644 --- a/litellm/_lazy_imports_registry.py +++ b/litellm/_lazy_imports_registry.py @@ -162,6 +162,7 @@ LLM_CONFIG_NAMES: Final = ( "OpenRouterDecisionsConfig", "CloudflareDecisionsConfig", "StrandsDeciderDecisionsConfig", + "DatabricksDecisionsConfig", "HostedVLLMDecisionsConfig", "OpenAIDecisionsConfig", "NvidiaNimRerankConfig", @@ -724,6 +725,7 @@ _LLM_CONFIGS_IMPORT_MAP: Final = { ".llms.strands_decider.decisions.transformation", "StrandsDeciderDecisionsConfig", ), + "DatabricksDecisionsConfig": (".llms.databricks.decisions.transformation", "DatabricksDecisionsConfig"), "HostedVLLMDecisionsConfig": ( ".llms.hosted_vllm.decisions.transformation", "HostedVLLMDecisionsConfig", diff --git a/litellm/decisions/main.py b/litellm/decisions/main.py index 335980bed2b..4a4af2b580e 100644 --- a/litellm/decisions/main.py +++ b/litellm/decisions/main.py @@ -127,7 +127,6 @@ def _prepare_call( api_key=api_key, ) provider_config: Final = _provider_config(upstream_model, provider) - canonical_model: Final = provider_config.canonical_model(upstream_model) if not upstream_model: raise litellm.BadRequestError( message="A model name is required for the Decisions API", @@ -140,6 +139,10 @@ def _prepare_call( model=model, llm_provider=provider, ) + try: + canonical_model: Final = provider_config.canonical_model(upstream_model) + except ValueError as error: + raise litellm.BadRequestError(message=str(error), model=model, llm_provider=provider) from error try: request: Final = _validate_request( state=state, questions=questions, decision_input=decision_input, safety_identifier=safety_identifier diff --git a/litellm/llms/databricks/decisions/transformation.py b/litellm/llms/databricks/decisions/transformation.py new file mode 100644 index 00000000000..a454da83e50 --- /dev/null +++ b/litellm/llms/databricks/decisions/transformation.py @@ -0,0 +1,57 @@ +import re +from collections.abc import Mapping +from dataclasses import dataclass, field +from types import MappingProxyType +from typing import Final + +from litellm.llms.base_llm.decisions.transformation import BaseDecisionsConfig + +_SERVING_ENDPOINT_NAME: Final = re.compile(r"[A-Za-z0-9_-][A-Za-z0-9._-]*") + + +def validate_serving_endpoint_name(name: str) -> str: + if _SERVING_ENDPOINT_NAME.fullmatch(name) is None: + raise ValueError( + f"Databricks serving endpoint name {name!r} must be the bare endpoint name (letters, digits, '-', '_' " + "and '.', with no '/', '?', '#', spaces, or a leading '.'), e.g. databricks-openjev-qwen35-4b" + ) + return name + + +@dataclass(frozen=True, slots=True) +class DatabricksDecisionsConnection: + api_base: str + api_key: str = field(repr=False) + + +class DatabricksDecisionsConfig(BaseDecisionsConfig): + api_key_env = ("DATABRICKS_API_KEY", "DATABRICKS_TOKEN") + api_base_env = ("DATABRICKS_API_BASE",) + + def missing_api_base_message(self, custom_llm_provider: str) -> str: + return ( + f"api_base is required for Decisions provider '{custom_llm_provider}': set DATABRICKS_API_BASE to " + "https:///serving-endpoints" + ) + + def canonical_model(self, model: str) -> str: + return validate_serving_endpoint_name(model) + + def get_complete_url(self, api_base: str, model: str) -> str: + return f"{api_base.rstrip('/')}/{validate_serving_endpoint_name(model)}/invocations" + + def classifier_response(self, body: Mapping[str, object], requested_model: str) -> Mapping[str, object]: + return MappingProxyType({**body, "model": requested_model}) + + def connection(self, api_base: str | None, api_key: str | None) -> DatabricksDecisionsConnection: + base: Final = self.resolve_api_base(api_base) + key: Final = api_key if api_base is not None else self.resolve_api_key(api_key) + if not base or not key: + raise ValueError( + "Databricks requires api_key or DATABRICKS_API_KEY (or DATABRICKS_TOKEN) and api_base or " + "DATABRICKS_API_BASE pointing to https:///serving-endpoints" + ) + return DatabricksDecisionsConnection(api_base=base.rstrip("/"), api_key=key) + + +DATABRICKS_DECISIONS_CONFIG: Final[DatabricksDecisionsConfig] = DatabricksDecisionsConfig() diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 5aaab3f3544..90f6bb11c4a 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -21147,6 +21147,18 @@ "output_vector_size": 1024, "source": "https://www.databricks.com/product/pricing/foundation-model-serving" }, + "databricks/databricks-openjev-qwen35-4b": { + "cache_creation_input_token_cost": 0.0, + "cache_read_input_token_cost": 0.0, + "input_cost_per_token": 0.0, + "litellm_provider": "databricks", + "metadata": { + "notes": "Pay-per-token Foundation Model API decision endpoint (OpenJev, Qwen 3.5 4B) queried through the System One API. Databricks published no DBU rate for it as of 2026-10-07, so a deployment's model_info prices set the billed cost." + }, + "mode": "evaluation", + "output_cost_per_token": 0.0, + "source": "https://docs.databricks.com/aws/en/machine-learning/foundation-model-apis/supported-models" + }, "dataforseo/search": { "input_cost_per_query": 0.003, "litellm_provider": "dataforseo", diff --git a/litellm/proxy/common_utils/registry_read_through.py b/litellm/proxy/common_utils/registry_read_through.py index 0a47de96645..8f64072b4d0 100644 --- a/litellm/proxy/common_utils/registry_read_through.py +++ b/litellm/proxy/common_utils/registry_read_through.py @@ -7,18 +7,24 @@ immediately can land on a sibling that has never heard of it and fail 400/404. On a registry miss, callers here fetch the missing row from the DB and load it into the local registry before giving up. A short negative-result TTL per key plus a global resync budget per window bound the DB load from lookups of -genuinely unknown names. +genuinely unknown names. An auto router's tiers are deployment rows of their +own, so a miss on an auto router reconciles every row instead of just its own, +or the first routed request on this replica would pick a tier it has not loaded. """ import asyncio import time -from collections.abc import Awaitable, Callable +from collections.abc import Awaitable, Callable, Mapping, Sequence from typing import TYPE_CHECKING, Final +from pydantic import TypeAdapter, ValidationError + from litellm._logging import verbose_proxy_logger from litellm.caching.in_memory_cache import InMemoryCache +from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper if TYPE_CHECKING: + from prisma import models as prisma_models from prisma.types import ( LiteLLM_AgentsTableInclude, LiteLLM_AgentsTableWhereUniqueInput, @@ -129,6 +135,11 @@ async def _resync_model_deployments(model_name: str) -> bool: prisma_client=prisma_client, proxy_logging_obj=proxy_server.proxy_logging_obj ) return proxy_server.llm_router is not None + if _holds_an_auto_router(rows): + await proxy_server.proxy_config.add_deployment( + prisma_client=prisma_client, proxy_logging_obj=proxy_server.proxy_logging_obj + ) + return _model_is_loaded(model_name) async with proxy_server.MODEL_RECONCILE_LOCK: await proxy_server.proxy_config.get_credentials(prisma_client=prisma_client) proxy_server.proxy_config._add_deployment(db_models=rows) @@ -136,6 +147,25 @@ async def _resync_model_deployments(model_name: str) -> bool: return True +def _holds_an_auto_router(rows: Sequence["prisma_models.LiteLLM_ProxyModelTable"]) -> bool: + return any(_names_an_auto_router(row.litellm_params) for row in rows) + + +_STORED_LITELLM_PARAMS: Final = TypeAdapter(Mapping[str, object]) + + +def _names_an_auto_router(stored_litellm_params: object) -> bool: + try: + params: Final = _STORED_LITELLM_PARAMS.validate_python(stored_litellm_params) + except ValidationError: + return False + stored_model: Final = params.get("model") + if not isinstance(stored_model, str): + return False + model: Final = decrypt_value_helper(value=stored_model, key="model", return_original_value=True) + return isinstance(model, str) and model.startswith("auto_router/") + + async def _resync_guardrails(guardrail_name: str) -> bool: from litellm.proxy import proxy_server from litellm.proxy.guardrails.guardrail_registry import ( diff --git a/litellm/proxy/management_helpers/auto_router_permissions.py b/litellm/proxy/management_helpers/auto_router_permissions.py index 16f6b512e0d..e74578688c3 100644 --- a/litellm/proxy/management_helpers/auto_router_permissions.py +++ b/litellm/proxy/management_helpers/auto_router_permissions.py @@ -76,7 +76,7 @@ class _MemberOpenSourceClassifierConfig(LiteLLMBaseModel): model_config = ConfigDict(extra="forbid") - provider: Literal["jev", "laya", "bespoke"] = "jev" + provider: Literal["jev", "laya", "bespoke", "databricks"] = "jev" model: str api_key: None = None api_base: None = None diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index acc9770ab41..a5f00aa7e25 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -59,6 +59,7 @@ from litellm.litellm_core_utils.sensitive_data_masker import mask_credentials_in from litellm.llms.anthropic.common_utils import is_claude_code_user_agent from litellm.llms.base_llm.base_utils import type_to_response_format_param from litellm.llms.custom_httpx.http_handler import get_async_httpx_client +from litellm.llms.databricks.decisions.transformation import DATABRICKS_DECISIONS_CONFIG from litellm.router_strategy.adaptive_router.classifier import classify_prompt from litellm.router_strategy.complexity_router.context_compaction import compaction_pending from litellm.router_strategy.complexity_router.tier_predictor import ( @@ -1342,6 +1343,14 @@ class ComplexityRouter(CustomLogger): http_client=get_async_httpx_client(httpxSpecialProvider.PassThroughEndpoint), provider=config.provider, ) + if config.provider == "databricks": + databricks: Final = DATABRICKS_DECISIONS_CONFIG.connection(config.api_base, config.api_key) + return HttpJevClassifierClient( + api_key=databricks.api_key, + api_base=databricks.api_base, + http_client=get_async_httpx_client(httpxSpecialProvider.PassThroughEndpoint), + provider="databricks", + ) api_key: Final = config.api_key or get_secret_str("TYPESAFE_API_KEY") if not api_key: raise ValueError( diff --git a/litellm/router_strategy/complexity_router/config.py b/litellm/router_strategy/complexity_router/config.py index 797db460422..36c19e37b8d 100644 --- a/litellm/router_strategy/complexity_router/config.py +++ b/litellm/router_strategy/complexity_router/config.py @@ -705,10 +705,18 @@ def normalize_classifier_config_aliases(config: Mapping[str, object]) -> Mapping return normalized +_ENVIRONMENT_KEY_SCOPE: Final = MappingProxyType( + { + "jev": "TYPESAFE_API_KEY is only sent to TYPESAFE_API_BASE or https://api.typesafe.ai", + "databricks": "DATABRICKS_API_KEY or DATABRICKS_TOKEN is only sent to DATABRICKS_API_BASE", + } +) + + class OpenSourceClassifierConfig(LiteLLMBaseModel): model_config = ConfigDict(extra="forbid", frozen=True) - provider: Literal["jev", "laya", "bespoke"] = "jev" + provider: Literal["jev", "laya", "bespoke", "databricks"] = "jev" model: str = "jev-latest" api_key: str | None = Field(default=None, description="Provider API key; optional for self-hosted providers") api_base: str | None = Field( @@ -753,10 +761,19 @@ class OpenSourceClassifierConfig(LiteLLMBaseModel): if self.api_base is not None: _ = validate_oss_api_base(self.provider, self.api_base) return self + if self.provider == "databricks": + from litellm.llms.databricks.decisions.transformation import validate_serving_endpoint_name + + if "model" not in self.model_fields_set: + raise ValueError( + "opensource_classifier_config.model is required for provider 'databricks': the serving endpoint " + "name, e.g. databricks-openjev-qwen35-4b" + ) + _ = validate_serving_endpoint_name(self.model) if self.api_base is not None and self.api_key is None: raise ValueError( - "opensource_classifier_config.api_base requires opensource_classifier_config.api_key: TYPESAFE_API_KEY is only sent " - "to TYPESAFE_API_BASE or https://api.typesafe.ai" + "opensource_classifier_config.api_base requires opensource_classifier_config.api_key: " + f"{_ENVIRONMENT_KEY_SCOPE[self.provider]}" ) return self diff --git a/litellm/router_strategy/complexity_router/jev_classifier.py b/litellm/router_strategy/complexity_router/jev_classifier.py index 93c8b0c1bd3..f68462af8e0 100644 --- a/litellm/router_strategy/complexity_router/jev_classifier.py +++ b/litellm/router_strategy/complexity_router/jev_classifier.py @@ -18,6 +18,7 @@ from litellm.litellm_core_utils.internal_call_metadata import ( from litellm.litellm_core_utils.litellm_logging import Logging from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler +from litellm.llms.databricks.decisions.transformation import DATABRICKS_DECISIONS_CONFIG 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, @@ -27,6 +28,7 @@ from litellm.types.llms.base import LiteLLMBaseModel from litellm.types.utils import AUTOROUTER_CLASSIFIER_CALL_ORIGIN JevProbability: TypeAlias = Annotated[float, Field(ge=0.0, le=1.0)] +ClassifierProvider: TypeAlias = Literal["typesafe", "laya", "bespoke", "databricks"] DEFAULT_JEV_INSTRUCTIONS: Final = _DEFAULT_JEV_INSTRUCTIONS @@ -85,12 +87,12 @@ class HttpJevClassifierClient: api_key: str | None, api_base: str, http_client: AsyncHTTPHandler, - provider: Literal["typesafe", "laya", "bespoke"] = "typesafe", + provider: ClassifierProvider = "typesafe", ) -> None: self._api_key = api_key self._api_base = api_base.rstrip("/") self._http_client = http_client - self._provider = provider + self._provider: ClassifierProvider = provider async def evaluate( self, @@ -103,33 +105,43 @@ class HttpJevClassifierClient: MappingProxyType({"Authorization": f"Bearer {self._api_key}"}) if self._api_key else MappingProxyType({}) ) response: Final = await self._http_client.post( # pyright: ignore[reportUnknownMemberType] # AsyncHTTPHandler has a dynamic post signature - f"{self._api_base}/v1/systemone", + self._request_url(request.model), json=request.model_dump(mode="json"), headers=MappingProxyType({**authorization, "Content-Type": "application/json"}), # pyright: ignore[reportArgumentType] # HTTP headers are not mutated by AsyncHTTPHandler timeout=timeout_s, ) response.raise_for_status() body: Final = TypeAdapter(dict[str, object]).validate_json(response.content) - normalized_body: Final = ( - MappingProxyType({**body, "model": laya_response_model(body, request.model)}) - if self._provider == "laya" - else body - ) + normalized_body: Final = self._normalized_body(body, request.model) try: - self._log_response(request, response, request_kwargs, start_time) + self._log_response(request, response, normalized_body, 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(normalized_body) + def _request_url(self, model: str) -> str: + if self._provider == "databricks": + return DATABRICKS_DECISIONS_CONFIG.get_complete_url(self._api_base, model) + return f"{self._api_base}/v1/systemone" + + def _normalized_body(self, body: Mapping[str, object], requested_model: str) -> Mapping[str, object]: + match self._provider: + case "laya": + return MappingProxyType({**body, "model": laya_response_model(body, requested_model)}) + case "databricks": + return DATABRICKS_DECISIONS_CONFIG.classifier_response(body, requested_model) + case "typesafe" | "bespoke": + return body + def _log_response( self, request: JevSystemOneRequest, response: httpx.Response, + body: Mapping[str, object], request_kwargs: Mapping[str, object] | None, start_time: datetime, ) -> None: try: - body: Final = TypeAdapter(dict[str, object]).validate_json(response.content) _ = TypeAdapter(JevUsage | None).validate_python(body.get("usage")) except ValidationError: return @@ -202,7 +214,7 @@ class JevVerdict(NamedTuple): confidence: float model: str cost: float | None - provider: Literal["typesafe", "laya", "bespoke"] = "typesafe" + provider: ClassifierProvider = "typesafe" class _RegistryPricing(LiteLLMBaseModel): @@ -226,7 +238,7 @@ def build_jev_request( def jev_classifier_cost( - response: JevSystemOneResponse, configured_model: str, provider: Literal["typesafe", "laya", "bespoke"] = "typesafe" + response: JevSystemOneResponse, configured_model: str, provider: ClassifierProvider = "typesafe" ) -> float | None: usage: Final = response.usage if usage is None: diff --git a/litellm/utils.py b/litellm/utils.py index 25e0b3aba8c..52209dbfce8 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -9015,6 +9015,8 @@ class ProviderConfigManager: return litellm.CloudflareDecisionsConfig() if provider == LlmProviders.STRANDS_DECIDER: return litellm.StrandsDeciderDecisionsConfig() + if provider == LlmProviders.DATABRICKS: + return litellm.DatabricksDecisionsConfig() if provider == LlmProviders.HOSTED_VLLM: return litellm.HostedVLLMDecisionsConfig() if provider == LlmProviders.OPENAI: diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 5aaab3f3544..90f6bb11c4a 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -21147,6 +21147,18 @@ "output_vector_size": 1024, "source": "https://www.databricks.com/product/pricing/foundation-model-serving" }, + "databricks/databricks-openjev-qwen35-4b": { + "cache_creation_input_token_cost": 0.0, + "cache_read_input_token_cost": 0.0, + "input_cost_per_token": 0.0, + "litellm_provider": "databricks", + "metadata": { + "notes": "Pay-per-token Foundation Model API decision endpoint (OpenJev, Qwen 3.5 4B) queried through the System One API. Databricks published no DBU rate for it as of 2026-10-07, so a deployment's model_info prices set the billed cost." + }, + "mode": "evaluation", + "output_cost_per_token": 0.0, + "source": "https://docs.databricks.com/aws/en/machine-learning/foundation-model-apis/supported-models" + }, "dataforseo/search": { "input_cost_per_query": 0.003, "litellm_provider": "dataforseo", diff --git a/tests/integration/management/test_auto_router_databricks_classifier_writes.py b/tests/integration/management/test_auto_router_databricks_classifier_writes.py new file mode 100644 index 00000000000..6f541f5e622 --- /dev/null +++ b/tests/integration/management/test_auto_router_databricks_classifier_writes.py @@ -0,0 +1,334 @@ +import asyncio +import json +import uuid +from collections.abc import Callable, Mapping +from typing import Final + +import httpx +from integration._support.client import Gateway, Scenario, eventually, object_value, string_value +from integration._support.database import read_rows +from integration._support.openai_wire import answering_model_discovery, chat_reply +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue + +_ENDPOINT: Final = "databricks-openjev-qwen35-4b" +_SECOND_ENDPOINT: Final = "databricks-openjev-qwen35-4b-canary" +_API_KEY: Final = "synthetic-databricks-key" +_BARE_NAME_RULE: Final = ( + "must be the bare endpoint name (letters, digits, '-', '_' and '.', with no '/', '?', '#', spaces, or a " + "leading '.'), e.g. databricks-openjev-qwen35-4b" +) +_REJECTED_AT_LOAD: Final = "The router would drop this deployment at load time, so the write is rejected instead." +_MODEL_REQUIRED: Final = ( + "opensource_classifier_config.model is required for provider 'databricks': the serving endpoint name, " + "e.g. databricks-openjev-qwen35-4b" +) +_BASE_NEEDS_KEY: Final = ( + "opensource_classifier_config.api_base requires opensource_classifier_config.api_key: DATABRICKS_API_KEY or " + "DATABRICKS_TOKEN is only sent to DATABRICKS_API_BASE" +) +_BLANK_KEY: Final = ( + "opensource_classifier_config.api_key must be non-empty; omit it to use the provider environment key" +) +_CLASSIFIER_ANSWER: Final[dict[str, JsonValue]] = { + "answers": { + "tier": { + "type": "choice", + "choice": "COMPLEX", + "confidence": 0.91, + "probabilities": {"SIMPLE": 0.03, "MEDIUM": 0.06, "COMPLEX": 0.91}, + } + }, + "usage": {"input_tokens": 12, "output_tokens": 1}, +} +_ROWS_QUERY: Final = ( + "SELECT request_id, status, metadata->>'internal_call_origin' AS origin " + 'FROM "LiteLLM_SpendLogs" WHERE model_group = %s' +) + + +def _classifier(request: Request) -> Reply: + return Reply(body=json.dumps(_CLASSIFIER_ANSWER).encode()) + + +def _tier(text: str) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + body: Final = json.loads(request.body) + return chat_reply(f"{text}-{uuid.uuid4().hex[:8]}", "tier-model", text, stream=bool(body.get("stream"))) + + return answering_model_discovery(respond) + + +def _classifier_config(api_base: str, **overrides: JsonValue) -> dict[str, JsonValue]: + return { + "provider": "databricks", + "model": _ENDPOINT, + "api_base": api_base, + "api_key": _API_KEY, + "timeout_ms": 20000, + "circuit_breaker_enabled": False, + **overrides, + } + + +def _without(config: Mapping[str, JsonValue], field: str) -> dict[str, JsonValue]: + return {name: value for name, value in config.items() if name != field} + + +def _router_config(classifier: Mapping[str, JsonValue], simple: str, complex_: str) -> dict[str, JsonValue]: + return { + "classifier_type": "oss_classifier", + "opensource_classifier_config": dict(classifier), + "tiers": {"SIMPLE": simple, "MEDIUM": simple, "COMPLEX": complex_, "REASONING": complex_}, + } + + +def _validate(gateway: Gateway, config: Mapping[str, JsonValue]) -> dict[str, JsonValue]: + response: Final = gateway.request( + "POST", "/auto_router/validate_complexity_router_config", {"complexity_router_config": dict(config)} + ) + assert response.status_code == 200, response.text + return object_value(response.json()) + + +def _create_router(scenario: Scenario, config: Mapping[str, JsonValue]) -> tuple[str, str]: + name: Final = f"router-{uuid.uuid4().hex[:10]}" + created: Final = scenario.gateway.post( + "/model/new", + { + "model_name": name, + "litellm_params": {"model": "auto_router/complexity_router", "complexity_router_config": dict(config)}, + "model_info": {}, + }, + ) + identity: Final = string_value(object_value(created["model_info"])["id"]) + scenario.cleanups.callback(scenario.delete_model, identity) + return name, identity + + +def _stored_classifier(gateway: Gateway, identity: str) -> dict[str, JsonValue]: + stored: Final = gateway.request("GET", "/model/info", params={"litellm_model_id": identity}) + assert stored.status_code == 200, stored.text + (entry,) = stored.json()["data"] + router_config: Final = object_value(object_value(object_value(entry)["litellm_params"])["complexity_router_config"]) + return object_value(router_config["opensource_classifier_config"]) + + +def _judge_targets(judge: Wire) -> tuple[str, ...]: + return tuple(request.target for request in judge.drain() if request.method == "POST") + + +def _chat_body(model: str) -> dict[str, JsonValue]: + return {"model": model, "messages": [{"role": "user", "content": f"design a distributed cache {uuid.uuid4().hex}"}]} + + +async def _burst(base_url: str, key: str, model: str, count: int) -> tuple[httpx.Response, ...]: + async with httpx.AsyncClient( + base_url=base_url, timeout=60, trust_env=False, headers={"Authorization": f"Bearer {key}"} + ) as client: + return tuple( + await asyncio.gather(*(client.post("/v1/chat/completions", json=_chat_body(model)) for _ in range(count))) + ) + + +def _rows(router: str, *, want: int) -> tuple[dict[str, JsonValue], ...]: + return tuple(eventually(lambda: read_rows(_ROWS_QUERY, (router,)), lambda found: len(found) >= want, seconds=70)) + + +def test_validate_accepts_a_databricks_classifier_with_a_bare_endpoint_name_of_any_length(gateway: Gateway) -> None: + config: Final = _router_config( + _classifier_config("https://dbc-example.cloud.databricks.com/serving-endpoints"), "simple", "complex" + ) + assert _validate(gateway, config) == {"valid": True, "error": None} + long_name: Final = _router_config( + _classifier_config("https://dbc-example.cloud.databricks.com/serving-endpoints", model="e" * 5000), + "simple", + "complex", + ) + assert _validate(gateway, long_name) == {"valid": True, "error": None} + + +def test_validate_names_each_databricks_classifier_misconfiguration(gateway: Gateway) -> None: + base: Final = _classifier_config("https://dbc-example.cloud.databricks.com/serving-endpoints") + cases: Final[tuple[tuple[str, dict[str, JsonValue], str], ...]] = ( + ("model omitted", _without(base, "model"), f"at opensource_classifier_config: Value error, {_MODEL_REQUIRED}"), + ( + "api_base without api_key", + _without(base, "api_key"), + f"at opensource_classifier_config: Value error, {_BASE_NEEDS_KEY}", + ), + ("slash in the name", {**base, "model": "a/b"}, f"Databricks serving endpoint name 'a/b' {_BARE_NAME_RULE}"), + ("space in the name", {**base, "model": "a b"}, f"Databricks serving endpoint name 'a b' {_BARE_NAME_RULE}"), + ("leading dot", {**base, "model": ".hidden"}, f"Databricks serving endpoint name '.hidden' {_BARE_NAME_RULE}"), + ("empty name", {**base, "model": ""}, f"Databricks serving endpoint name '' {_BARE_NAME_RULE}"), + ( + "blank api_key", + {**base, "api_key": ""}, + f"at opensource_classifier_config.api_key: Value error, {_BLANK_KEY}", + ), + ("int name", {**base, "model": 5}, "at opensource_classifier_config.model: Input should be a valid string"), + ( + "list name", + {**base, "model": ["a"]}, + "at opensource_classifier_config.model: Input should be a valid string", + ), + ("null name", {**base, "model": None}, "at opensource_classifier_config.model: Input should be a valid string"), + ) + for label, classifier, expected in cases: + verdict: Final = _validate(gateway, _router_config(classifier, "simple", "complex")) + assert verdict["valid"] is False, (label, verdict) + error: Final = string_value(verdict["error"]) + assert error.startswith("complexity_router_config is invalid "), (label, error) + assert expected in error, (label, error) + assert error.endswith(_REJECTED_AT_LOAD), (label, error) + + +def test_test_routing_classifies_through_the_databricks_endpoint_without_calling_the_tier_it_picked( + gateway: Gateway, +) -> None: + prompt: Final = f"design a distributed cache {uuid.uuid4().hex}" + with ( + wire_server(_classifier) as judge, + wire_server(_tier("simple answer")) as simple, + wire_server(_tier("complex answer")) as complex_tier, + gateway.scenario() as scenario, + ): + simple_model: Final = scenario.model(model="openai/simple-tier", api_base=simple.url) + complex_model: Final = scenario.model(model="openai/complex-tier", api_base=complex_tier.url) + config: Final = _router_config(_classifier_config(judge.url), simple_model, complex_model) + response: Final = gateway.request( + "POST", "/auto_router/test_routing", {"prompt": prompt, "complexity_router_config": config} + ) + assert response.status_code == 200, response.text + assert response.json() == { + "routed_model": complex_model, + "routed_model_configured": True, + "routing_decision": { + "router_model_name": "auto_router_routing_test", + "router_type": "complexity", + "routed_model": complex_model, + "cause": "jev_classifier", + "tier": "COMPLEX", + "signals": [ + "databricks-classifier:COMPLEX", + "databricks-confidence=0.910000", + "tier-probability:SIMPLE=0.030000", + "tier-probability:MEDIUM=0.060000", + "tier-probability:COMPLEX=0.910000", + ], + "classifier_model": f"databricks/{_ENDPOINT}", + "classifier_cost": 0.0, + "classifier_probabilities": {"SIMPLE": 0.03, "MEDIUM": 0.06, "COMPLEX": 0.91}, + "classifier_confidence": 0.91, + "conversation_continuing": False, + }, + }, response.text + (call,) = [request for request in judge.drain() if request.method == "POST"] + assert call.target == f"/{_ENDPOINT}/invocations", call.target + assert call.headers.get("authorization") == f"Bearer {_API_KEY}", call.headers + body: Final = json.loads(call.body) + assert body["model"] == _ENDPOINT, body + assert body["state"] == f"\nClassify this message:\n{prompt}", body + assert _judge_targets(simple) == () and _judge_targets(complex_tier) == () + + +def test_model_new_refuses_a_databricks_classifier_without_an_endpoint_name_and_stores_nothing( + gateway: Gateway, +) -> None: + name: Final = f"router-{uuid.uuid4().hex[:10]}" + config: Final = _router_config( + _without(_classifier_config("https://dbc-example.cloud.databricks.com/serving-endpoints"), "model"), + "simple", + "complex", + ) + response: Final = gateway.request( + "POST", + "/model/new", + { + "model_name": name, + "litellm_params": {"model": "auto_router/complexity_router", "complexity_router_config": config}, + "model_info": {}, + }, + ) + assert response.status_code == 400, response.text + error: Final = object_value(object_value(response.json())["error"]) + assert _MODEL_REQUIRED in string_value(error["message"]), error + assert (error["type"], error["param"]) == ("validation_error", "litellm_params.model"), error + listed: Final = gateway.get("/model/info")["data"] + assert isinstance(listed, list) + assert [entry for entry in map(object_value, listed) if entry["model_name"] == name] == [] + + +def test_model_update_refuses_dropping_the_endpoint_name_and_keeps_the_stored_classifier_serving( + gateway: Gateway, +) -> None: + with ( + wire_server(_classifier) as judge, + wire_server(_tier("simple answer")) as simple, + wire_server(_tier("complex answer")) as complex_tier, + gateway.scenario() as scenario, + ): + simple_model: Final = scenario.model(model="openai/simple-tier", api_base=simple.url) + complex_model: Final = scenario.model(model="openai/complex-tier", api_base=complex_tier.url) + classifier: Final = _classifier_config(judge.url) + name, identity = _create_router(scenario, _router_config(classifier, simple_model, complex_model)) + dropped: Final = _router_config(_without(classifier, "model"), simple_model, complex_model) + response: Final = gateway.request( + "PATCH", + f"/model/{identity}/update", + {"litellm_params": {"model": "auto_router/complexity_router", "complexity_router_config": dropped}}, + ) + assert response.status_code == 400, response.text + assert _MODEL_REQUIRED in string_value(object_value(object_value(response.json())["error"])["message"]), ( + response.text + ) + assert _stored_classifier(gateway, identity)["model"] == _ENDPOINT + chat: Final = gateway.request("POST", "/v1/chat/completions", _chat_body(name)) + assert chat.status_code == 200, chat.text + assert chat.json()["choices"][0]["message"]["content"] == "complex answer", chat.text + assert chat.headers["x-litellm-complexity-router-cause"] == "jev_classifier", dict(chat.headers) + assert _judge_targets(judge) == (f"/{_ENDPOINT}/invocations",) + + +async def test_renaming_the_endpoint_under_a_burst_moves_the_classifier_and_logs_every_request_once( + gateway: Gateway, +) -> None: + base_url: Final = str(gateway.client.base_url) + with ( + wire_server(_classifier) as judge, + wire_server(_tier("simple answer")) as simple, + wire_server(_tier("complex answer")) as complex_tier, + gateway.scenario() as scenario, + ): + simple_model: Final = scenario.model(model="openai/simple-tier", api_base=simple.url) + complex_model: Final = scenario.model(model="openai/complex-tier", api_base=complex_tier.url) + classifier: Final = _classifier_config(judge.url) + name, identity = _create_router(scenario, _router_config(classifier, simple_model, complex_model)) + burst: Final = asyncio.create_task(_burst(base_url, gateway.key, name, 24)) + renamed: Final = _router_config({**classifier, "model": _SECOND_ENDPOINT}, simple_model, complex_model) + patched: Final = gateway.request( + "PATCH", + f"/model/{identity}/update", + {"litellm_params": {"model": "auto_router/complexity_router", "complexity_router_config": renamed}}, + ) + assert patched.status_code == 200, patched.text + during: Final = await burst + assert [response.status_code for response in during] == [200] * 24, [response.text for response in during] + assert set(_judge_targets(judge)) <= {f"/{_ENDPOINT}/invocations", f"/{_SECOND_ENDPOINT}/invocations"} + assert _stored_classifier(gateway, identity)["model"] == _SECOND_ENDPOINT + + def classifier_target() -> tuple[str, ...]: + chat: Final = gateway.request("POST", "/v1/chat/completions", _chat_body(name)) + assert chat.status_code == 200, chat.text + return _judge_targets(judge) + + eventually(classifier_target, lambda targets: targets == (f"/{_SECOND_ENDPOINT}/invocations",), seconds=30) + after: Final = await _burst(base_url, gateway.key, name, 6) + assert [response.status_code for response in after] == [200] * 6, [response.text for response in after] + assert set(_judge_targets(judge)) == {f"/{_SECOND_ENDPOINT}/invocations"} + identities: Final = tuple(string_value(response.json()["id"]) for response in (*during, *after)) + assert len(set(identities)) == len(identities), identities + rows: Final = _rows(name, want=len(identities)) + request_rows: Final = [row for row in rows if row["origin"] != "autorouter_classifier"] + assert sorted(string_value(row["request_id"]) for row in request_rows) == sorted(identities), request_rows + assert all(row["status"] == "success" for row in rows), rows diff --git a/tests/integration/management/test_model_health_check.py b/tests/integration/management/test_model_health_check.py index 11b902b540a..075c6051353 100644 --- a/tests/integration/management/test_model_health_check.py +++ b/tests/integration/management/test_model_health_check.py @@ -186,3 +186,43 @@ def test_evaluation_mode_health_check_of_hosted_vllm_sends_a_choice_probe(gatewa }, ) ] + + +_DATABRICKS_PROBE_REPLY: Final = JsonResponse( + content_type="application/json", + body={ + "model": "databricks-openjev-qwen35-4b", + "answers": {"reachable": {"type": "noul", "noul": 1.0}}, + "usage": {"input_tokens": 10, "output_tokens": 1}, + }, +) + + +def test_evaluation_mode_health_check_of_a_databricks_serving_endpoint_resolves_the_mode_from_the_cost_map( + gateway: Gateway, +) -> None: + with gateway.scenario() as scenario: + handle: Final = register_scenario(f"health-decisions-{uuid.uuid4().hex[:12]}", _DATABRICKS_PROBE_REPLY) + scenario.cleanups.callback(delete_scenario, handle) + model: Final = scenario.model( + model="databricks/databricks-openjev-qwen35-4b", + api_base=handle.api_base(), + api_key="synthetic-databricks-key", + ) + listed: Final = gateway.request("GET", "/v2/model/info", params={"model": model}) + assert listed.status_code == 200, listed.text + assert [object_value(object_value(entry)["model_info"])["mode"] for entry in listed.json()["data"]] == [ + "evaluation" + ], listed.text + report: Final = _health_report(gateway, model) + assert (report["healthy_count"], report["unhealthy_count"]) == (1, 0), report + assert _probes_sent_to(gateway, handle) == [ + ( + f"/{handle.scenario_id}/databricks-openjev-qwen35-4b/invocations", + { + "model": "databricks-openjev-qwen35-4b", + "state": os.environ.get("DEFAULT_HEALTH_CHECK_PROMPT", "test from litellm"), + "questions": {"reachable": {"type": "noul", "instructions": "Is the service reachable?"}}, + }, + ) + ] diff --git a/tests/integration/providers/test_databricks_decisions_wire.py b/tests/integration/providers/test_databricks_decisions_wire.py new file mode 100644 index 00000000000..5c3b7d1b1a1 --- /dev/null +++ b/tests/integration/providers/test_databricks_decisions_wire.py @@ -0,0 +1,519 @@ +import threading +import time +import uuid +from collections.abc import Mapping, Sequence +from pathlib import Path +from typing import Final + +import httpx +import pytest +from integration._support.client import Gateway, Scenario, eventually, object_value, string_value +from integration._support.database import read_rows +from integration._support.process import graceful_stop_seconds, owned_proxy_process +from integration._support.upstream import ScenarioHandle, delete_scenario, register_scenario +from integration.cost_calculation.cost_tracking_case import JsonResponse +from pydantic import JsonValue + +_ENDPOINT: Final = "databricks-openjev-qwen35-4b" +_API_KEY: Final = "synthetic-databricks-key" +_ENV_KEY: Final = "synthetic-databricks-env-key" +_ENV_TOKEN: Final = "synthetic-databricks-env-token" +_STATE: Final[dict[str, JsonValue]] = {"ticket": "The export job hangs at 99%", "component": "billing"} +_TIER_CRITERIA: Final[dict[str, JsonValue]] = {"SIMPLE": "a lookup", "MEDIUM": "some work", "COMPLEX": "deep work"} +_QUESTIONS: Final[dict[str, JsonValue]] = {"tier": {"type": "choice", "criteria": _TIER_CRITERIA}} +_OPENAI_INPUT: Final = "The export job hangs at 99% in billing" +_OPENAI_TIER_QUESTION: Final[dict[str, JsonValue]] = { + "type": "choice", + "name": "tier", + "instructions": "Which tier?", + "choices": [{"value": name, "description": description} for name, description in _TIER_CRITERIA.items()], +} +_ANSWERS: Final[dict[str, JsonValue]] = { + "tier": { + "type": "choice", + "choice": "COMPLEX", + "confidence": 0.91, + "probabilities": {"SIMPLE": 0.03, "MEDIUM": 0.06, "COMPLEX": 0.91}, + } +} +_USAGE: Final[dict[str, JsonValue]] = {"input_tokens": 367, "output_tokens": 3} +_ANSWER_BODY: Final[dict[str, JsonValue]] = {"model": _ENDPOINT, "answers": _ANSWERS, "usage": _USAGE} +_BARE_NAME_RULE: Final = ( + "must be the bare endpoint name (letters, digits, '-', '_' and '.', with no '/', '?', '#', spaces, or a " + "leading '.'), e.g. databricks-openjev-qwen35-4b" +) +_MISSING_KEY: Final = "Missing API key for Decisions provider 'databricks'" +_NON_BARE_NAMES: Final[tuple[tuple[str, str], ...]] = ( + ("slash", "foo/bar"), + ("space", "a b"), + ("leading-dot", ".hidden"), + ("question-mark", "a?b"), + ("hash", "a#b"), +) +_DATABRICKS_ERRORS: Final[tuple[tuple[int, dict[str, JsonValue]], ...]] = ( + (400, {"error_code": "BAD_REQUEST", "message": "Invalid input: questions must not be empty"}), + (404, {"error_code": "RESOURCE_DOES_NOT_EXIST", "message": f"Endpoint with name '{_ENDPOINT}' does not exist."}), +) +_SPEND_QUERY: Final = ( + 'SELECT spend, status, call_type, model_group, custom_llm_provider, api_base FROM "LiteLLM_SpendLogs" ' + "WHERE request_id = %s" +) +_OWNED_PROXY_BUDGET: Final = int(2 * graceful_stop_seconds() + 120) +_MEGABYTE_NAME: Final = "a" * 1_000_000 + "/" + + +def _number(value: JsonValue) -> float: + assert isinstance(value, (int, float)) and not isinstance(value, bool), value + return float(value) + + +def _register(scenario: Scenario, body: dict[str, JsonValue], *, status: int = 200) -> ScenarioHandle: + handle: Final = register_scenario( + f"databricks-{uuid.uuid4().hex[:12]}", JsonResponse(content_type="application/json", body=body, status=status) + ) + scenario.cleanups.callback(delete_scenario, handle) + return handle + + +def _decide(gateway: Gateway, model: str, *, key: str | None = None) -> httpx.Response: + return gateway.request("POST", "/v1/systemone", {"model": model, "state": _STATE, "questions": _QUESTIONS}, key=key) + + +def _observed(gateway: Gateway) -> tuple[dict[str, JsonValue], ...]: + with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream: + return tuple(map(object_value, upstream.get("/__observations").json()["requests"])) + + +def _calls_to(requests: Sequence[dict[str, JsonValue]], handle: ScenarioHandle) -> list[dict[str, JsonValue]]: + return [request for request in requests if string_value(request["path"]).startswith(f"/{handle.scenario_id}/")] + + +def _upstream_calls(gateway: Gateway, handle: ScenarioHandle) -> list[dict[str, JsonValue]]: + return _calls_to(_observed(gateway), handle) + + +def _spend_row(call_id: str) -> dict[str, JsonValue]: + rows: Final = eventually(lambda: read_rows(_SPEND_QUERY, (call_id,)), lambda found: len(found) == 1, seconds=70) + return rows[0] + + +def _wildcard_deployment(scenario: Scenario, **parameters: JsonValue) -> str: + prefix: Final = f"databricks-{uuid.uuid4().hex[:8]}" + created: Final = scenario.gateway.post( + "/model/new", + {"model_name": f"{prefix}/*", "litellm_params": {"model": "databricks/*", **parameters}, "model_info": {}}, + ) + scenario.cleanups.callback(scenario.delete_model, string_value(object_value(created["model_info"])["id"])) + return prefix + + +def _router_config(classifier: Mapping[str, JsonValue], simple: str, complex_: str) -> dict[str, JsonValue]: + return { + "classifier_type": "oss_classifier", + "opensource_classifier_config": dict(classifier), + "tiers": {"SIMPLE": simple, "MEDIUM": simple, "COMPLEX": complex_, "REASONING": complex_}, + } + + +def _environment_classifier() -> dict[str, JsonValue]: + return {"provider": "databricks", "model": _ENDPOINT, "timeout_ms": 20000, "circuit_breaker_enabled": False} + + +def _router(scenario: Scenario, classifier: Mapping[str, JsonValue]) -> str: + simple: Final = scenario.model(model="openai/simple-tier") + complex_: Final = scenario.model(model="openai/complex-tier") + name: Final = f"router-{uuid.uuid4().hex[:10]}" + created: Final = scenario.gateway.post( + "/model/new", + { + "model_name": name, + "litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_config": _router_config(classifier, simple, complex_), + }, + "model_info": {}, + }, + ) + scenario.cleanups.callback(scenario.delete_model, string_value(object_value(created["model_info"])["id"])) + return name + + +def _chat(gateway: Gateway, model: str, *, key: str | None = None) -> httpx.Response: + return gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": f"design a distributed cache {uuid.uuid4().hex}"}]}, + key=key, + ) + + +def _assert_classified_by_the_complex_tier(response: httpx.Response, router: str) -> None: + assert response.status_code == 200, response.text + assert response.json()["choices"][0]["message"]["content"], response.text + assert ( + response.headers["x-litellm-model-group"], + response.headers["x-litellm-model-name"], + response.headers["x-litellm-complexity-router-tier"], + response.headers["x-litellm-complexity-router-cause"], + response.headers["x-litellm-classifier-cost"], + ) == (router, "openai/complex-tier", "COMPLEX", "jev_classifier", "0.0"), dict(response.headers) + + +def _assert_classifier_call( + call: Mapping[str, JsonValue], *, path: str, authorization: str, prompt_marker: str +) -> None: + assert call["path"] == path, call + assert call["authorization"] == authorization, call + body: Final = object_value(call["body"]) + assert body["model"] == _ENDPOINT, body + assert prompt_marker in string_value(body["state"]), body + assert sorted(object_value(object_value(body["questions"])["tier"])) == ["criteria", "instructions", "type"], body + + +def _assert_refused_before_any_upstream_call(gateway: Gateway, model: str, message: str) -> None: + with gateway.scenario() as scenario: + handle: Final = _register(scenario, _ANSWER_BODY) + deployment: Final = scenario.model(model=model, api_base=handle.api_base(), api_key=_API_KEY) + response: Final = _decide(gateway, deployment) + assert response.status_code == 400, response.text + assert message in response.text, response.text + assert _upstream_calls(gateway, handle) == [] + row: Final = _spend_row(response.headers["x-litellm-call-id"]) + assert (row["status"], row["call_type"], _number(row["spend"])) == ("failure", "adecisions", 0.0), row + + +@pytest.mark.parametrize("name", [name for _, name in _NON_BARE_NAMES], ids=[label for label, _ in _NON_BARE_NAMES]) +def test_a_deployment_naming_a_non_bare_endpoint_is_refused_before_any_upstream_call( + gateway: Gateway, name: str +) -> None: + _assert_refused_before_any_upstream_call( + gateway, f"databricks/{name}", f"Databricks serving endpoint name {name!r} {_BARE_NAME_RULE}" + ) + + +def test_a_deployment_naming_no_endpoint_is_refused_before_any_upstream_call(gateway: Gateway) -> None: + _assert_refused_before_any_upstream_call(gateway, "databricks/", "A model name is required for the Decisions API") + + +def test_an_openai_format_request_at_v1_decisions_reaches_the_serving_endpoint_as_system_one(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + handle: Final = _register(scenario, _ANSWER_BODY) + deployment: Final = scenario.model( + model=f"databricks/{_ENDPOINT}", api_base=handle.api_base(), api_key=_API_KEY + ) + response: Final = gateway.request( + "POST", + "/v1/decisions", + { + "model": deployment, + "input": _OPENAI_INPUT, + "questions": [_OPENAI_TIER_QUESTION], + "safety_identifier": "end-user-1", + }, + ) + assert response.status_code == 200, response.text + assert response.json() == { + "model": _ENDPOINT, + "answers": [ + { + "type": "choice", + "name": "tier", + "choice": "COMPLEX", + "probabilities": [ + {"value": "SIMPLE", "probability": 0.03}, + {"value": "MEDIUM", "probability": 0.06}, + {"value": "COMPLEX", "probability": 0.91}, + ], + "confidence": 0.91, + } + ], + "usage": { + "input_tokens": 367, + "input_tokens_details": {"cached_tokens": 0, "cache_write_tokens": 0}, + "output_tokens": 3, + "output_tokens_details": {"reasoning_tokens": 0}, + "total_tokens": 370, + }, + }, response.text + assert response.headers["x-litellm-model-name"] == f"databricks/{_ENDPOINT}", dict(response.headers) + (call,) = _upstream_calls(gateway, handle) + assert call["path"] == f"/{handle.scenario_id}/{_ENDPOINT}/invocations", call + assert call["authorization"] == f"Bearer {_API_KEY}", call + assert call["body"] == { + "model": _ENDPOINT, + "state": _OPENAI_INPUT, + "questions": {"tier": {"type": "choice", "instructions": "Which tier?", "criteria": _TIER_CRITERIA}}, + }, call + row: Final = _spend_row(response.headers["x-litellm-call-id"]) + assert ( + row["status"], + row["call_type"], + row["custom_llm_provider"], + row["model_group"], + row["api_base"], + _number(row["spend"]), + ) == ("success", "adecisions", "databricks", deployment, f"{handle.api_base()}/{_ENDPOINT}/invocations", 0.0), ( + row + ) + + +def test_a_wildcard_deployment_sends_each_request_to_the_named_endpoint_and_bills_each_once(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + handle: Final = _register(scenario, _ANSWER_BODY) + prefix: Final = _wildcard_deployment(scenario, api_base=handle.api_base(), api_key=_API_KEY) + model: Final = f"{prefix}/{_ENDPOINT}" + responses: Final = tuple(_decide(gateway, model) for _ in range(2)) + for response in responses: + assert response.status_code == 200, response.text + assert response.json() == _ANSWER_BODY, response.text + assert response.headers["x-litellm-model-group"] == model, dict(response.headers) + assert response.headers["x-litellm-model-name"] == f"databricks/{_ENDPOINT}", dict(response.headers) + call_ids: Final = tuple(response.headers["x-litellm-call-id"] for response in responses) + assert len(set(call_ids)) == 2, call_ids + calls: Final = _upstream_calls(gateway, handle) + assert len(calls) == 2, calls + for call in calls: + assert call["path"] == f"/{handle.scenario_id}/{_ENDPOINT}/invocations", call + assert call["authorization"] == f"Bearer {_API_KEY}", call + assert call["body"] == {"model": _ENDPOINT, "state": _STATE, "questions": _QUESTIONS}, call + for call_id in call_ids: + row: Final = _spend_row(call_id) + assert ( + row["status"], + row["call_type"], + row["custom_llm_provider"], + row["model_group"], + row["api_base"], + _number(row["spend"]), + ) == ("success", "adecisions", "databricks", model, f"{handle.api_base()}/{_ENDPOINT}/invocations", 0.0), ( + row + ) + + +@pytest.mark.parametrize("model", (5, ["a"]), ids=("int", "list")) +def test_a_non_string_model_is_refused_at_the_gateway_without_an_upstream_call( + gateway: Gateway, model: JsonValue +) -> None: + with gateway.scenario() as scenario: + handle: Final = _register(scenario, _ANSWER_BODY) + _wildcard_deployment(scenario, api_base=handle.api_base(), api_key=_API_KEY) + response: Final = gateway.request( + "POST", "/v1/systemone", {"model": model, "state": _STATE, "questions": _QUESTIONS} + ) + assert 400 <= response.status_code < 500, response.text + assert _upstream_calls(gateway, handle) == [] + + +def test_a_five_kilobyte_endpoint_name_reaches_the_upstream_and_its_not_found_answer_is_kept(gateway: Gateway) -> None: + name: Final = "n" * 5000 + with gateway.scenario() as scenario: + handle: Final = _register( + scenario, + {"error_code": "RESOURCE_DOES_NOT_EXIST", "message": f"Endpoint with name '{name}' does not exist."}, + status=404, + ) + prefix: Final = _wildcard_deployment(scenario, api_base=handle.api_base(), api_key=_API_KEY) + response: Final = _decide(gateway, f"{prefix}/{name}") + assert response.status_code == 404, response.text[:500] + assert "RESOURCE_DOES_NOT_EXIST" in response.text, response.text[:500] + (call,) = _upstream_calls(gateway, handle) + assert call["path"] == f"/{handle.scenario_id}/{name}/invocations", call["path"][:200] + assert call["authorization"] == f"Bearer {_API_KEY}", call["authorization"] + assert call["body"] == {"model": name, "state": _STATE, "questions": _QUESTIONS} + row: Final = _spend_row(response.headers["x-litellm-call-id"]) + assert (row["status"], row["call_type"], row["model_group"], _number(row["spend"])) == ( + "failure", + "adecisions", + f"{prefix}/{name}", + 0.0, + ), row + + +@pytest.mark.parametrize("status,body", _DATABRICKS_ERRORS, ids=("400", "404")) +def test_a_databricks_error_keeps_its_status_and_message_and_bills_nothing( + gateway: Gateway, status: int, body: dict[str, JsonValue] +) -> None: + with gateway.scenario() as scenario: + handle: Final = _register(scenario, body, status=status) + model: Final = scenario.model(model=f"databricks/{_ENDPOINT}", api_base=handle.api_base(), api_key=_API_KEY) + response: Final = _decide(gateway, model) + assert response.status_code == status, response.text + assert string_value(body["message"]) in response.text, response.text + (call,) = _upstream_calls(gateway, handle) + assert call["path"] == f"/{handle.scenario_id}/{_ENDPOINT}/invocations", call + row: Final = _spend_row(response.headers["x-litellm-call-id"]) + assert (row["status"], row["call_type"], row["model_group"], _number(row["spend"])) == ( + "failure", + "adecisions", + model, + 0.0, + ), row + + +def test_a_deployment_without_a_key_and_no_environment_key_is_refused_before_any_upstream_call( + gateway: Gateway, +) -> None: + with gateway.scenario() as scenario: + handle: Final = _register(scenario, _ANSWER_BODY) + model: Final = scenario.model(model=f"databricks/{_ENDPOINT}", api_base=handle.api_base(), api_key=None) + response: Final = _decide(gateway, model) + assert response.status_code == 401, response.text + assert _MISSING_KEY in response.text, response.text + assert _upstream_calls(gateway, handle) == [] + + +@pytest.mark.timeout(_OWNED_PROXY_BUDGET) +def test_the_environment_key_and_base_serve_a_deployment_a_classifier_and_a_team_member_that_leave_them_unset( + gateway: Gateway, tmp_path: Path +) -> None: + with gateway.scenario() as rig_scenario: + handle: Final = _register(rig_scenario, _ANSWER_BODY) + environment_base: Final = f"{handle.api_base()}/serving-endpoints" + environment_path: Final = f"/{handle.scenario_id}/serving-endpoints/{_ENDPOINT}/invocations" + environment: Final = { + "DATABRICKS_API_KEY": _ENV_KEY, + "DATABRICKS_TOKEN": _ENV_TOKEN, + "DATABRICKS_API_BASE": environment_base, + } + with owned_proxy_process(gateway, tmp_path, environment) as owned, owned.gateway.scenario() as scenario: + candidate: Final = owned.gateway + with_base: Final = scenario.model(model=f"databricks/{_ENDPOINT}", api_base=handle.api_base(), api_key=None) + bare: Final = scenario.model(model=f"databricks/{_ENDPOINT}", api_base=None, api_key=None) + decided: Final = tuple(_decide(candidate, model) for model in (with_base, bare)) + for response in decided: + assert response.status_code == 200, response.text + assert response.json() == _ANSWER_BODY, response.text + assert [(call["path"], call["authorization"]) for call in _upstream_calls(candidate, handle)] == [ + (f"/{handle.scenario_id}/{_ENDPOINT}/invocations", f"Bearer {_ENV_KEY}"), + (environment_path, f"Bearer {_ENV_KEY}"), + ] + for response, model in zip(decided, (with_base, bare)): + row: Final = _spend_row(response.headers["x-litellm-call-id"]) + assert (row["status"], row["model_group"], _number(row["spend"])) == ("success", model, 0.0), row + assert ( + _spend_row(decided[1].headers["x-litellm-call-id"])["api_base"] + == f"{environment_base}/{_ENDPOINT}/invocations" + ) + + router: Final = _router(scenario, _environment_classifier()) + routed: Final = _chat(candidate, router) + _assert_classified_by_the_complex_tier(routed, router) + (judge_call,) = _upstream_calls(candidate, handle) + _assert_classifier_call( + judge_call, + path=environment_path, + authorization=f"Bearer {_ENV_KEY}", + prompt_marker="design a distributed cache", + ) + + team: Final = scenario.team(team_member_permissions=["/auto_router/manage"]) + member: Final = scenario.member(team) + member_key: Final = scenario.key(user_id=member, team_id=team) + simple: Final = scenario.model(model="openai/simple-tier") + complex_: Final = scenario.model(model="openai/complex-tier") + team_router: Final = f"team-router-{uuid.uuid4().hex[:8]}" + created: Final = candidate.request( + "POST", + "/model/new", + { + "model_name": team_router, + "litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_config": _router_config(_environment_classifier(), simple, complex_), + }, + "model_info": {"team_id": team}, + }, + key=member_key, + ) + assert created.status_code == 200, created.text + team_router_id: Final = string_value(object_value(object_value(created.json())["model_info"])["id"]) + scenario.cleanups.callback(scenario.delete_model, team_router_id) + stored: Final = candidate.request("GET", "/model/info", params={"litellm_model_id": team_router_id}) + assert stored.status_code == 200, stored.text + (entry,) = stored.json()["data"] + stored_classifier: Final = object_value( + object_value(object_value(object_value(entry)["litellm_params"])["complexity_router_config"])[ + "opensource_classifier_config" + ] + ) + assert (stored_classifier["provider"], stored_classifier["model"]) == ("databricks", _ENDPOINT), entry + assert "api_key" not in stored_classifier and "api_base" not in stored_classifier, entry + member_chat: Final = _chat(candidate, team_router, key=member_key) + _assert_classified_by_the_complex_tier(member_chat, team_router) + (member_judge_call,) = _upstream_calls(candidate, handle) + _assert_classifier_call( + member_judge_call, + path=environment_path, + authorization=f"Bearer {_ENV_KEY}", + prompt_marker="design a distributed cache", + ) + + +@pytest.mark.timeout(_OWNED_PROXY_BUDGET) +def test_the_environment_token_serves_a_deployment_and_a_classifier_when_no_environment_key_is_set( + gateway: Gateway, tmp_path: Path +) -> None: + with gateway.scenario() as rig_scenario: + handle: Final = _register(rig_scenario, _ANSWER_BODY) + environment_path: Final = f"/{handle.scenario_id}/serving-endpoints/{_ENDPOINT}/invocations" + environment: Final = { + "DATABRICKS_TOKEN": _ENV_TOKEN, + "DATABRICKS_API_BASE": f"{handle.api_base()}/serving-endpoints", + } + with ( + owned_proxy_process(gateway, tmp_path, environment, remove_environment=("DATABRICKS_API_KEY",)) as owned, + owned.gateway.scenario() as scenario, + ): + candidate: Final = owned.gateway + bare: Final = scenario.model(model=f"databricks/{_ENDPOINT}", api_base=None, api_key=None) + decided: Final = _decide(candidate, bare) + assert decided.status_code == 200, decided.text + assert decided.json() == _ANSWER_BODY, decided.text + (decision_call,) = _upstream_calls(candidate, handle) + assert (decision_call["path"], decision_call["authorization"]) == (environment_path, f"Bearer {_ENV_TOKEN}") + assert _spend_row(decided.headers["x-litellm-call-id"])["status"] == "success" + router: Final = _router(scenario, _environment_classifier()) + routed: Final = _chat(candidate, router) + _assert_classified_by_the_complex_tier(routed, router) + (judge_call,) = _upstream_calls(candidate, handle) + _assert_classifier_call( + judge_call, + path=environment_path, + authorization=f"Bearer {_ENV_TOKEN}", + prompt_marker="design a distributed cache", + ) + + +def test_a_megabyte_endpoint_name_is_refused_in_linear_time_while_liveliness_stays_fast(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + handle: Final = _register(scenario, _ANSWER_BODY) + prefix: Final = _wildcard_deployment(scenario, api_base=handle.api_base(), api_key=_API_KEY) + model: Final = f"{prefix}/{_MEGABYTE_NAME}" + anonymous: Final = gateway.client.post( + "/v1/systemone", json={"model": model, "state": _STATE, "questions": _QUESTIONS} + ) + assert anonymous.status_code == 401, anonymous.text[:300] + liveliness: Final[list[tuple[int, float]]] = [] + stop: Final = threading.Event() + + def poll() -> None: + while not stop.is_set(): + started: Final = time.perf_counter() + probe: Final = gateway.client.get("/health/liveliness") + liveliness.append((probe.status_code, time.perf_counter() - started)) + + poller: Final = threading.Thread(target=poll) + poller.start() + started: Final = time.perf_counter() + response: Final = _decide(gateway, model) + elapsed: Final = time.perf_counter() - started + stop.set() + poller.join() + assert elapsed < 3, elapsed + assert liveliness and all(status == 200 for status, _ in liveliness), liveliness + assert max(seconds for _, seconds in liveliness) < 1, liveliness + assert response.status_code == 400, response.text[:300] + assert _BARE_NAME_RULE in response.text, response.text[-400:] + assert _upstream_calls(gateway, handle) == [] + row: Final = _spend_row(response.headers["x-litellm-call-id"]) + assert (row["status"], row["call_type"], _number(row["spend"])) == ("failure", "adecisions", 0.0), row diff --git a/tests/integration/providers/test_decisions_wire.py b/tests/integration/providers/test_decisions_wire.py index 4830f0ecc3d..ac4da38ef37 100644 --- a/tests/integration/providers/test_decisions_wire.py +++ b/tests/integration/providers/test_decisions_wire.py @@ -108,6 +108,15 @@ _PROVIDERS: Final = ( True, "cloudflare/@cf/cloudflare/clef", ), + _Provider( + "databricks", + "databricks/databricks-openjev-qwen35-4b", + "/databricks-openjev-qwen35-4b/invocations", + "databricks-openjev-qwen35-4b", + _API_KEY, + False, + None, + ), ) _PERPLEXITY: Final = _PROVIDERS[0] _OPENROUTER: Final = _PROVIDERS[2] diff --git a/tests/integration/routing/test_complexity_router_databricks_classifier.py b/tests/integration/routing/test_complexity_router_databricks_classifier.py new file mode 100644 index 00000000000..25611b4ec08 --- /dev/null +++ b/tests/integration/routing/test_complexity_router_databricks_classifier.py @@ -0,0 +1,1023 @@ +import asyncio +import json +import os +import re +import signal +import threading +import time +import uuid +from collections.abc import Callable, Iterator, Mapping, Sequence +from contextlib import ExitStack, contextmanager +from dataclasses import dataclass +from pathlib import Path +from queue import SimpleQueue +from types import MappingProxyType +from typing import Final +from urllib.parse import urlsplit + +import anthropic +import httpx +import openai +import psutil +import pytest +import yaml +from integration._support.client import Gateway, Scenario, eventually, object_value, string_value +from integration._support.database import read_rows +from integration._support.openai_wire import answering_model_discovery, chat_reply, responses_reply +from integration._support.process import graceful_stop_seconds, owned_proxy_process +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue, TypeAdapter + +from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_if_encrypted_with +from litellm.router_strategy.complexity_router.config import DEFAULT_JEV_INSTRUCTIONS + +_ENDPOINT: Final = "databricks-openjev-qwen35-4b" +_API_KEY: Final = "synthetic-databricks-key" +_CLASSIFIER_MODEL: Final = f"databricks/{_ENDPOINT}" +_TYPESAFE_MODEL: Final = "jev-1.13.0" +_TYPESAFE_KEY: Final = "synthetic-typesafe-key" +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") +_ANTHROPIC_VERSION: Final = MappingProxyType({"anthropic-version": "2023-06-01"}) +_SALT: Final = os.environ.get("LITELLM_SALT_KEY", "sk-integration-salt") +_TIER_CRITERIA: Final[dict[str, JsonValue]] = { + "SIMPLE": "Greetings, chitchat, or short factual lookups with known answers", + "MEDIUM": "Everyday requests needing explanation, light reasoning, or minor technical work", + "COMPLEX": "Non-trivial code, architecture, multi-step work, or specialized domain depth", + "REASONING": "Open-ended analysis, proofs, tradeoffs, or tasks requiring careful thought", +} +_PROBABILITIES: Final[dict[str, JsonValue]] = {"SIMPLE": 0.03, "MEDIUM": 0.06, "COMPLEX": 0.91} +_USAGE: Final[dict[str, JsonValue]] = {"input_tokens": 12, "output_tokens": 1} +_ANSWER: Final[dict[str, JsonValue]] = { + "answers": {"tier": {"type": "choice", "choice": "COMPLEX", "confidence": 0.91, "probabilities": _PROBABILITIES}}, + "usage": _USAGE, +} +_ROWS_QUERY: Final = ( + "SELECT request_id, status, model, call_type, cache_hit, custom_llm_provider, spend, " + "metadata->>'internal_call_origin' AS origin " + 'FROM "LiteLLM_SpendLogs" WHERE model_group = %s ORDER BY "startTime"' +) +_SIMPLE_TEXT: Final = "simple answer" +_COMPLEX_TEXT: Final = "complex answer" +_CALL_TYPES: Final = MappingProxyType( + {"/v1/chat/completions": "acompletion", "/v1/responses": "aresponses", "/v1/messages": "anthropic_messages"} +) + + +@dataclass(frozen=True, slots=True) +class _Routed: + name: str + identity: str + simple: str + complex_: str + + +@dataclass(frozen=True, slots=True) +class _Served: + route: str + stream: bool + status: int + identity: str + text: str + cause: str + classifier_cost: str | None + + +@dataclass(frozen=True, slots=True) +class _Failure: + label: str + respond: Callable[[Request], Reply] + overrides: Mapping[str, JsonValue] + reason: str + error_type: str + classifier_rows: int + + +def _number(value: JsonValue) -> float: + assert isinstance(value, (int, float)) and not isinstance(value, bool), value + return float(value) + + +def _prompt() -> str: + return f"design a distributed cache {uuid.uuid4().hex}" + + +def _short_prompt() -> str: + return f"hi {uuid.uuid4().hex[:8]}" + + +def _answering(body: Mapping[str, JsonValue]) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + return Reply(body=json.dumps(body).encode()) + + return respond + + +def _failing(status: int) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + return Reply(status=status, body=json.dumps({"error_code": "FAILED", "message": f"scripted {status}"}).encode()) + + return respond + + +def _slow(seconds: float) -> Callable[[Request], Reply]: + answer: Final = _answering(_ANSWER) + + def respond(request: Request) -> Reply: + time.sleep(seconds) + return answer(request) + + return respond + + +def _tier(prefix: str, text: str) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + body: Final = _JSON_OBJECT.validate_json(request.body) if request.body else {} + identity: Final = f"{prefix}-{uuid.uuid4().hex[:8]}" + if request.target == "/chat/completions": + return chat_reply(identity, "tier-model", text, stream=bool(body.get("stream"))) + if request.target == "/responses": + return responses_reply(identity, "tier-model", text, stream=bool(body.get("stream"))) + return Reply(status=404, body=json.dumps({"error": {"message": f"no route {request.target}"}}).encode()) + + return answering_model_discovery(respond) + + +@contextmanager +def _wires(classifier: Callable[[Request], Reply] | None = None) -> Iterator[tuple[Wire, Wire, Wire]]: + with ExitStack() as stack: + judge: Final = stack.enter_context(wire_server(classifier or _answering(_ANSWER))) + simple: Final = stack.enter_context(wire_server(_tier("simple", _SIMPLE_TEXT))) + complex_: Final = stack.enter_context(wire_server(_tier("complex", _COMPLEX_TEXT))) + yield judge, simple, complex_ + + +def _classifier_config(judge_url: str, **overrides: JsonValue) -> dict[str, JsonValue]: + return { + "provider": "databricks", + "model": _ENDPOINT, + "api_base": judge_url, + "api_key": _API_KEY, + "timeout_ms": 20000, + "circuit_breaker_enabled": False, + **overrides, + } + + +def _router_config(classifier: Mapping[str, JsonValue], simple: str, complex_: str) -> dict[str, JsonValue]: + return { + "classifier_type": "oss_classifier", + "opensource_classifier_config": dict(classifier), + "tiers": {"SIMPLE": simple, "MEDIUM": simple, "COMPLEX": complex_, "REASONING": complex_}, + } + + +def _deploy( + scenario: Scenario, judge: Wire, simple: Wire, complex_: Wire, classifier: Mapping[str, JsonValue] | None = None +) -> _Routed: + simple_model: Final = scenario.model(model="openai/simple-tier", api_base=simple.url, api_key="synthetic-tier-key") + complex_model: Final = scenario.model( + model="openai/complex-tier", api_base=complex_.url, api_key="synthetic-tier-key" + ) + name: Final = f"router-{uuid.uuid4().hex[:10]}" + created: Final = scenario.gateway.post( + "/model/new", + { + "model_name": name, + "litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_config": _router_config( + classifier if classifier is not None else _classifier_config(judge.url), simple_model, complex_model + ), + }, + "model_info": {}, + }, + ) + identity: Final = string_value(object_value(created["model_info"])["id"]) + scenario.cleanups.callback(scenario.delete_model, identity) + return _Routed(name=name, identity=identity, simple=simple_model, complex_=complex_model) + + +def _posts(wire: Wire) -> tuple[Request, ...]: + return tuple(request for request in wire.drain() if request.method == "POST") + + +def _state(prompt: str, system: str | None) -> str: + if system is None: + return f"\nClassify this message:\n{prompt}" + return f"\nCaller system prompt, quoted as task context:\n{system}\n\nClassify this message:\n{prompt}" + + +def _assert_classifier_call( + call: Request, *, prompt: str, system: str | None = None, endpoint: str = _ENDPOINT, key: str = _API_KEY +) -> None: + assert call.target == f"/{endpoint}/invocations", call.target + assert call.headers.get("authorization") == f"Bearer {key}", call.headers + assert json.loads(call.body) == { + "state": _state(prompt, system), + "model": endpoint, + "questions": {"tier": {"type": "choice", "instructions": DEFAULT_JEV_INSTRUCTIONS, "criteria": _TIER_CRITERIA}}, + }, call.body + + +def _assert_classified(headers: Mapping[str, str], *, router: str, cost: str | None = "0.0") -> None: + assert ( + headers.get("x-litellm-model-group"), + headers.get("x-litellm-model-name"), + headers.get("x-litellm-complexity-router-tier"), + headers.get("x-litellm-complexity-router-cause"), + headers.get("x-litellm-classifier-cost"), + ) == (router, "openai/complex-tier", "COMPLEX", "jev_classifier", cost), dict(headers) + + +def _assert_fell_back(headers: Mapping[str, str], *, router: str) -> None: + assert ( + headers.get("x-litellm-model-group"), + headers.get("x-litellm-model-name"), + headers.get("x-litellm-complexity-router-tier"), + headers.get("x-litellm-complexity-router-cause"), + headers.get("x-litellm-classifier-cost"), + ) == (router, "openai/simple-tier", "SIMPLE", "heuristic_scorer", None), dict(headers) + + +def _rows(router: str, *, want: int) -> tuple[dict[str, JsonValue], ...]: + return tuple(eventually(lambda: read_rows(_ROWS_QUERY, (router,)), lambda found: len(found) >= want, seconds=70)) + + +def _classifier_rows(rows: Sequence[Mapping[str, JsonValue]]) -> tuple[Mapping[str, JsonValue], ...]: + return tuple(row for row in rows if row["origin"] == "autorouter_classifier") + + +def _request_rows(rows: Sequence[Mapping[str, JsonValue]]) -> tuple[Mapping[str, JsonValue], ...]: + return tuple(row for row in rows if row["origin"] != "autorouter_classifier") + + +def _assert_classifier_row( + row: Mapping[str, JsonValue], *, model: str = _CLASSIFIER_MODEL, provider: str = "databricks", spend: float = 0.0 +) -> None: + assert (row["status"], row["call_type"], row["model"], row["custom_llm_provider"], row["cache_hit"]) == ( + "success", + "pass_through_endpoint", + model, + provider, + "False", + ), row + assert _number(row["spend"]) == pytest.approx(spend), row + + +def _issued_response_id(caller_id: str) -> str: + decrypted: Final = decrypt_if_encrypted_with(caller_id.removeprefix("resp_"), _SALT) + assert decrypted is not None, caller_id + return decrypted.split(";")[0].split("response_id:")[-1] + + +def _logged_ids(route: str, caller_id: str) -> frozenset[str]: + if route != "/v1/responses": + return frozenset({caller_id}) + return frozenset({caller_id, _issued_response_id(caller_id)}) + + +def _assert_logged_once(rows: Sequence[Mapping[str, JsonValue]], *, route: str, request_id: str) -> None: + (classifier_row,) = _classifier_rows(rows) + _assert_classifier_row(classifier_row) + (request_row,) = _request_rows(rows) + assert request_row["request_id"] in _logged_ids(route, request_id), (request_row, request_id) + assert (request_row["call_type"], request_row["status"]) == (_CALL_TYPES[route], "success"), request_row + + +def _data_frames(lines: Sequence[str]) -> tuple[dict[str, JsonValue], ...]: + return tuple( + _JSON_OBJECT.validate_json(line.removeprefix("data:").strip()) + for line in lines + if line.startswith("data:") and line.removeprefix("data:").strip() != "[DONE]" + ) + + +def _delta_texts(frames: Sequence[Mapping[str, JsonValue]]) -> Iterator[str]: + for frame in frames: + choices: Final = frame["choices"] + if not isinstance(choices, list): + continue + for choice in choices: + content: Final = object_value(object_value(choice)["delta"]).get("content") + if isinstance(content, str): + yield content + + +def _chat_stream_identity_and_text(frames: Sequence[Mapping[str, JsonValue]]) -> tuple[str, str]: + identities: Final = {string_value(frame["id"]) for frame in frames} + assert len(identities) == 1, identities + return identities.pop(), "".join(_delta_texts(frames)) + + +def _responses_stream_identity_and_text(frames: Sequence[Mapping[str, JsonValue]]) -> tuple[str, str]: + created: Final = [frame for frame in frames if frame["type"] == "response.created"] + assert len(created) == 1, frames + identity: Final = string_value(object_value(created[0]["response"])["id"]) + text: Final = "".join( + string_value(frame["delta"]) for frame in frames if frame["type"] == "response.output_text.delta" + ) + return identity, text + + +def _messages_stream_identity_and_text(frames: Sequence[Mapping[str, JsonValue]]) -> tuple[str, str]: + started: Final = [frame for frame in frames if frame["type"] == "message_start"] + assert len(started) == 1, frames + identity: Final = string_value(object_value(started[0]["message"])["id"]) + deltas: Final = [object_value(frame["delta"]) for frame in frames if frame["type"] == "content_block_delta"] + text: Final = "".join(string_value(delta["text"]) for delta in deltas if delta.get("type") == "text_delta") + return identity, text + + +def _body(route: str, model: str, prompt: str, *, stream: bool) -> dict[str, JsonValue]: + if route == "/v1/chat/completions": + return {"model": model, "messages": [{"role": "user", "content": prompt}], "stream": stream} + if route == "/v1/responses": + return {"model": model, "input": prompt, "stream": stream} + return {"model": model, "max_tokens": 64, "messages": [{"role": "user", "content": prompt}], "stream": stream} + + +def _identity_and_text(route: str, body: Mapping[str, JsonValue]) -> tuple[str, str]: + if route == "/v1/chat/completions": + (choice,) = body["choices"] if isinstance(body["choices"], list) else () + return string_value(body["id"]), string_value(object_value(object_value(choice)["message"])["content"]) + if route == "/v1/responses": + (item,) = body["output"] if isinstance(body["output"], list) else () + (content,) = object_value(item)["content"] if isinstance(object_value(item)["content"], list) else () + return string_value(body["id"]), string_value(object_value(content)["text"]) + (block,) = body["content"] if isinstance(body["content"], list) else () + return string_value(body["id"]), string_value(object_value(block)["text"]) + + +def _stream_identity_and_text(route: str, lines: Sequence[str]) -> tuple[str, str]: + frames: Final = _data_frames(lines) + if route == "/v1/chat/completions": + return _chat_stream_identity_and_text(frames) + if route == "/v1/responses": + return _responses_stream_identity_and_text(frames) + return _messages_stream_identity_and_text(frames) + + +async def _send(client: httpx.AsyncClient, route: str, model: str, prompt: str, *, stream: bool) -> _Served: + body: Final = _body(route, model, prompt, stream=stream) + headers: Final = dict(_ANTHROPIC_VERSION) if route == "/v1/messages" else {} + if not stream: + response: Final = await client.post(route, json=body, headers=headers) + assert response.status_code == 200, response.text + identity, text = _identity_and_text(route, _JSON_OBJECT.validate_json(response.content)) + return _Served( + route=route, + stream=False, + status=response.status_code, + identity=identity, + text=text, + cause=response.headers.get("x-litellm-complexity-router-cause", ""), + classifier_cost=response.headers.get("x-litellm-classifier-cost"), + ) + async with client.stream("POST", route, json=body, headers=headers) as streamed: + lines: Final = [line async for line in streamed.aiter_lines()] + assert streamed.status_code == 200, lines + identity, text = _stream_identity_and_text(route, lines) + return _Served( + route=route, + stream=True, + status=streamed.status_code, + identity=identity, + text=text, + cause=streamed.headers.get("x-litellm-complexity-router-cause", ""), + classifier_cost=streamed.headers.get("x-litellm-classifier-cost"), + ) + + +def _mixed_calls(count: int) -> tuple[tuple[str, bool], ...]: + routes: Final = tuple(_CALL_TYPES) + return tuple((routes[index % len(routes)], index % 2 == 1) for index in range(count)) + + +async def _burst( + base_url: str, + key: str, + model: str, + calls: Sequence[tuple[str, bool]], + prompts: Sequence[str], + *, + tolerate_transport_errors: bool = False, +) -> tuple[_Served, ...]: + async with httpx.AsyncClient( + base_url=base_url, timeout=90, trust_env=False, headers={"Authorization": f"Bearer {key}"} + ) as client: + results: Final = await asyncio.gather( + *(_send(client, route, model, prompt, stream=stream) for (route, stream), prompt in zip(calls, prompts)), + return_exceptions=tolerate_transport_errors, + ) + for result in results: + assert not isinstance(result, BaseException) or isinstance(result, httpx.TransportError), repr(result) + return tuple(result for result in results if isinstance(result, _Served)) + + +def _served_by_logged_id(served: Sequence[_Served]) -> Iterator[tuple[str, _Served]]: + for item in served: + for identity in _logged_ids(item.route, item.identity): + yield identity, item + + +def _assert_each_logged_once(rows: Sequence[Mapping[str, JsonValue]], served: Sequence[_Served]) -> None: + request_rows: Final = _request_rows(rows) + accepted: Final = dict(_served_by_logged_id(served)) + assert all(string_value(row["request_id"]) in accepted for row in request_rows), (request_rows, accepted) + matched: Final = tuple(accepted[string_value(row["request_id"])] for row in request_rows) + assert sorted(item.identity for item in matched) == sorted(item.identity for item in served), (matched, served) + for row, item in zip(request_rows, matched, strict=True): + assert (row["call_type"], row["status"]) == (_CALL_TYPES[item.route], "success"), (item, row) + + +def _test_routing(gateway: Gateway, config: Mapping[str, JsonValue], prompt: str) -> dict[str, JsonValue]: + response: Final = gateway.request( + "POST", "/auto_router/test_routing", {"prompt": prompt, "complexity_router_config": dict(config)} + ) + assert response.status_code == 200, response.text + return object_value(object_value(response.json())["routing_decision"]) + + +def _sdk_base(gateway: Gateway) -> str: + return str(gateway.client.base_url).rstrip("/") + + +def test_chat_completion_over_httpx_is_classified_by_the_databricks_endpoint_with_the_system_prompt_quoted( + gateway: Gateway, +) -> None: + prompt: Final = _prompt() + with _wires() as (judge, simple, complex_), gateway.scenario() as scenario: + routed: Final = _deploy(scenario, judge, simple, complex_) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + { + "model": routed.name, + "messages": [{"role": "system", "content": "be terse"}, {"role": "user", "content": prompt}], + }, + ) + assert response.status_code == 200, response.text + identity, text = _identity_and_text("/v1/chat/completions", _JSON_OBJECT.validate_json(response.content)) + assert text == _COMPLEX_TEXT, response.text + _assert_classified(response.headers, router=routed.name) + (classifier_call,) = _posts(judge) + _assert_classifier_call(classifier_call, prompt=prompt, system="be terse") + (tier_call,) = _posts(complex_) + assert tier_call.target == "/chat/completions", tier_call.target + tier_body: Final = _JSON_OBJECT.validate_json(tier_call.body) + assert tier_body["model"] == "complex-tier", tier_body + assert tier_body["messages"] == [ + {"role": "system", "content": "be terse"}, + {"role": "user", "content": prompt}, + ], tier_body + assert _posts(simple) == () + _assert_logged_once(_rows(routed.name, want=2), route="/v1/chat/completions", request_id=identity) + + +def test_chat_completion_stream_over_httpx_is_classified_and_logged_under_the_chunk_id(gateway: Gateway) -> None: + prompt: Final = _prompt() + with _wires() as (judge, simple, complex_), gateway.scenario() as scenario: + routed: Final = _deploy(scenario, judge, simple, complex_) + with gateway.client.stream( + "POST", + "/v1/chat/completions", + json=_body("/v1/chat/completions", routed.name, prompt, stream=True), + headers={"Authorization": f"Bearer {gateway.key}"}, + ) as streamed: + lines: Final = list(streamed.iter_lines()) + assert streamed.status_code == 200, lines + _assert_classified(streamed.headers, router=routed.name) + identity, text = _stream_identity_and_text("/v1/chat/completions", lines) + assert text == _COMPLEX_TEXT, lines + (classifier_call,) = _posts(judge) + _assert_classifier_call(classifier_call, prompt=prompt) + (tier_call,) = _posts(complex_) + assert _JSON_OBJECT.validate_json(tier_call.body)["stream"] is True, tier_call.body + assert _posts(simple) == () + _assert_logged_once(_rows(routed.name, want=2), route="/v1/chat/completions", request_id=identity) + + +def test_chat_completion_through_the_openai_sdk_is_classified_and_logged_once(gateway: Gateway) -> None: + prompt: Final = _prompt() + with _wires() as (judge, simple, complex_), gateway.scenario() as scenario: + routed: Final = _deploy(scenario, judge, simple, complex_) + client: Final = openai.OpenAI(base_url=f"{_sdk_base(gateway)}/v1", api_key=gateway.key, max_retries=0) + raw: Final = client.chat.completions.with_raw_response.create( + model=routed.name, messages=[{"role": "user", "content": prompt}] + ) + completion: Final = raw.parse() + assert completion.choices[0].message.content == _COMPLEX_TEXT, completion + _assert_classified(raw.headers, router=routed.name) + (classifier_call,) = _posts(judge) + _assert_classifier_call(classifier_call, prompt=prompt) + assert len(_posts(complex_)) == 1 and _posts(simple) == () + _assert_logged_once(_rows(routed.name, want=2), route="/v1/chat/completions", request_id=completion.id) + + +async def test_chat_completion_stream_through_the_async_openai_sdk_is_classified_and_logged_once( + gateway: Gateway, +) -> None: + prompt: Final = _prompt() + with _wires() as (judge, simple, complex_), gateway.scenario() as scenario: + routed: Final = _deploy(scenario, judge, simple, complex_) + client: Final = openai.AsyncOpenAI(base_url=f"{_sdk_base(gateway)}/v1", api_key=gateway.key, max_retries=0) + raw: Final = await client.chat.completions.with_raw_response.create( + model=routed.name, messages=[{"role": "user", "content": prompt}], stream=True + ) + chunks: Final = [chunk async for chunk in raw.parse()] + identities: Final = {chunk.id for chunk in chunks} + assert len(identities) == 1, identities + assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices) == _COMPLEX_TEXT + _assert_classified(raw.headers, router=routed.name) + (classifier_call,) = _posts(judge) + _assert_classifier_call(classifier_call, prompt=prompt) + assert len(_posts(complex_)) == 1 and _posts(simple) == () + _assert_logged_once(_rows(routed.name, want=2), route="/v1/chat/completions", request_id=identities.pop()) + + +def test_responses_request_over_httpx_is_classified_and_forwarded_to_the_tier_responses_route(gateway: Gateway) -> None: + prompt: Final = _prompt() + with _wires() as (judge, simple, complex_), gateway.scenario() as scenario: + routed: Final = _deploy(scenario, judge, simple, complex_) + response: Final = gateway.request("POST", "/v1/responses", {"model": routed.name, "input": prompt}) + assert response.status_code == 200, response.text + identity, text = _identity_and_text("/v1/responses", _JSON_OBJECT.validate_json(response.content)) + assert text == _COMPLEX_TEXT, response.text + _assert_classified(response.headers, router=routed.name) + (classifier_call,) = _posts(judge) + _assert_classifier_call(classifier_call, prompt=prompt) + (tier_call,) = _posts(complex_) + assert tier_call.target == "/responses", tier_call.target + tier_body: Final = _JSON_OBJECT.validate_json(tier_call.body) + assert (tier_body["model"], tier_body["input"]) == ("complex-tier", prompt), tier_body + assert _posts(simple) == () + _assert_logged_once(_rows(routed.name, want=2), route="/v1/responses", request_id=identity) + + +async def test_responses_stream_through_the_async_openai_sdk_is_classified_and_logged_under_the_created_id( + gateway: Gateway, +) -> None: + prompt: Final = _prompt() + with _wires() as (judge, simple, complex_), gateway.scenario() as scenario: + routed: Final = _deploy(scenario, judge, simple, complex_) + client: Final = openai.AsyncOpenAI(base_url=f"{_sdk_base(gateway)}/v1", api_key=gateway.key, max_retries=0) + raw: Final = await client.responses.with_raw_response.create(model=routed.name, input=prompt, stream=True) + events: Final = [event async for event in raw.parse()] + created: Final = [event for event in events if event.type == "response.created"] + assert len(created) == 1, [event.type for event in events] + text: Final = "".join(event.delta for event in events if event.type == "response.output_text.delta") + assert text == _COMPLEX_TEXT, [event.type for event in events] + _assert_classified(raw.headers, router=routed.name) + (classifier_call,) = _posts(judge) + _assert_classifier_call(classifier_call, prompt=prompt) + (tier_call,) = _posts(complex_) + assert tier_call.target == "/responses", tier_call.target + assert _posts(simple) == () + _assert_logged_once(_rows(routed.name, want=2), route="/v1/responses", request_id=created[0].response.id) + + +def test_messages_request_through_the_anthropic_sdk_is_classified_and_bridged_to_the_tier_responses_route( + gateway: Gateway, +) -> None: + prompt: Final = _prompt() + with _wires() as (judge, simple, complex_), gateway.scenario() as scenario: + routed: Final = _deploy(scenario, judge, simple, complex_) + client: Final = anthropic.Anthropic(base_url=_sdk_base(gateway), api_key=gateway.key, max_retries=0) + raw: Final = client.messages.with_raw_response.create( + model=routed.name, max_tokens=64, messages=[{"role": "user", "content": prompt}] + ) + message: Final = raw.parse() + assert [block.text for block in message.content if block.type == "text"] == [_COMPLEX_TEXT], message + _assert_classified(raw.headers, router=routed.name) + (classifier_call,) = _posts(judge) + _assert_classifier_call(classifier_call, prompt=prompt) + (tier_call,) = _posts(complex_) + assert tier_call.target == "/responses", tier_call.target + assert prompt in tier_call.body.decode(), tier_call.body + assert _posts(simple) == () + _assert_logged_once(_rows(routed.name, want=2), route="/v1/messages", request_id=message.id) + + +def test_messages_stream_over_httpx_is_classified_and_logged_under_the_message_start_id(gateway: Gateway) -> None: + prompt: Final = _prompt() + with _wires() as (judge, simple, complex_), gateway.scenario() as scenario: + routed: Final = _deploy(scenario, judge, simple, complex_) + with gateway.client.stream( + "POST", + "/v1/messages", + json=_body("/v1/messages", routed.name, prompt, stream=True), + headers={"Authorization": f"Bearer {gateway.key}", **_ANTHROPIC_VERSION}, + ) as streamed: + lines: Final = list(streamed.iter_lines()) + assert streamed.status_code == 200, lines + _assert_classified(streamed.headers, router=routed.name) + identity, text = _stream_identity_and_text("/v1/messages", lines) + assert text == _COMPLEX_TEXT, lines + (classifier_call,) = _posts(judge) + _assert_classifier_call(classifier_call, prompt=prompt) + (tier_call,) = _posts(complex_) + assert tier_call.target == "/responses", tier_call.target + assert _posts(simple) == () + _assert_logged_once(_rows(routed.name, want=2), route="/v1/messages", request_id=identity) + + +def test_a_non_catalog_endpoint_name_is_classified_without_a_cost_header_and_logged_at_zero(gateway: Gateway) -> None: + prompt: Final = _prompt() + endpoint: Final = "my-jev-endpoint" + with _wires() as (judge, simple, complex_), gateway.scenario() as scenario: + routed: Final = _deploy(scenario, judge, simple, complex_, _classifier_config(judge.url, model=endpoint)) + response: Final = gateway.request( + "POST", "/v1/chat/completions", _body("/v1/chat/completions", routed.name, prompt, stream=False) + ) + assert response.status_code == 200, response.text + identity, text = _identity_and_text("/v1/chat/completions", _JSON_OBJECT.validate_json(response.content)) + assert text == _COMPLEX_TEXT, response.text + _assert_classified(response.headers, router=routed.name, cost=None) + (classifier_call,) = _posts(judge) + _assert_classifier_call(classifier_call, prompt=prompt, endpoint=endpoint) + assert len(_posts(complex_)) == 1 and _posts(simple) == () + rows: Final = _rows(routed.name, want=2) + (classifier_row,) = _classifier_rows(rows) + _assert_classifier_row(classifier_row, model=f"databricks/{endpoint}") + (request_row,) = _request_rows(rows) + assert (request_row["request_id"], request_row["status"]) == (identity, "success"), request_row + + +def test_a_response_cache_hit_still_runs_the_classifier_and_logs_the_hit_row_at_zero(gateway: Gateway) -> None: + prompt: Final = _prompt() + with _wires() as (judge, simple, complex_), gateway.scenario() as scenario: + routed: Final = _deploy(scenario, judge, simple, complex_) + body: Final = _body("/v1/chat/completions", routed.name, prompt, stream=False) + first: Final = gateway.request("POST", "/v1/chat/completions", body) + second: Final = gateway.request("POST", "/v1/chat/completions", body) + assert (first.status_code, second.status_code) == (200, 200), (first.text, second.text) + first_identity, first_text = _identity_and_text( + "/v1/chat/completions", _JSON_OBJECT.validate_json(first.content) + ) + second_identity, second_text = _identity_and_text( + "/v1/chat/completions", _JSON_OBJECT.validate_json(second.content) + ) + assert (first_text, second_text) == (_COMPLEX_TEXT, _COMPLEX_TEXT) + assert second_identity == first_identity, (first_identity, second_identity) + assert "x-litellm-cache-key" not in first.headers, dict(first.headers) + assert "x-litellm-cache-key" in second.headers, dict(second.headers) + _assert_classified(first.headers, router=routed.name) + _assert_classified(second.headers, router=routed.name) + classifier_calls: Final = _posts(judge) + assert len(classifier_calls) == 2, classifier_calls + for call in classifier_calls: + _assert_classifier_call(call, prompt=prompt) + assert len(_posts(complex_)) == 1 and _posts(simple) == () + rows: Final = _rows(routed.name, want=4) + classifier_rows: Final = _classifier_rows(rows) + assert len(classifier_rows) == 2, rows + for row in classifier_rows: + _assert_classifier_row(row) + request_rows: Final = _request_rows(rows) + assert [row["request_id"] for row in request_rows if row["cache_hit"] != "True"] == [first_identity], rows + (hit_row,) = [row for row in request_rows if row["cache_hit"] == "True"] + assert string_value(hit_row["request_id"]).startswith(f"{first_identity}_cache_hit"), hit_row + assert (hit_row["status"], _number(hit_row["spend"])) == ("success", 0.0), hit_row + + +_FAILURES: Final = ( + _Failure("500", _failing(500), {}, "classifier_error", "MaskedHTTPStatusError", 0), + _Failure("401", _failing(401), {}, "classifier_error", "MaskedHTTPStatusError", 0), + _Failure( + "malformed", + _answering({"answers": {"tier": {"type": "choice", "choice": "COMPLEX", "confidence": 0.91}}, "usage": _USAGE}), + {}, + "invalid_response", + "ValidationError", + 1, + ), + _Failure( + "unknown-tier", + _answering( + { + "answers": { + "tier": { + "type": "choice", + "choice": "GALAXY", + "confidence": 0.91, + "probabilities": {"GALAXY": 0.91}, + } + }, + "usage": _USAGE, + } + ), + {}, + "classifier_error", + "ValueError", + 1, + ), + _Failure("timeout", _slow(3), {"timeout_ms": 1000}, "timeout", "TimeoutError", 0), +) + + +@pytest.mark.parametrize("failure", _FAILURES, ids=[failure.label for failure in _FAILURES]) +def test_a_failing_databricks_classifier_falls_back_to_the_heuristic_and_names_the_failure( + gateway: Gateway, failure: _Failure +) -> None: + prompt: Final = _short_prompt() + with _wires(failure.respond) as (judge, simple, complex_), gateway.scenario() as scenario: + classifier: Final = _classifier_config(judge.url, **failure.overrides) + routed: Final = _deploy(scenario, judge, simple, complex_, classifier) + decision: Final = _test_routing(gateway, _router_config(classifier, routed.simple, routed.complex_), prompt) + assert ( + decision["cause"], + decision["tier"], + decision["routed_model"], + decision["classifier_failure_reason"], + decision["classifier_error_type"], + ) == ("heuristic_scorer", "SIMPLE", routed.simple, failure.reason, failure.error_type), decision + assert "classifier_cost" not in decision, decision + response: Final = gateway.request( + "POST", "/v1/chat/completions", _body("/v1/chat/completions", routed.name, prompt, stream=False) + ) + assert response.status_code == 200, response.text + identity, text = _identity_and_text("/v1/chat/completions", _JSON_OBJECT.validate_json(response.content)) + assert text == _SIMPLE_TEXT, response.text + _assert_fell_back(response.headers, router=routed.name) + classifier_calls: Final = _posts(judge) + assert len(classifier_calls) == 2, classifier_calls + for call in classifier_calls: + _assert_classifier_call(call, prompt=prompt) + assert len(_posts(simple)) == 1 and _posts(complex_) == () + rows: Final = _rows(routed.name, want=1 + failure.classifier_rows) + (request_row,) = _request_rows(rows) + assert (request_row["request_id"], request_row["call_type"], request_row["status"]) == ( + identity, + "acompletion", + "success", + ), request_row + classifier_rows: Final = _classifier_rows(rows) + assert len(classifier_rows) == failure.classifier_rows, rows + for row in classifier_rows: + _assert_classifier_row(row) + + +def test_a_five_kilobyte_prompt_with_a_system_message_reaches_the_classifier_whole(gateway: Gateway) -> None: + prompt: Final = "p" * 5000 + uuid.uuid4().hex + with _wires() as (judge, simple, complex_), gateway.scenario() as scenario: + routed: Final = _deploy(scenario, judge, simple, complex_) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + { + "model": routed.name, + "messages": [{"role": "system", "content": "answer in one line"}, {"role": "user", "content": prompt}], + }, + ) + assert response.status_code == 200, response.text[:400] + identity, text = _identity_and_text("/v1/chat/completions", _JSON_OBJECT.validate_json(response.content)) + assert text == _COMPLEX_TEXT, response.text[:400] + _assert_classified(response.headers, router=routed.name) + (classifier_call,) = _posts(judge) + _assert_classifier_call(classifier_call, prompt=prompt, system="answer in one line") + assert len(_posts(complex_)) == 1 and _posts(simple) == () + _assert_logged_once(_rows(routed.name, want=2), route="/v1/chat/completions", request_id=identity) + + +def _typesafe_cost() -> float: + prices: Final = object_value( + json.loads(Path("model_prices_and_context_window.json").read_text())[f"typesafe/{_TYPESAFE_MODEL}"] + ) + return _number(_USAGE["input_tokens"]) * _number(prices["input_cost_per_token"]) + _number( + _USAGE["output_tokens"] + ) * _number(prices["output_cost_per_token"]) + + +def test_the_default_jev_provider_still_posts_to_systemone_and_bills_from_the_cost_map(gateway: Gateway) -> None: + prompt: Final = _prompt() + expected_cost: Final = _typesafe_cost() + with ( + _wires(_answering({**_ANSWER, "model": _TYPESAFE_MODEL})) as (judge, simple, complex_), + gateway.scenario() as scenario, + ): + classifier: Final[dict[str, JsonValue]] = { + "model": _TYPESAFE_MODEL, + "api_base": judge.url, + "api_key": _TYPESAFE_KEY, + "timeout_ms": 20000, + "circuit_breaker_enabled": False, + } + routed: Final = _deploy(scenario, judge, simple, complex_, classifier) + response: Final = gateway.request( + "POST", "/v1/chat/completions", _body("/v1/chat/completions", routed.name, prompt, stream=False) + ) + assert response.status_code == 200, response.text + identity, text = _identity_and_text("/v1/chat/completions", _JSON_OBJECT.validate_json(response.content)) + assert text == _COMPLEX_TEXT, response.text + assert ( + response.headers["x-litellm-model-name"], + response.headers["x-litellm-complexity-router-tier"], + response.headers["x-litellm-complexity-router-cause"], + ) == ("openai/complex-tier", "COMPLEX", "jev_classifier"), dict(response.headers) + assert float(response.headers["x-litellm-classifier-cost"]) == pytest.approx(expected_cost) + (classifier_call,) = _posts(judge) + assert classifier_call.target == "/v1/systemone", classifier_call.target + assert classifier_call.headers.get("authorization") == f"Bearer {_TYPESAFE_KEY}", classifier_call.headers + assert prompt in classifier_call.body.decode(), classifier_call.body + assert len(_posts(complex_)) == 1 and _posts(simple) == () + rows: Final = _rows(routed.name, want=2) + (classifier_row,) = _classifier_rows(rows) + _assert_classifier_row( + classifier_row, model=f"typesafe/{_TYPESAFE_MODEL}", provider="typesafe", spend=expected_cost + ) + (request_row,) = _request_rows(rows) + assert (request_row["request_id"], request_row["status"]) == (identity, "success"), request_row + + +async def test_a_classifier_outage_mid_traffic_falls_back_and_recovers_with_every_request_logged_once( + gateway: Gateway, +) -> None: + base_url: Final = _sdk_base(gateway) + calls: Final = _mixed_calls(12) + with ( + wire_server(_tier("simple", _SIMPLE_TEXT)) as simple, + wire_server(_tier("complex", _COMPLEX_TEXT)) as complex_, + gateway.scenario() as scenario, + ): + with wire_server(_answering(_ANSWER)) as judge: + routed: Final = _deploy(scenario, judge, simple, complex_) + config: Final = _router_config(_classifier_config(judge.url), routed.simple, routed.complex_) + port: Final = urlsplit(judge.url).port + assert port is not None + before_prompts: Final = tuple(_prompt() for _ in calls) + before: Final = await _burst(base_url, gateway.key, routed.name, calls, before_prompts) + assert [item.cause for item in before] == ["jev_classifier"] * 12, before + assert [item.classifier_cost for item in before] == ["0.0"] * 12, before + assert [item.text for item in before] == [_COMPLEX_TEXT] * 12, before + assert sorted(json.loads(call.body)["state"] for call in _posts(judge)) == sorted( + _state(prompt, None) for prompt in before_prompts + ) + during_prompts: Final = tuple(_short_prompt() for _ in calls) + during: Final = await _burst(base_url, gateway.key, routed.name, calls, during_prompts) + assert [item.cause for item in during] == ["heuristic_scorer"] * 12, during + assert [item.classifier_cost for item in during] == [None] * 12, during + assert [item.text for item in during] == [_SIMPLE_TEXT] * 12, during + decision: Final = _test_routing(gateway, config, _short_prompt()) + assert (decision["cause"], decision["classifier_failure_reason"]) == ( + "heuristic_scorer", + "classifier_error", + ), decision + with wire_server(_answering(_ANSWER), port=port) as recovered: + after_prompts: Final = tuple(_prompt() for _ in calls) + after: Final = await _burst(base_url, gateway.key, routed.name, calls, after_prompts) + assert [item.cause for item in after] == ["jev_classifier"] * 12, after + assert [item.text for item in after] == [_COMPLEX_TEXT] * 12, after + assert sorted(json.loads(call.body)["state"] for call in _posts(recovered)) == sorted( + _state(prompt, None) for prompt in after_prompts + ) + served: Final = (*before, *during, *after) + assert len({item.identity for item in served}) == 36, served + rows: Final = _rows(routed.name, want=60) + _assert_each_logged_once(rows, served) + classifier_rows: Final = _classifier_rows(rows) + assert len(classifier_rows) == 24, rows + for row in classifier_rows: + _assert_classifier_row(row) + + +async def test_a_slow_classifier_under_a_burst_times_out_every_call_into_the_heuristic_without_classifier_rows( + gateway: Gateway, +) -> None: + base_url: Final = _sdk_base(gateway) + calls: Final = tuple(("/v1/chat/completions", False) for _ in range(20)) + prompts: Final = tuple(_short_prompt() for _ in calls) + with _wires(_slow(3)) as (judge, simple, complex_), gateway.scenario() as scenario: + routed: Final = _deploy(scenario, judge, simple, complex_, _classifier_config(judge.url, timeout_ms=500)) + served: Final = await _burst(base_url, gateway.key, routed.name, calls, prompts) + assert [item.cause for item in served] == ["heuristic_scorer"] * 20, served + assert [item.text for item in served] == [_SIMPLE_TEXT] * 20, served + assert len({item.identity for item in served}) == 20, served + seen: Final[list[Request]] = [] + + def drained() -> int: + seen.extend(_posts(judge)) + return len(seen) + + eventually(drained, lambda count: count == 20, seconds=10) + assert sorted(json.loads(call.body)["state"] for call in seen) == sorted( + _state(prompt, None) for prompt in prompts + ) + assert len(_posts(simple)) == 20 and _posts(complex_) == () + rows: Final = _rows(routed.name, want=20) + _assert_each_logged_once(rows, served) + assert _classifier_rows(rows) == (), rows + + +def _worker_startups(log: Path) -> tuple[tuple[int, ...], int]: + text: Final = log.read_text() + return tuple(int(pid) for pid in _STARTED_WORKER.findall(text)), text.count("Application startup complete.") + + +def _open_upstream_connections(pid: int, upstream: str) -> int: + port: Final = urlsplit(upstream).port + return sum( + 1 + for connection in psutil.Process(pid).net_connections(kind="tcp") + if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port + ) + + +def _chaos_config(judge: Wire, simple: Wire, complex_: Wire, tmp_path: Path, router: str) -> Path: + config: Final = { + **_JSON_OBJECT.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())), + "model_list": [ + { + "model_name": "chaos-simple", + "litellm_params": { + "model": "openai/simple-tier", + "api_base": simple.url, + "api_key": "synthetic-tier-key", + }, + }, + { + "model_name": "chaos-complex", + "litellm_params": { + "model": "openai/complex-tier", + "api_base": complex_.url, + "api_key": "synthetic-tier-key", + }, + }, + { + "model_name": router, + "litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_config": _router_config( + _classifier_config(judge.url, timeout_ms=60000), "chaos-simple", "chaos-complex" + ), + }, + }, + ], + } + path: Final = tmp_path / "databricks-classifier-chaos.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +@pytest.mark.timeout(int(2 * graceful_stop_seconds() + 120 + 180)) +async def test_worker_sigkill_mid_burst_leaves_the_sibling_classifying_through_the_databricks_endpoint( + gateway: Gateway, tmp_path: Path +) -> None: + router: Final = f"chaos-router-{uuid.uuid4().hex[:8]}" + calls: Final = tuple(("/v1/chat/completions", False) for _ in range(20)) + prompts: Final = tuple(_prompt() for _ in calls) + release: Final = threading.Event() + held_states: Final[SimpleQueue[str]] = SimpleQueue() + answer: Final = _answering(_ANSWER) + + def held(request: Request) -> Reply: + held_states.put(string_value(_JSON_OBJECT.validate_json(request.body)["state"])) + assert release.wait(timeout=120), "The burst was never released" + return answer(request) + + with ( + wire_server(held) as judge, + wire_server(_tier("simple", _SIMPLE_TEXT)) as simple, + wire_server(_tier("complex", _COMPLEX_TEXT)) as complex_, + ): + path: Final = _chaos_config(judge, simple, complex_, tmp_path, router) + with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned: + candidate: Final = owned.gateway + base_url: Final = str(candidate.client.base_url) + workers, _ = eventually( + lambda: _worker_startups(owned.log), lambda found: len(found[0]) == 2 and found[1] == 2, seconds=120 + ) + burst: Final = asyncio.create_task( + _burst(base_url, candidate.key, router, calls, prompts, tolerate_transport_errors=True) + ) + await asyncio.to_thread(eventually, held_states.qsize, lambda size: size == 20, 90) + held_by: Final = MappingProxyType({pid: _open_upstream_connections(pid, judge.url) for pid in workers}) + assert sum(held_by.values()) == 20, held_by + victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__) + victim: Final = psutil.Process(victim_pid) + victim.suspend() + victim.send_signal(signal.SIGKILL) + release.set() + served: Final = await burst + assert len(served) == held_by[survivor_pid], (held_by, len(served)) + assert [item.cause for item in served] == ["jev_classifier"] * len(served), served + follow_up_prompt: Final = _prompt() + (answered,) = await _burst( + base_url, candidate.key, router, (("/v1/chat/completions", False),), (follow_up_prompt,) + ) + assert (answered.cause, answered.text) == ("jev_classifier", _COMPLEX_TEXT), answered + await asyncio.to_thread( + eventually, lambda: _worker_startups(owned.log), lambda found: len(found[0]) == 3 and found[1] == 3, 180 + ) + everything: Final = (*served, answered) + assert len({item.identity for item in everything}) == len(everything), everything + rows: Final = _rows(router, want=2 * len(everything)) + _assert_each_logged_once(rows, everything) + classifier_rows: Final = _classifier_rows(rows) + assert len(classifier_rows) == len(everything), rows + for row in classifier_rows: + _assert_classifier_row(row) + states: Final[set[str]] = set() + while not held_states.empty(): + states.add(held_states.get()) + assert states == {_state(prompt, None) for prompt in (*prompts, follow_up_prompt)}, states diff --git a/tests/unit/decisions/test_main.py b/tests/unit/decisions/test_main.py index b7e3f8687d8..99d7d59407f 100644 --- a/tests/unit/decisions/test_main.py +++ b/tests/unit/decisions/test_main.py @@ -22,6 +22,8 @@ from litellm.types.decisions import ( OpenAIDecisionResponse, ScoreAnswer, ) +from litellm.types.utils import LlmProviders +from litellm.utils import ProviderConfigManager _QUESTIONS: Final[Mapping[str, object]] = MappingProxyType( { @@ -68,6 +70,20 @@ _STRANDS_RESPONSE: Final[Mapping[str, object]] = { "usage": {"input_tokens": 216, "output_tokens": 3}, "latency_ms": 3722.17, } +_DATABRICKS_ENDPOINT: Final = "databricks-openjev-qwen35-4b" +_DATABRICKS_RESPONSE: Final[Mapping[str, object]] = { + "model": "/mosaicml/local_model", + "answers": { + "is_defect": {"type": "noul", "noul": 0.9}, + "sentiment": { + "type": "choice", + "choice": "positive", + "confidence": 0.8, + "probabilities": {"positive": 0.8, "negative": 0.2}, + }, + }, + "usage": {"input_tokens": 214, "output_tokens": 0}, +} _PROVIDERS: Final[tuple[tuple[str, str, str, str], ...]] = ( ( "perplexity", @@ -867,6 +883,164 @@ async def test_hosted_vllm_choice_question_goes_out_as_a_jev_choice_with_the_ser } +@pytest.mark.asyncio +async def test_databricks_requires_api_base_before_http( + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + monkeypatch.delenv("DATABRICKS_API_BASE", raising=False) + monkeypatch.setenv("DATABRICKS_API_KEY", "dapi-key") + + with pytest.raises(litellm.BadRequestError, match="DATABRICKS_API_BASE"): + await litellm.adecisions( + model=f"databricks/{_DATABRICKS_ENDPOINT}", + state="review", + questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + ) + + assert len(respx_mock.calls) == 0 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "endpoint_name", ["serving-endpoints/openjev", "openjev?x=1", "openjev#frag", "..", "open jev"] +) +async def test_databricks_rejects_a_serving_endpoint_name_that_would_rewrite_the_url( + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, + endpoint_name: str, +) -> None: + monkeypatch.setenv("DATABRICKS_API_BASE", "https://workspace.example/serving-endpoints") + monkeypatch.setenv("DATABRICKS_API_KEY", "dapi-key") + + with pytest.raises(litellm.BadRequestError, match="bare endpoint name"): + await litellm.adecisions( + model=f"databricks/{endpoint_name}", + state="review", + questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + ) + + assert len(respx_mock.calls) == 0 + + +@pytest.mark.asyncio +async def test_databricks_requires_a_key_before_http( + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + monkeypatch.setenv("DATABRICKS_API_BASE", "https://workspace.example/serving-endpoints") + monkeypatch.delenv("DATABRICKS_API_KEY", raising=False) + monkeypatch.delenv("DATABRICKS_TOKEN", raising=False) + + with pytest.raises(litellm.AuthenticationError, match="Missing API key for Decisions provider 'databricks'"): + await litellm.adecisions( + model=f"databricks/{_DATABRICKS_ENDPOINT}", + state="review", + questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + ) + + assert len(respx_mock.calls) == 0 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "api_base", ["https://workspace.example/serving-endpoints", "https://workspace.example/serving-endpoints/"] +) +async def test_databricks_posts_the_jev_body_to_the_serving_endpoint_invocations_route( + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, + api_base: str, +) -> None: + monkeypatch.setenv("DATABRICKS_API_BASE", api_base) + monkeypatch.delenv("DATABRICKS_API_KEY", raising=False) + monkeypatch.setenv("DATABRICKS_TOKEN", "dapi-token") + route: Final = respx_mock.post( + f"https://workspace.example/serving-endpoints/{_DATABRICKS_ENDPOINT}/invocations" + ).respond(json=_DATABRICKS_RESPONSE) + questions: Final = { + "is_defect": {"type": "noul", "instructions": "Is this a defect?"}, + "sentiment": {"type": "choice", "criteria": {"positive": None, "negative": "unhappy"}}, + } + + response: Final = await litellm.adecisions( + model=f"databricks/{_DATABRICKS_ENDPOINT}", + state={"ticket": "export hangs"}, + questions=questions, + ) + + assert route.called + sent: Final = respx_mock.calls[0].request + assert sent.headers["authorization"] == "Bearer dapi-token" + assert json.loads(sent.content) == { + "model": _DATABRICKS_ENDPOINT, + "state": {"ticket": "export hangs"}, + "questions": questions, + } + assert response.model == "/mosaicml/local_model" + assert isinstance(response.answers["is_defect"], NoulAnswer) + assert isinstance(response.answers["sentiment"], ChoiceAnswer) + assert response.hidden_params["custom_llm_provider"] == "databricks" + assert response.hidden_params["model"] == f"databricks/{_DATABRICKS_ENDPOINT}" + + +@pytest.mark.asyncio +async def test_databricks_router_deployment_sends_its_own_key_to_its_own_base( + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + monkeypatch.setenv("DATABRICKS_API_BASE", "https://other-workspace.example/serving-endpoints") + monkeypatch.setenv("DATABRICKS_API_KEY", "env-key-stays-home") + router: Final = litellm.Router( + model_list=[ + { + "model_name": "decider", + "litellm_params": { + "model": f"databricks/{_DATABRICKS_ENDPOINT}", + "api_base": "https://workspace.example/serving-endpoints", + "api_key": "deployment-key", + }, + } + ] + ) + route: Final = respx_mock.post( + f"https://workspace.example/serving-endpoints/{_DATABRICKS_ENDPOINT}/invocations" + ).respond(json=_DATABRICKS_RESPONSE) + + response: Final = await router.adecisions( + model="decider", + state="review", + questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + ) + + assert route.called + assert respx_mock.calls[0].request.headers["authorization"] == "Bearer deployment-key" + assert response.answers["is_defect"] == NoulAnswer(type="noul", noul=0.9) + + +_PROVIDERS_WHOSE_CHAT_MODELS_ALSO_DECIDE: Final = frozenset( + {LlmProviders.OPENAI.value, LlmProviders.OPENROUTER.value, LlmProviders.HOSTED_VLLM.value} +) +_DEDICATED_DECIDER_PROVIDERS: Final[tuple[str, ...]] = tuple( + sorted( + provider.value + for provider in LlmProviders + if provider.value not in _PROVIDERS_WHOSE_CHAT_MODELS_ALSO_DECIDE + and ProviderConfigManager.get_provider_decisions_config(model="", provider=provider) is not None + ) +) + + +@pytest.mark.parametrize("provider", _DEDICATED_DECIDER_PROVIDERS) +def test_every_dedicated_decider_provider_ships_an_evaluation_mode_cost_map_entry(provider: str) -> None: + evaluation_entries: Final = tuple( + name + for name, info in litellm.model_cost.items() + if isinstance(info, Mapping) and info.get("litellm_provider") == provider and info.get("mode") == "evaluation" + ) + + assert evaluation_entries, f"a {provider} decider would be health-checked as chat without a mode: evaluation entry" + + @pytest.mark.asyncio @pytest.mark.parametrize( ("api_base", "url"), diff --git a/tests/unit/llms/databricks/decisions/__init__.py b/tests/unit/llms/databricks/decisions/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/databricks/decisions/test_transformation.py b/tests/unit/llms/databricks/decisions/test_transformation.py new file mode 100644 index 00000000000..e6de1d1a635 --- /dev/null +++ b/tests/unit/llms/databricks/decisions/test_transformation.py @@ -0,0 +1,56 @@ +from typing import Final + +import pytest + +from litellm.llms.databricks.decisions.transformation import DatabricksDecisionsConfig + +_BASE: Final = "https://workspace.example/serving-endpoints" +_CONFIG: Final = DatabricksDecisionsConfig() + + +@pytest.mark.parametrize("api_base", [_BASE, f"{_BASE}/"]) +def test_the_complete_url_is_the_serving_endpoint_invocations_route(api_base: str) -> None: + assert _CONFIG.get_complete_url(api_base, "openjev") == f"{_BASE}/openjev/invocations" + + +@pytest.mark.parametrize( + "name", + ["", "serving-endpoints/openjev", "openjev?x=1", "openjev#frag", ".", "..", ".openjev", "open jev", "openjev%2Fx"], +) +def test_a_name_that_would_rewrite_the_url_is_rejected_before_any_request(name: str) -> None: + with pytest.raises(ValueError, match="bare endpoint name"): + _ = _CONFIG.get_complete_url(_BASE, name) + with pytest.raises(ValueError, match="bare endpoint name"): + _ = _CONFIG.canonical_model(name) + + +def test_connection_prefers_the_configured_values_over_the_environment(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("DATABRICKS_API_BASE", "https://env.example/serving-endpoints") + monkeypatch.setenv("DATABRICKS_API_KEY", "dapi-env") + connection: Final = _CONFIG.connection(f"{_BASE}/", "dapi-configured") + assert (connection.api_base, connection.api_key) == (_BASE, "dapi-configured") + + +def test_connection_falls_back_to_the_workspace_token_when_the_api_key_variable_is_unset( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("DATABRICKS_API_BASE", _BASE) + monkeypatch.delenv("DATABRICKS_API_KEY", raising=False) + monkeypatch.setenv("DATABRICKS_TOKEN", "dapi-token") + connection: Final = _CONFIG.connection(None, None) + assert (connection.api_base, connection.api_key) == (_BASE, "dapi-token") + + +def test_connection_never_sends_the_environment_key_to_a_configured_base(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("DATABRICKS_API_BASE", _BASE) + monkeypatch.setenv("DATABRICKS_API_KEY", "dapi-env") + with pytest.raises(ValueError, match="DATABRICKS_API_BASE"): + _ = _CONFIG.connection("https://collector.example/serving-endpoints", None) + + +def test_classifier_response_accounts_the_serving_endpoint_instead_of_the_container_model() -> None: + body: Final = {"model": "/mosaicml/local_model", "answers": {}, "usage": {"input_tokens": 3, "output_tokens": 0}} + normalized: Final = _CONFIG.classifier_response(body, "openjev") + assert normalized["model"] == "openjev" + assert normalized["usage"] == body["usage"] + assert body["model"] == "/mosaicml/local_model" diff --git a/tests/unit/proxy/common_utils/test_registry_read_through.py b/tests/unit/proxy/common_utils/test_registry_read_through.py index 35f448c4fcf..dd0dca328ac 100644 --- a/tests/unit/proxy/common_utils/test_registry_read_through.py +++ b/tests/unit/proxy/common_utils/test_registry_read_through.py @@ -203,9 +203,7 @@ def clean_agent_registry(): @pytest.mark.asyncio -async def test_get_agent_with_read_through_recovers_agent_created_on_sibling_replica( - clean_agent_registry, monkeypatch -): +async def test_get_agent_with_read_through_recovers_agent_created_on_sibling_replica(clean_agent_registry, monkeypatch): from unittest.mock import AsyncMock, MagicMock import litellm.proxy.proxy_server as proxy_server @@ -601,7 +599,9 @@ async def test_resync_guardrails_syncs_decrypted_litellm_params(monkeypatch): synced: list[dict] = [] monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) monkeypatch.setattr(proxy_server, "store_model_in_db", True) - monkeypatch.setattr(IN_MEMORY_GUARDRAIL_HANDLER, "sync_guardrail_from_db", lambda guardrail: synced.append(guardrail)) + monkeypatch.setattr( + IN_MEMORY_GUARDRAIL_HANDLER, "sync_guardrail_from_db", lambda guardrail: synced.append(guardrail) + ) monkeypatch.setattr(read_through_module, "_initialized_guardrail", lambda guardrail_name: MagicMock()) assert await _resync_guardrails("enc-guardrail") is True @@ -610,15 +610,21 @@ async def test_resync_guardrails_syncs_decrypted_litellm_params(monkeypatch): @pytest.mark.asyncio @pytest.mark.parametrize("lookup", ["agent-id", "Agent name"]) -async def test_agent_read_through_hydrates_identity_binding(lookup, clean_agent_registry, fresh_agent_read_through, monkeypatch): +async def test_agent_read_through_hydrates_identity_binding( + lookup, clean_agent_registry, fresh_agent_read_through, monkeypatch +): from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock from litellm.proxy.common_utils.registry_read_through import get_agent_with_read_through binding = { - "agent_id": "agent-id", "provider": "microsoft_entra", "tenant_id": "tenant", "client_id": "client", - "issuer": "https://login.microsoftonline.com/tenant/v2.0", "revision": "revision", + "agent_id": "agent-id", + "provider": "microsoft_entra", + "tenant_id": "tenant", + "client_id": "client", + "issuer": "https://login.microsoftonline.com/tenant/v2.0", + "revision": "revision", } async def load_row(*, where, include): @@ -733,3 +739,49 @@ async def test_agent_read_through_answers_a_loaded_agent_without_reading_the_db( assert await agent_registry_read_through.attempt("wired-loaded-agent-id") is True assert await agent_registry_read_through.attempt("wired-loaded-agent") is True + + +@pytest.mark.asyncio +@pytest.mark.parametrize("router_holds_the_name_after_reconcile", [True, False]) +async def test_resync_model_deployments_reconciles_every_db_row_when_the_missed_model_is_an_auto_router( + monkeypatch: pytest.MonkeyPatch, router_holds_the_name_after_reconcile: bool +) -> None: + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy import proxy_server + from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper + from litellm.proxy.common_utils.registry_read_through import _resync_model_deployments + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-auto-router-read-through") + router_name: Final = "auto-router-created-on-a-sibling-replica" + row: Final = MagicMock() + row.model_name = router_name + row.litellm_params = { + "model": encrypt_value_helper("auto_router/complexity-router"), + "auto_router_config": {"model_tiers": {"SIMPLE": "simple-tier", "COMPLEX": "complex-tier"}}, + } + prisma_client: Final = MagicMock() + prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[row]) + router: Final = MagicMock() + router.model_names = [] + router.has_model_id.return_value = False + reconciled: Final = MagicMock() + + async def reconcile_every_row(prisma_client: object, proxy_logging_obj: object) -> None: + reconciled(prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj) + router.model_names = ( + [router_name, "simple-tier", "complex-tier"] if router_holds_the_name_after_reconcile else [] + ) + + def reject_single_row_path(db_models: object) -> None: + raise AssertionError(f"an auto router row took the single-row path: {db_models}") + + monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) + monkeypatch.setattr(proxy_server, "store_model_in_db", True) + monkeypatch.setattr(proxy_server, "llm_router", router) + monkeypatch.setattr(proxy_server, "llm_model_list", None) + monkeypatch.setattr(proxy_server.proxy_config, "add_deployment", reconcile_every_row) + monkeypatch.setattr(proxy_server.proxy_config, "_add_deployment", reject_single_row_path) + + assert await _resync_model_deployments(router_name) is router_holds_the_name_after_reconcile + reconciled.assert_called_once_with(prisma_client=prisma_client, proxy_logging_obj=proxy_server.proxy_logging_obj) diff --git a/tests/unit/proxy/decisions_endpoints/test_endpoints.py b/tests/unit/proxy/decisions_endpoints/test_endpoints.py index ccfa12e10d6..6ce15636f1e 100644 --- a/tests/unit/proxy/decisions_endpoints/test_endpoints.py +++ b/tests/unit/proxy/decisions_endpoints/test_endpoints.py @@ -257,6 +257,51 @@ def test_proxy_decisions_dispatches_strands_decider( assert "authorization" not in upstream.calls[0].request.headers +def test_proxy_systemone_dispatches_databricks_serving_endpoint( + client: TestClient, + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + monkeypatch.delenv("DATABRICKS_API_BASE", raising=False) + monkeypatch.delenv("DATABRICKS_API_KEY", raising=False) + monkeypatch.delenv("DATABRICKS_TOKEN", raising=False) + router: Final = litellm.Router( + model_list=[ + { + "model_name": "databricks-decider", + "litellm_params": { + "model": "databricks/databricks-openjev-qwen35-4b", + "api_base": "https://workspace.example/serving-endpoints", + "api_key": "dapi-deployment", + }, + } + ] + ) + monkeypatch.setattr(litellm.proxy.proxy_server, "llm_router", router) + upstream: Final = respx_mock.post( + "https://workspace.example/serving-endpoints/databricks-openjev-qwen35-4b/invocations" + ).respond(json={**_RESPONSE, "model": "/mosaicml/local_model"}) + + response: Final = client.post( + "/v1/systemone", + json={ + "model": "databricks-decider", + "state": {"source": "proxy-test"}, + "questions": {"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + }, + ) + + assert response.status_code == 200, response.text + assert response.json()["answers"] == _RESPONSE["answers"] + assert upstream.called + assert json.loads(upstream.calls[0].request.content) == { + "model": "databricks-openjev-qwen35-4b", + "state": {"source": "proxy-test"}, + "questions": {"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + } + assert upstream.calls[0].request.headers["authorization"] == "Bearer dapi-deployment" + + def test_proxy_decisions_without_model_uses_the_proxy_default_model( client: TestClient, monkeypatch: pytest.MonkeyPatch, diff --git a/tests/unit/proxy/management_helpers/test_auto_router_permissions.py b/tests/unit/proxy/management_helpers/test_auto_router_permissions.py index 1ccfbab7b1f..b19445c95d9 100644 --- a/tests/unit/proxy/management_helpers/test_auto_router_permissions.py +++ b/tests/unit/proxy/management_helpers/test_auto_router_permissions.py @@ -146,6 +146,8 @@ def test_tier_config_is_normalized_and_unknown_router_extras_are_rejected() -> N ({"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"), + ({"provider": "databricks", "model": "my-openjev", "api_base": "https://collector.invalid"}, "opensource_classifier_config"), + ({"provider": "databricks", "model": "my-openjev", "api_key": "dapi-member"}, "api_key"), ], ) @pytest.mark.parametrize("legacy", [False, True]) @@ -164,7 +166,10 @@ 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"), ("bespoke", "nimble-latest")]) +@pytest.mark.parametrize( + ("provider", "model"), + [("typesafe", "jev-preview"), ("laya", "english"), ("bespoke", "nimble-latest"), ("databricks", "my-openjev")], +) @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( diff --git a/tests/unit/router_strategy/complexity_router/test_jev_classifier.py b/tests/unit/router_strategy/complexity_router/test_jev_classifier.py index 418bf522b6a..fea283cbe9c 100644 --- a/tests/unit/router_strategy/complexity_router/test_jev_classifier.py +++ b/tests/unit/router_strategy/complexity_router/test_jev_classifier.py @@ -451,7 +451,14 @@ 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"), ("bespoke", "nimble-latest", "bespoke")], + [ + (None, "jev-latest", "jev"), + ("typesafe", "jev-latest", "jev"), + ("jev", "jev-latest", "jev"), + ("laya", "english", "laya"), + ("bespoke", "nimble-latest", "bespoke"), + ("databricks", "databricks-openjev-qwen35-4b", "databricks"), + ], ) def test_classifier_aliases_load_and_serialize_one_canonical_config( classifier_type: str, config_key: str, provider: str | None, model: str, canonical_provider: str @@ -533,6 +540,85 @@ async def test_oss_routes_with_its_own_credentials_and_accounts_the_checkpoint( assert recorder.calls[0]["response_cost"] == pytest.approx(0.31) +def test_databricks_requires_the_bare_serving_endpoint_name_as_the_model() -> None: + with pytest.raises(ValueError, match="model is required for provider 'databricks'"): + JevClassifierConfig.model_validate({"provider": "databricks"}) + with pytest.raises(ValueError, match="bare endpoint name"): + JevClassifierConfig.model_validate({"provider": "databricks", "model": "serving-endpoints/my-openjev"}) + with pytest.raises(ValueError, match="bare endpoint name"): + JevClassifierConfig.model_validate({"provider": "databricks", "model": ".."}) + assert JevClassifierConfig.model_validate({"provider": "databricks", "model": "my-openjev"}).model == "my-openjev" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("configured_key", [False, True]) +@pytest.mark.parametrize("legacy", [False, True]) +async def test_databricks_routes_through_the_serving_endpoint_and_accounts_the_endpoint_model( + monkeypatch: pytest.MonkeyPatch, configured_key: bool, legacy: bool +) -> None: + endpoint: Final = "databricks-openjev-qwen35-4b" + monkeypatch.setenv("TYPESAFE_API_KEY", "never-send-typesafe-key") + monkeypatch.setenv("DATABRICKS_API_BASE", "https://workspace.test/serving-endpoints") + monkeypatch.delenv("DATABRICKS_API_KEY", raising=False) + monkeypatch.setenv("DATABRICKS_TOKEN", "dapi-env-token") + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + monkeypatch.setitem(litellm.model_cost, f"databricks/{endpoint}", {"input_cost_per_token": 0.01}) + recorder: Final = _UsageRecorder(f"databricks/{endpoint}") + monkeypatch.setattr(litellm, "_async_success_callback", [recorder]) + router: Final = ComplexityRouter( + "databricks-route", + litellm.Router(model_list=[]), + { + "classifier_type": "jev" if legacy else "oss_classifier", + "jev_classifier_config" if legacy else "opensource_classifier_config": { + "provider": "databricks", + "model": endpoint, + **({"api_key": "dapi-configured", "api_base": "https://workspace.test/serving-endpoints"} if configured_key else {}), + }, + "tiers": {"SIMPLE": "cheap"}, + }, + derive_savings_baseline=False, + ) + with respx.mock(assert_all_called=True) as upstream: + route: Final = upstream.post(f"https://workspace.test/serving-endpoints/{endpoint}/invocations").respond( + 200, + json={ + "model": "/mosaicml/local_model", + "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) == ("databricks", endpoint) + assert outcome.classifier_cost == pytest.approx(0.31) + sent: Final = route.calls.last.request + assert sent.headers.get("authorization") == ("Bearer dapi-configured" if configured_key else "Bearer dapi-env-token") + assert json.loads(sent.content)["model"] == endpoint + assert len(recorder.calls) == 1 + assert recorder.calls[0]["response_cost"] == pytest.approx(0.31) + + +def test_databricks_without_a_base_or_key_fails_at_client_build(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("DATABRICKS_API_BASE", raising=False) + monkeypatch.delenv("DATABRICKS_API_KEY", raising=False) + monkeypatch.delenv("DATABRICKS_TOKEN", raising=False) + with pytest.raises(ValueError, match="DATABRICKS_API_BASE"): + ComplexityRouter( + "databricks-route", + litellm.Router(model_list=[]), + { + "classifier_type": "oss_classifier", + "opensource_classifier_config": {"provider": "databricks", "model": "databricks-openjev-qwen35-4b"}, + "tiers": {"SIMPLE": "cheap"}, + }, + derive_savings_baseline=False, + ) + + def test_jev_config_is_rejected_for_other_classifier_types() -> None: with pytest.raises(ValueError, match="has no effect"): ComplexityRouterConfig.model_validate( @@ -570,6 +656,22 @@ def test_jev_api_base_without_its_own_key_is_rejected_so_the_environment_key_sta assert JevClassifierConfig(api_key="sk-own").api_base is None +def test_databricks_api_base_without_its_own_key_is_rejected_so_the_workspace_token_stays_home() -> None: + with pytest.raises(ValueError, match="DATABRICKS_API_KEY or DATABRICKS_TOKEN is only sent to DATABRICKS_API_BASE"): + JevClassifierConfig.model_validate( + {"provider": "databricks", "model": "my-openjev", "api_base": "https://collector.invalid/serving-endpoints"} + ) + paired: Final = JevClassifierConfig.model_validate( + { + "provider": "databricks", + "model": "my-openjev", + "api_base": "https://workspace.invalid/serving-endpoints", + "api_key": "dapi-own", + } + ) + assert (paired.api_base, paired.api_key) == ("https://workspace.invalid/serving-endpoints", "dapi-own") + + @pytest.mark.parametrize( ("probabilities", "confidence"), [ diff --git a/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.tsx b/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.tsx index a8d492a6ced..f21dc3d795a 100644 --- a/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.tsx @@ -16,7 +16,11 @@ import { type ClassifierType, type ComplexityRouterConfigValue, } from "./ComplexityRouterConfig"; -import { defaultJevClassifierConfig, normalizeJevClassifierConfig } from "./jev_classifier_config"; +import { + defaultJevClassifierConfig, + isOssClassifierProvider, + normalizeJevClassifierConfig, +} from "./jev_classifier_config"; import { transitionClassifierType } from "./classifier_type_transition"; import { isForecastClassifier } from "./forecast_classifier_config"; import { @@ -150,7 +154,7 @@ const AutoRouterClassifierTabs: React.FC = ({ val if (next === "jev") changeType("jev"); }; const changeProvider = (provider: unknown) => { - if (provider !== "jev" && provider !== "laya" && provider !== "bespoke") return; + if (!isOssClassifierProvider(provider)) return; const defaults = defaultJevClassifierConfig(provider); onChange({ ...value, @@ -173,7 +177,11 @@ const AutoRouterClassifierTabs: React.FC = ({ 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, Laya, or Bespoke Nimble to choose a tier" }, + { + value: "jev", + label: "OSS Classifier", + description: "Use Jev, Laya, Bespoke Nimble, or Databricks to choose a tier", + }, ].map((option) => ( + )} diff --git a/ui/litellm-dashboard/src/components/add_model/ClassifierTypeRadios.tsx b/ui/litellm-dashboard/src/components/add_model/ClassifierTypeRadios.tsx index 2eda2c63f64..9deb53753bd 100644 --- a/ui/litellm-dashboard/src/components/add_model/ClassifierTypeRadios.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ClassifierTypeRadios.tsx @@ -55,7 +55,9 @@ const ClassifierTypeRadios: React.FC = ({ value, clas OSS Classifier{" "} - uses Jev or Laya to decide the tier + + uses Jev, Laya, Bespoke, or a Databricks endpoint to decide the tier + diff --git a/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.integration.test.tsx b/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.integration.test.tsx index 99726a8f071..15a736de7b2 100644 --- a/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.integration.test.tsx @@ -101,9 +101,11 @@ describe("JEV classifier editor", () => { ["jev", "Jev", "jev-test"], ["laya", "Laya", "multilingual"], ["bespoke", "Bespoke Nimble", "bespokelabs/Bespoke-Nimble-9B"], + ["databricks", "Databricks", "databricks-openjev-qwen35-4b"], ] as const)( "preserves %s, custom tiers and context through save, reload and probe", async (provider, label, model) => { + const typesTheModel = provider === "jev" || provider === "databricks"; renderWithProviders(
); expect(screen.getByLabelText("Judge model")).toBeInTheDocument(); expect(screen.getByText("Reasoning Effort")).toBeInTheDocument(); @@ -123,10 +125,11 @@ describe("JEV classifier editor", () => { expect(screen.getByLabelText("Classifier Model")).toHaveValue("jev-latest"); fireEvent.click(screen.getByRole("radio", { name: label })); if (provider === "bespoke") expect(screen.getByLabelText("Classifier Model")).toHaveTextContent("nimble-latest"); - if (provider !== "jev") { - await chooseSelectOption(userEvent, screen.getByLabelText("Classifier Model"), model); - } else { + if (provider === "databricks") expect(screen.getByLabelText("Classifier Model")).toHaveValue(""); + if (typesTheModel) { fireEvent.change(screen.getByLabelText("Classifier Model"), { target: { value: model } }); + } else { + await chooseSelectOption(userEvent, screen.getByLabelText("Classifier Model"), model); } fireEvent.change(screen.getByLabelText("Classifier Timeout (ms)"), { target: { value: "4200" } }); fireEvent.change(screen.getByLabelText("Context Window Size"), { target: { value: "6" } }); @@ -136,8 +139,8 @@ describe("JEV classifier editor", () => { fireEvent.click(screen.getByRole("button", { name: "Save and reload" })); expect(screen.getByRole("radio", { name: /^OSS Classifier$/ })).toBeChecked(); 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); + if (typesTheModel) expect(screen.getByLabelText("Classifier Model")).toHaveValue(model); + else expect(screen.getByLabelText("Classifier Model")).toHaveTextContent(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(); diff --git a/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.tsx b/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.tsx index 3818e3ae9c2..86d2d0f8348 100644 --- a/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.tsx +++ b/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.tsx @@ -7,13 +7,15 @@ 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, OSS_CLASSIFIER_MODELS } from "./jev_classifier_config"; +import { defaultJevClassifierConfig, fixedClassifierModels } 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.", + databricks: + "Uses a Databricks ai_decide serving endpoint with your configured tiers. Set DATABRICKS_API_BASE and DATABRICKS_API_KEY on the gateway, and enter the serving endpoint name as the classifier model.", }; export default function JevClassifierConfig({ @@ -25,7 +27,7 @@ export default function JevClassifierConfig({ }) { const id = useId(); const config = value.jev_classifier_config ?? defaultJevClassifierConfig(); - const models = config.provider && config.provider !== "jev" ? OSS_CLASSIFIER_MODELS[config.provider] : undefined; + const models = fixedClassifierModels(config.provider); const update = (patch: Partial) => onChange({ ...value, jev_classifier_config: { ...config, ...patch } }); @@ -48,7 +50,14 @@ export default function JevClassifierConfig({ ) : ( - update({ model: event.target.value })} /> + update({ model: event.target.value })} + /> )}
diff --git a/ui/litellm-dashboard/src/components/add_model/build_auto_router_routing_test_request.test.ts b/ui/litellm-dashboard/src/components/add_model/build_auto_router_routing_test_request.test.ts index 4e3ef490079..9ce99d39564 100644 --- a/ui/litellm-dashboard/src/components/add_model/build_auto_router_routing_test_request.test.ts +++ b/ui/litellm-dashboard/src/components/add_model/build_auto_router_routing_test_request.test.ts @@ -52,6 +52,8 @@ describe("buildAutoRouterRoutingTestRequest", () => { ["bespoke", "nimble-latest", "json"], ["bespoke", "bespokelabs/Bespoke-Nimble-9B", "object"], ["bespoke", "bespokelabs/Bespoke-Nimble-9B", "json"], + ["databricks", "databricks-openjev-qwen35-4b", "object"], + ["databricks", "databricks-openjev-qwen35-4b", "json"], ])("probes saved %s/%s %s configuration with custom tiers and team context", (provider, model, format) => { const config = { classifier_type: "oss_classifier", @@ -71,10 +73,15 @@ 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) => { + it.each([ + ["laya", "unsupported"], + ["bespoke", "unsupported"], + ["databricks", "serving-endpoints/openjev"], + ["databricks", ".."], + ])("does not probe unsupported %s model %s", (provider, model) => { const config = { classifier_type: "oss_classifier", - opensource_classifier_config: { provider, model: "unsupported" }, + opensource_classifier_config: { provider, model }, tiers: CONFIG.tiers, }; expect(buildSavedJevConnectionTestRequest(config, "saved-id")).toBeUndefined(); diff --git a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts index e1153e34bc8..23dedbbc860 100644 --- a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts +++ b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts @@ -141,6 +141,9 @@ describe("buildComplexityRouterConfig", () => { { model: " " }, { provider: "laya" as const, model: "unsupported" }, { provider: "bespoke" as const, model: "unsupported" }, + { provider: "databricks" as const, model: "" }, + { provider: "databricks" as const, model: "serving-endpoints/openjev" }, + { provider: "databricks" as const, model: ".." }, { timeout_ms: 0 }, { timeout_ms: 1.5 }, { timeout_ms: Number.NaN }, @@ -159,6 +162,7 @@ describe("buildComplexityRouterConfig", () => { ["bespoke", "nimble-latest"], ["bespoke", "nimble"], ["bespoke", "bespokelabs/Bespoke-Nimble-9B"], + ["databricks", "databricks-openjev-qwen35-4b"], ["jev", "custom-jev-model"], [undefined, "custom-jev-model"], ] as const)("accepts %s model %s before saving or testing", (provider, model) => { @@ -177,6 +181,8 @@ describe("buildComplexityRouterConfig", () => { ["laya", "english", true], ["bespoke", "nimble-latest", false], ["bespoke", "nimble-latest", true], + ["databricks", "databricks-openjev-qwen35-4b", false], + ["databricks", "databricks-openjev-qwen35-4b", true], ] as const)("serializes %s/%s with shared context and no LLM config, custom tiers: %s", (provider, model, custom) => { const params: BuildComplexityRouterConfigParams = { ...baseParams, diff --git a/ui/litellm-dashboard/src/components/add_model/jev_classifier_config.test.ts b/ui/litellm-dashboard/src/components/add_model/jev_classifier_config.test.ts new file mode 100644 index 00000000000..05a10dd4927 --- /dev/null +++ b/ui/litellm-dashboard/src/components/add_model/jev_classifier_config.test.ts @@ -0,0 +1,44 @@ +import { describe, expect, it } from "vitest"; +import { defaultJevClassifierConfig, hydrateOssClassifier, jevClassifierConfigSchema } from "./jev_classifier_config"; + +describe("hydrateOssClassifier", () => { + it("keeps a stored Databricks classifier instead of falling back to Jev", () => { + const stored = { + provider: "databricks", + model: "databricks-openjev-qwen35-4b", + timeout_ms: 4200, + circuit_breaker_enabled: true, + circuit_breaker_cooldown_seconds: 50, + }; + + const hydrated = hydrateOssClassifier({ classifier_type: "oss_classifier", opensource_classifier_config: stored }); + + expect(hydrated).toEqual({ classifier_type: "jev", jev_classifier_config: stored }); + }); + + it("falls back to Jev only when the stored classifier cannot be parsed", () => { + const hydrated = hydrateOssClassifier({ + classifier_type: "oss_classifier", + opensource_classifier_config: { provider: "databricks", model: "serving-endpoints/openjev" }, + }); + + expect(hydrated.jev_classifier_config).toEqual(defaultJevClassifierConfig()); + }); +}); + +describe("jevClassifierConfigSchema", () => { + it("starts a Databricks classifier with an empty endpoint name that does not pass until typed", () => { + const fresh = defaultJevClassifierConfig("databricks"); + + expect(fresh).toEqual({ provider: "databricks", model: "", timeout_ms: 3000 }); + expect(jevClassifierConfigSchema.safeParse(fresh).success).toBe(false); + expect(jevClassifierConfigSchema.safeParse({ ...fresh, model: "my-openjev" }).success).toBe(true); + }); + + it.each(["serving-endpoints/openjev", "openjev?x=1", "openjev#frag", ".", "..", ".openjev", "open jev"])( + "rejects the Databricks endpoint name %j because it would change the request URL", + (model) => { + expect(jevClassifierConfigSchema.safeParse({ provider: "databricks", model }).success).toBe(false); + }, + ); +}); diff --git a/ui/litellm-dashboard/src/components/add_model/jev_classifier_config.ts b/ui/litellm-dashboard/src/components/add_model/jev_classifier_config.ts index a9a7cc18711..02cd8bcbfe9 100644 --- a/ui/litellm-dashboard/src/components/add_model/jev_classifier_config.ts +++ b/ui/litellm-dashboard/src/components/add_model/jev_classifier_config.ts @@ -6,13 +6,33 @@ export const OSS_CLASSIFIER_MODELS = { bespoke: ["nimble-latest", "nimble", "bespokelabs/Bespoke-Nimble-9B"], } as const; +const OSS_CLASSIFIER_PROVIDERS = ["jev", "laya", "bespoke", "databricks"] as const; +export type OssClassifierProvider = (typeof OSS_CLASSIFIER_PROVIDERS)[number]; +export const isOssClassifierProvider = (value: unknown): value is OssClassifierProvider => + OSS_CLASSIFIER_PROVIDERS.some((provider) => provider === value); + +const DEFAULT_TIMEOUT_MS = 3000; +const SERVING_ENDPOINT_NAME = /^[A-Za-z0-9_-][A-Za-z0-9._-]*$/; + +export const fixedClassifierModels = (provider: OssClassifierProvider | undefined): readonly string[] | undefined => + provider === "laya" || provider === "bespoke" ? OSS_CLASSIFIER_MODELS[provider] : undefined; + +const defaultClassifierModel = (provider: OssClassifierProvider | undefined): string => + fixedClassifierModels(provider)?.[0] ?? (provider === "databricks" ? "" : "jev-latest"); + +const isSupportedClassifierModel = (provider: OssClassifierProvider | undefined, model: string): boolean => { + if (provider === "databricks") return SERVING_ENDPOINT_NAME.test(model); + const fixed = fixedClassifierModels(provider); + return fixed === undefined || fixed.some((candidate) => candidate === model); +}; + const jevClassifierConfigFields = { provider: z.preprocess( (value) => (value === "typesafe" ? "jev" : value), - z.enum(["jev", "laya", "bespoke"]).optional(), + z.enum(["jev", "laya", "bespoke", "databricks"]).optional(), ), model: z.string().trim().min(1).optional(), - timeout_ms: z.number().int().positive().default(3000), + timeout_ms: z.number().int().positive().default(DEFAULT_TIMEOUT_MS), instructions: z .string() .nullish() @@ -24,24 +44,18 @@ const jevClassifierConfigFields = { export const jevClassifierConfigSchema = z .object(jevClassifierConfigFields) - .transform((config) => ({ - ...config, - model: - config.model ?? - (config.provider && config.provider !== "jev" ? OSS_CLASSIFIER_MODELS[config.provider][0] : "jev-latest"), - })) - .refine( - (config) => - !config.provider || - config.provider === "jev" || - OSS_CLASSIFIER_MODELS[config.provider].some((model) => model === config.model), - { error: "Select a supported classifier model", path: ["model"] }, - ); + .transform((config) => ({ ...config, model: config.model ?? defaultClassifierModel(config.provider) })) + .refine((config) => isSupportedClassifierModel(config.provider, config.model), { + error: "Select a supported classifier model, or enter the bare Databricks serving endpoint name", + path: ["model"], + }); export type JevClassifierConfig = z.infer; -export const defaultJevClassifierConfig = (provider: JevClassifierConfig["provider"] = "jev"): JevClassifierConfig => - jevClassifierConfigSchema.parse({ provider }); +export const defaultJevClassifierConfig = (provider: OssClassifierProvider = "jev"): JevClassifierConfig => + provider === "databricks" + ? { provider, model: "", timeout_ms: DEFAULT_TIMEOUT_MS } + : jevClassifierConfigSchema.parse({ provider }); export const hydrateOssClassifier = (config: { classifier_type?: ClassifierType | "oss_classifier"; diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index cd13779a2f1..6a34399236c 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -38868,7 +38868,7 @@ export interface components { * @default jev * @enum {string} */ - provider: "jev" | "laya" | "bespoke"; + provider: "jev" | "laya" | "bespoke" | "databricks"; /** * Timeout Ms * @default 3000