mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
22ad3c5fff
commit
883e3210ee
33 changed files with 2675 additions and 60 deletions
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
57
litellm/llms/databricks/decisions/transformation.py
Normal file
57
litellm/llms/databricks/decisions/transformation.py
Normal 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()
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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?"}},
|
||||
},
|
||||
)
|
||||
]
|
||||
|
|
|
|||
519
tests/integration/providers/test_databricks_decisions_wire.py
Normal file
519
tests/integration/providers/test_databricks_decisions_wire.py
Normal 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
|
||||
|
|
@ -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
|
|
@ -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"),
|
||||
|
|
|
|||
0
tests/unit/llms/databricks/decisions/__init__.py
Normal file
0
tests/unit/llms/databricks/decisions/__init__.py
Normal file
56
tests/unit/llms/databricks/decisions/test_transformation.py
Normal file
56
tests/unit/llms/databricks/decisions/test_transformation.py
Normal 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"
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
[
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
)}
|
||||
|
|
|
|||
|
|
@ -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}>
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
},
|
||||
);
|
||||
});
|
||||
|
|
@ -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";
|
||||
|
|
|
|||
2
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
2
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -38868,7 +38868,7 @@ export interface components {
|
|||
* @default jev
|
||||
* @enum {string}
|
||||
*/
|
||||
provider: "jev" | "laya" | "bespoke";
|
||||
provider: "jev" | "laya" | "bespoke" | "databricks";
|
||||
/**
|
||||
* Timeout Ms
|
||||
* @default 3000
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue