feat(decisions): add Databricks ai_decide as a /v1/decisions provider and auto-router decider (#45200)

* feat(decisions): add Databricks ai_decide as a /v1/decisions provider and auto-router decider

* fix(decisions): reject Databricks endpoint names carrying URL separators and keep the provider wiring under llms/

* fix(decisions): reject dot-segment Databricks endpoint names and add Databricks to the dashboard's OSS classifier editor

* fix(ui): name every OSS classifier provider in the classifier radio copy

* feat(decisions): register the Databricks OpenJev decider as an evaluation-mode cost-map entry

* fix(cost-map): declare zero cache rates on the databricks-openjev-qwen35-4b entry

* fix(proxy): reconcile every DB deployment when a registry miss lands on an auto router

A /model/new for an auto router lands on one replica. A request routed to
that auto router on a sibling replica missed it, the model read-through
loaded only the auto router's own row, and the pre-routing hook then picked
a tier the sibling had not loaded, so the caller got a 400 "no healthy
deployments for <tier>" until the next periodic reload. An auto router miss
now runs the full add_deployment reconcile, tiers included. Found by the
audit's two-worker rig on this PR's routing cells; CircleCI's one-worker
proxy never opens the window.

* test(integration): audit cells for Databricks decisions and the Databricks auto-router classifier

Wire cells for /v1/decisions with databricks/<endpoint> deployments, routing
cells for the complexity router's databricks classifier on all three
endpoints with the Postgres rows read back, management cells for the
classifier's spend writes and for the evaluation-mode health check of a
Databricks serving endpoint.
This commit is contained in:
Mateo Wang 2026-10-09 16:19:23 -07:00 • committed by GitHub
parent 22ad3c5fff
commit 883e3210ee
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
33 changed files with 2675 additions and 60 deletions

View file

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

View file

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

View file

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

View file

@ -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://<workspace-host>/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://<workspace-host>/serving-endpoints"
)
return DatabricksDecisionsConnection(api_base=base.rstrip("/"), api_key=key)
DATABRICKS_DECISIONS_CONFIG: Final[DatabricksDecisionsConfig] = DatabricksDecisionsConfig()

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -18,6 +18,7 @@ from litellm.litellm_core_utils.internal_call_metadata import (
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.llms.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:

View file

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

View file

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

View file

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

View file

@ -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?"}},
},
)
]

View file

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

View file

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

File diff suppressed because it is too large Load diff

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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"),
[

View file

@ -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<AutoRouterClassifierTabsProps> = ({ 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<AutoRouterClassifierTabsProps> = ({ val
{[
{ value: "heuristics", label: "Heuristics", description: "Classify locally, with no API call" },
{ value: "llm", label: "LLM", description: "Use a judge model to choose a solver" },
{ value: "jev", label: "OSS Classifier", description: "Use Jev, 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) => (
<Label
key={option.value}
@ -218,6 +226,10 @@ const AutoRouterClassifierTabs: React.FC<AutoRouterClassifierTabsProps> = ({ val
<RadioGroupItem value="bespoke" />
Bespoke Nimble
</Label>
<Label>
<RadioGroupItem value="databricks" />
Databricks
</Label>
</RadioGroup>
</fieldset>
)}

View file

@ -55,7 +55,9 @@ const ClassifierTypeRadios: React.FC<ClassifierTypeRadiosProps> = ({ value, clas
<RadioGroupItem value="jev" className="mt-0.5" />
<span>
<strong className="font-semibold">OSS Classifier</strong>{" "}
<span className="text-muted-foreground">uses Jev or Laya to decide the tier</span>
<span className="text-muted-foreground">
uses Jev, Laya, Bespoke, or a Databricks endpoint to decide the tier
</span>
</span>
</Label>
<SimpleTooltip content={scorerLockedReason}>

View file

@ -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(<Form />);
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();

View file

@ -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<typeof config>) =>
onChange({ ...value, jev_classifier_config: { ...config, ...patch } });
@ -48,7 +50,14 @@ export default function JevClassifierConfig({
</SelectContent>
</Select>
) : (
<Input id={`${id}-model`} value={config.model} onChange={(event) => update({ model: event.target.value })} />
<Input
id={`${id}-model`}
value={config.model}
placeholder={
config.provider === "databricks" ? "Serving endpoint name, e.g. databricks-openjev-qwen35-4b" : undefined
}
onChange={(event) => update({ model: event.target.value })}
/>
)}
</div>
<div>

View file

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

View file

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

View file

@ -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);
},
);
});

View file

@ -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<typeof jevClassifierConfigSchema>;
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";

View file

@ -38868,7 +38868,7 @@ export interface components {
* @default jev
* @enum {string}
*/
provider: "jev" | "laya" | "bespoke";
provider: "jev" | "laya" | "bespoke" | "databricks";
/**
* Timeout Ms
* @default 3000