Merge pull request #39648 from BerriAI/litellm_internal_staging
Some checks are pending
Unit Tests: Proxy DB Operations / auth-checks (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / budgets (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / custom-logging (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / db-and-spend (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / endpoints-and-responses (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / guardrails-hooks (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / jwt-and-keys (push) Blocked by required conditions
CI Coverage / assert-ci-coverage (push) Waiting to run
CodeQL / Analyze (python) (push) Waiting to run
CodeQL / Analyze (actions) (push) Waiting to run
CodeQL / Analyze (javascript-typescript) (push) Waiting to run
CodSpeed Benchmarks / benchmarks (push) Waiting to run
Helm unit test / unit-test (push) Waiting to run
Scorecard supply-chain security / Scorecard analysis (push) Waiting to run
Code Quality Checks / code-quality (push) Waiting to run
Code Quality Checks / python-310-import-smoke (push) Waiting to run
Postgres Tests / proxy-security (push) Waiting to run
Postgres Tests / schema-migration (push) Waiting to run
Postgres Tests / proxy-behavior (push) Waiting to run
LiteLLM Rust / rustfmt, clippy, test (push) Waiting to run
LiteLLM Rust / release wheel (push) Waiting to run
Terraform Provider / gofmt, vet, build, test (push) Waiting to run
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Waiting to run
Unit Tests: Documentation Validation / documentation (push) Waiting to run
Unit Tests: Proxy DB Operations / assert-shard-coverage (push) Waiting to run
Unit Tests / enterprise-package (push) Waiting to run
Unit Tests / enterprise-routing (push) Waiting to run
Unit Tests / integrations (push) Waiting to run
Unit Tests / All Other Providers (push) Waiting to run
Unit Tests / Vertex AI (push) Waiting to run
Unit Tests / misc (push) Waiting to run
Unit Tests / proxy-auth (push) Waiting to run
Unit Tests / proxy-endpoints (push) Waiting to run
Unit Tests / proxy-extras (push) Waiting to run
Unit Tests / proxy-infra (push) Waiting to run
Unit Tests: Proxy DB Operations / key-generation (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / logging-misc (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / proxy-runtime (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / proxy-server-core (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / proxy-utils (push) Blocked by required conditions
Unit Tests / caching-local (push) Waiting to run
Unit Tests / core-utils (push) Waiting to run
Unit Tests / proxy-server (push) Waiting to run
Unit Tests / responses-caching-types (push) Waiting to run
GitHub Actions Security Analysis / zizmor (push) Waiting to run

chore(ci): promote internal staging to main
This commit is contained in:
yuneng-jiang 2026-09-03 14:21:37 -07:00 • committed by GitHub
commit 7d5b6456ba
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
65 changed files with 2461 additions and 805 deletions

View file

@ -40,6 +40,7 @@ ROUTER_SETTINGS_MANAGED_OUTSIDE_CONFIG: Final[frozenset[str]] = frozenset(
"router_general_settings",
"ignore_invalid_deployments",
"fallback_access_check",
"heuristic_v2_router_limit",
}
)
DEFAULT_BATCH_SIZE: Final = int(os.getenv("DEFAULT_BATCH_SIZE", 512))

View file

@ -18,9 +18,11 @@ caller's identity metadata, minus two things that must never be forwarded as-is:
from __future__ import annotations
from collections.abc import Mapping
from types import MappingProxyType
from typing import Final
from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY, NON_INFERENCE_CALL_TYPES
from litellm.litellm_core_utils.initialize_dynamic_callback_params import initialize_standard_callback_dynamic_params
from litellm.types.utils import BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN, InternalCallOrigin
BUDGET_RESERVATION_METADATA_KEYS: Final = frozenset({"user_api_key_budget_reservation"})
@ -142,6 +144,19 @@ def forwarded_internal_call_metadata(
}
def parent_session_kwargs(request_kwargs: Mapping[str, object] | None) -> Mapping[str, str]:
kwargs: Final = request_kwargs or MappingProxyType({})
return MappingProxyType(
{k: v for k in ("litellm_session_id", "litellm_trace_id") if isinstance(v := kwargs.get(k), str)}
)
def effective_turn_off_message_logging(request_kwargs: Mapping[str, object] | None) -> bool | None:
return initialize_standard_callback_dynamic_params(dict(request_kwargs) if request_kwargs else None).get(
"turn_off_message_logging"
)
def sanitized_forwardable_call_metadata(
parent_metadata: Mapping[str, object],
call_origin: InternalCallOrigin,

View file

@ -7,6 +7,7 @@ from litellm.exceptions import UnsupportedParamsError
from litellm.llms.openai.chat.gpt_5_transformation import (
OpenAIGPT5Config,
_get_effort_level,
is_gpt_reasoning_series_name,
)
from litellm.types.llms.openai import AllMessageValues
@ -35,26 +36,7 @@ class AzureOpenAIGPT5Config(AzureOpenAIConfig, OpenAIGPT5Config):
@classmethod
def is_model_gpt_5_model(cls, model: str) -> bool:
"""Check if the Azure model string refers to a gpt-5 variant.
Accepts both explicit gpt-5 model names and the ``gpt5_series/`` prefix
used for manual routing.
"""
# The gpt-5-chat* family (gpt-5-chat, gpt-5-chat-latest, gpt-5-chat-2025-08-07,
# …) are regular chat models: they support temperature and tool_choice but NOT
# reasoning_effort. They must NOT be routed through the GPT-5 reasoning path.
#
# Versioned chat models such as gpt-5.3-chat and gpt-5.1-chat ARE reasoning
# models and must stay on the GPT-5 path. The distinguishing feature is that
# the gpt-5-chat family has a literal "-chat" immediately after "gpt-5"
# (i.e. "gpt-5-chat…"), while versioned chat models interpose a minor version
# number (i.e. "gpt-5.<digit>-chat").
#
# Using a startswith("gpt-5-chat") prefix check on the normalized name (rather
# than a substring check) makes this boundary explicit and avoids any ambiguity
# if future model names coincidentally contain "gpt-5-chat" as an interior run.
_normalized: Final = model.split("/")[-1] # strip provider prefix, e.g. "azure/"
return ("gpt-5" in model and not _normalized.startswith("gpt-5-chat")) or "gpt5_series" in model
return is_gpt_reasoning_series_name(model) or "gpt5_series" in model
def get_supported_openai_params(self, model: str) -> list[str]:
"""Get supported parameters for Azure OpenAI GPT-5 models.

View file

@ -14,6 +14,7 @@ from litellm.litellm_core_utils.prompt_templates.factory import (
convert_to_azure_openai_messages,
)
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.openai.chat.gpt_5_transformation import GPT_REASONING_SERIES_MARKERS
from litellm.types.llms.azure import (
API_VERSION_MONTH_SUPPORTED_RESPONSE_FORMAT,
API_VERSION_YEAR_SUPPORTED_RESPONSE_FORMAT,
@ -139,7 +140,7 @@ class AzureOpenAIConfig(BaseConfig):
name family needs the rename, including the ``gpt-5-chat*`` models that are excluded from
the reasoning path by https://github.com/BerriAI/litellm/issues/13781.
"""
return "gpt-5" in model or "gpt5_series" in model
return any(marker in model for marker in GPT_REASONING_SERIES_MARKERS) or "gpt5_series" in model
def _is_response_format_supported_model(self, model: str) -> bool:
"""

View file

@ -61,6 +61,14 @@ def _get_effort_level(value: str | dict | None) -> str | None:
return None
GPT_REASONING_SERIES_MARKERS: Final = ("gpt-5", "gpt-6")
def is_gpt_reasoning_series_name(model: str) -> bool:
normalized: Final = model.split("/")[-1]
return any(marker in model for marker in GPT_REASONING_SERIES_MARKERS) and not normalized.startswith("gpt-5-chat")
class OpenAIGPT5Config(OpenAIGPTConfig):
"""Configuration for gpt-5 models including GPT-5-Codex variants.
@ -73,21 +81,7 @@ class OpenAIGPT5Config(OpenAIGPTConfig):
@classmethod
def is_model_gpt_5_model(cls, model: str) -> bool:
# The gpt-5-chat* family (gpt-5-chat, gpt-5-chat-latest, gpt-5-chat-2025-08-07,
# …) are regular chat models: they support temperature and tool_choice but NOT
# reasoning_effort. They must NOT be routed through the GPT-5 reasoning path.
#
# Versioned chat models such as gpt-5.3-chat and gpt-5.1-chat ARE reasoning
# models and must stay on the GPT-5 path. The distinguishing feature is that
# the gpt-5-chat family has a literal "-chat" immediately after "gpt-5"
# (i.e. "gpt-5-chat…"), while versioned chat models interpose a minor version
# number (i.e. "gpt-5.<digit>-chat").
#
# Using a startswith("gpt-5-chat") prefix check on the normalized name (rather
# than a substring check) makes this boundary explicit and avoids any ambiguity
# if future model names coincidentally contain "gpt-5-chat" as an interior run.
_normalized: Final = model.split("/")[-1] # strip provider prefix, e.g. "openai/"
return "gpt-5" in model and not _normalized.startswith("gpt-5-chat")
return is_gpt_reasoning_series_name(model)
@classmethod
def is_model_gpt_5_search_model(cls, model: str) -> bool:
@ -122,6 +116,8 @@ class OpenAIGPT5Config(OpenAIGPTConfig):
def is_model_gpt_5_4_plus_model(cls, model: str) -> bool:
"""Check if the model is gpt-5.4 or newer (5.4, 5.5, 5.6, etc., including pro)."""
model_name: Final = model.split("/")[-1]
if model_name.startswith("gpt-6"):
return True
if not model_name.startswith("gpt-5."):
return False
try:

View file

@ -15,6 +15,7 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo
)
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
from litellm.llms.openai.chat.gpt_5_transformation import is_gpt_reasoning_series_name
from litellm.responses.litellm_completion_transformation.custom_tools import TOOL_CALL_ITEM_ID_PREFIX_BY_TYPE
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import *
@ -88,7 +89,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
parts: Final = model.split("/")
if len(parts) > 1 and parts[0] not in ("openai",):
return False
return "gpt-5" in model and "gpt-5-chat" not in model
return is_gpt_reasoning_series_name(model)
@staticmethod
def _supports_reasoning_effort_none(model: str) -> bool:

View file

@ -998,6 +998,16 @@ def replace_project_and_location_in_route(requested_route: str, vertex_project:
return modified_route
def _api_version_for_route(requested_route: str) -> Literal["v1", "v1beta1"]:
return "v1beta1" if "cachedContent" in requested_route else "v1"
def _with_api_version(requested_route: str) -> str:
if not requested_route.startswith("/projects/"):
return requested_route
return f"/{_api_version_for_route(requested_route)}{requested_route}"
def construct_target_url(
base_url: str,
requested_route: str,
@ -1017,18 +1027,19 @@ def construct_target_url(
new_base_url: Final = httpx.URL(base_url)
if "locations" in requested_route: # contains the target project id + location
if vertex_project and vertex_location:
requested_route = replace_project_and_location_in_route(requested_route, vertex_project, vertex_location)
return new_base_url.copy_with(path=requested_route)
targeted_route: Final = (
replace_project_and_location_in_route(requested_route, vertex_project, vertex_location)
if vertex_project and vertex_location
else requested_route
)
return new_base_url.copy_with(path=_with_api_version(targeted_route))
"""
- Add endpoint version (e.g. v1beta for cachedContent, v1 for rest)
- Add default project id
- Add default location
"""
vertex_version: Literal["v1", "v1beta1"] = "v1"
if "cachedContent" in requested_route:
vertex_version = "v1beta1"
vertex_version: Literal["v1", "v1beta1"] = _api_version_for_route(requested_route)
# Check if the requested route starts with a version
# e.g. /v1beta1/publishers/google/models/gemini-3-pro-preview:streamGenerateContent

View file

@ -16,6 +16,10 @@ if TYPE_CHECKING:
from litellm.proxy._types import EnterpriseLicenseData
AUTO_ROUTER_LICENSE_FEATURE: Final = "auto_router"
HEURISTIC_V2_LICENSE_REMEDY: Final = "A LiteLLM license with the 'auto_router' feature lifts the limit."
class LicenseCheck:
"""
- Check if license in env
@ -149,6 +153,19 @@ class LicenseCheck:
return False
return team_count > _max_teams_in_license
def heuristic_v2_router_limit(self) -> int | None:
"""
How many heuristic_v2 auto-routers this proxy may hold: unlimited (None) only when the
signed license lists the auto_router feature, otherwise one. A license verified through
the API carries no feature list, so it does not lift the limit either.
"""
if self.airgapped_license_data is None:
return 1
allowed_features: Final = self.airgapped_license_data.get("allowed_features")
if isinstance(allowed_features, list) and AUTO_ROUTER_LICENSE_FEATURE in allowed_features:
return None
return 1
def verify_license_without_api_request(self, public_key, license_key):
try:
from cryptography.hazmat.primitives import hashes
@ -179,19 +196,21 @@ class LicenseCheck:
# Decode and parse the data
license_data: Final = json.loads(message.decode())
self.airgapped_license_data = EnterpriseLicenseData(**license_data)
# debug information provided in license data
verbose_proxy_logger.debug("License data: %s", license_data)
# Check expiration date
expiration_date: Final = datetime.strptime(license_data["expiration_date"], "%Y-%m-%d")
if expiration_date < datetime.now():
self.airgapped_license_data = None
return False, "License has expired"
self.airgapped_license_data = EnterpriseLicenseData(**license_data)
return True
except Exception as e:
self.airgapped_license_data = None
verbose_proxy_logger.debug(
"litellm.proxy.auth.litellm_license.py::verify_license_without_api_request - Unable to verify License locally. - %s",
e,

View file

@ -17,6 +17,7 @@ import json
import traceback
from collections.abc import Awaitable, Mapping, Sequence
from datetime import datetime, timezone
from types import MappingProxyType
from typing import Any, Final, Literal, Protocol, cast, overload
import fastapi
@ -77,6 +78,7 @@ from litellm.types.proxy.management_endpoints.internal_user_endpoints import (
BulkUpdateUserRequest,
BulkUpdateUserResponse,
UserListResponse,
UserSearchWhere,
UserUpdateResult,
)
from litellm.types.proxy.management_endpoints.scim_v2 import (
@ -2080,6 +2082,22 @@ async def _authorize_user_list_request(
return ",".join(allowed_org_ids)
_NO_SEARCH_WHERE: Final[Mapping[str, object]] = MappingProxyType({})
def _user_search_where(search: str | None) -> Mapping[str, object]:
"""Prisma predicate for `/user/list?search=`: user_id or user_email contains it, case-insensitive."""
if not search:
return _NO_SEARCH_WHERE
search_where: Final[UserSearchWhere] = {
"OR": (
{"user_id": {"contains": search, "mode": "insensitive"}},
{"user_email": {"contains": search, "mode": "insensitive"}},
)
}
return search_where
@router.get(
"/user/list",
tags=["Internal User management"],
@ -2091,6 +2109,10 @@ async def get_users(
user_ids: str | None = fastapi.Query(default=None, description="Get list of users by user_ids"),
sso_user_ids: str | None = fastapi.Query(default=None, description="Get list of users by sso_user_id"),
user_email: str | None = fastapi.Query(default=None, description="Filter users by partial email match"),
search: str | None = fastapi.Query(
default=None,
description="Combined search: matches users whose 'user_id' or 'user_email' contains the value (case-insensitive).",
),
team: str | None = fastapi.Query(default=None, description="Filter users by team id"),
page: int = fastapi.Query(default=1, ge=1, description="Page number"),
page_size: int = fastapi.Query(default=25, ge=1, le=100, description="Number of items per page"),
@ -2121,6 +2143,8 @@ async def get_users(
Get list of users by sso_ids. Comma separated list of sso_ids.
user_email: Optional[str]
Filter users by partial email match
search: Optional[str]
Combined search: matches users whose user_id or user_email contains the value (case-insensitive)
team: Optional[str]
Filter users by team id. Will match if user has this team in their teams array.
page: int
@ -2197,7 +2221,11 @@ async def get_users(
where_conditions["organization_memberships"] = {"some": {"organization_id": {"in": org_id_list}}}
## Filter any none fastapi.Query params - e.g. where_conditions: {'user_email': {'contains': Query(None), 'mode': 'insensitive'}, 'teams': {'has': Query(None)}}
where_conditions = {k: v for k, v in where_conditions.items() if v is not None}
where: Final[Mapping[str, object]] = {
key: value
for key, value in (*where_conditions.items(), *_user_search_where(search).items())
if value is not None
}
# Build order_by conditions
@ -2206,14 +2234,14 @@ async def get_users(
)
users: Final[Sequence[prisma_models.LiteLLM_UserTable]] = await UserRepository(prisma_client).table.find_many(
where=where_conditions,
where=where,
skip=skip,
take=page_size,
order=(order_by if order_by else {"created_at": "desc"}), # Default to created_at desc if no sort specified
)
# Get total count of user rows
total_count: Final[int] = await UserRepository(prisma_client).table.count(where=where_conditions)
total_count: Final[int] = await UserRepository(prisma_client).table.count(where=where)
# Get key count for each user
user_key_counts: Final = await get_user_key_counts(prisma_client, [user.user_id for user in users])

View file

@ -13,10 +13,11 @@ model/{model_id}/update - PATCH endpoint for model update.
import asyncio
import datetime
import json
from collections.abc import Awaitable, Mapping, Sequence
from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, Sequence
from contextlib import AbstractAsyncContextManager, asynccontextmanager
from json import JSONDecodeError
from types import MappingProxyType
from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, cast
from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, TypeVar, cast
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
from pydantic import BaseModel, ConfigDict, Field, ValidationError, field_validator
@ -49,6 +50,7 @@ from litellm.proxy._types import (
TeamModelDeleteRequest,
UserAPIKeyAuth,
)
from litellm.proxy.auth.litellm_license import HEURISTIC_V2_LICENSE_REMEDY
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.config_sync_pubsub import (
coordination_redis_cache,
@ -96,6 +98,9 @@ from litellm.router_strategy.complexity_router import (
from litellm.router_utils.auto_router_model_naming import (
STRATEGY_ROUTER_PARAM_FIELDS,
carries_complexity_router_settings,
count_heuristic_v2_routers,
heuristic_v2_limit_violation,
uses_heuristic_v2_classifier,
validate_complexity_router_config_placement,
validate_complexity_router_config_write,
validate_strategy_router_model_write,
@ -153,6 +158,8 @@ class _ProxyModelTable(Protocol):
def find_many(self, *, where: Mapping[str, object]) -> Awaitable[Sequence[_ProxyModelRow]]: ...
def create(self, *, data: Mapping[str, object]) -> Awaitable[_ProxyModelRow]: ...
def update(
self, *, where: Mapping[str, object], data: Mapping[str, object]
) -> Awaitable[_ProxyModelRow | None]: ...
@ -166,6 +173,9 @@ class _TxModelTables(Protocol):
litellm_proxymodeltable: _ProxyModelTable
_RowT = TypeVar("_RowT")
class _ExistingModelRow(Protocol):
@property
def litellm_params(self) -> Mapping[str, object]: ...
@ -269,6 +279,66 @@ def _raise_on_strategy_router_write_violation(
)
HEURISTIC_V2_SLOT_LOCK_KEY: Final = 5_872_301
_HEURISTIC_V2_LOCK_SQL: Final = "SELECT 1 AS locked FROM pg_advisory_xact_lock($1)"
_HEURISTIC_V2_DB_ROWS_SQL: Final = """
SELECT count(*)::int AS held FROM "LiteLLM_ProxyModelTable"
WHERE model_id <> $1
AND (CASE jsonb_typeof(litellm_params) WHEN 'string' THEN (litellm_params #>> '{}')::jsonb ELSE litellm_params END)
-> 'complexity_router_config' ->> 'classifier_type' = 'heuristic_v2'
"""
def _effective_complexity_router_config(
incoming_params: GenericLiteLLMParams | None, existing_params: GenericLiteLLMParams | None
) -> object:
"""The complexity config a write leaves on the row: the incoming one when the write carries it, else the stored one."""
incoming: Final = None if incoming_params is None else incoming_params.complexity_router_config
if incoming is not None or existing_params is None:
return incoming
return existing_params.complexity_router_config
@asynccontextmanager
async def _heuristic_v2_slot(
prisma_client: PrismaClient, *, effective_config: object, model_id: str | None
) -> AsyncGenerator[_ProxyModelTable, None]:
"""Hand out the model table to write through while the row's claim on a heuristic_v2 slot is settled.
A write that leaves the row on classifier_type heuristic_v2 under a limited license runs
inside one transaction that takes an advisory lock in its own statement before counting
(a statement's snapshot predates anything it locks), so pods cannot both pass the count:
the DB rows (any pod, either JSON shape) plus this proxy's config.yaml routers are judged
against the license limit and the write is refused with a 403 before it happens. The row
being edited keeps its own slot through ``model_id``. Every other write, and every write on
an unlimited license, goes through the repository table with no lock. Only the row write
itself may run inside: anything that needs a second connection (the team model bookkeeping)
must wait until the transaction has committed and the lock is released. The transaction
writes bypass the repository's publish-on-write, so the config change is published once
after commit, the way delete_team_models does.
"""
from litellm.proxy.proxy_server import _license_check, llm_router
limit: Final = _license_check.heuristic_v2_router_limit()
if limit is None or not uses_heuristic_v2_classifier(effective_config):
yield _proxy_model_table(prisma_client)
return
async with prisma_client.db.tx() as tx_ctx:
tables: Final[_TxModelTables] = tx_ctx
await tx_ctx.query_raw(_HEURISTIC_V2_LOCK_SQL, HEURISTIC_V2_SLOT_LOCK_KEY)
rows: Sequence[Mapping[str, object]] = await tx_ctx.query_raw(_HEURISTIC_V2_DB_ROWS_SQL, model_id or "")
db_held: Final = rows[0].get("held") if rows else 0
config_rows: Final = () if llm_router is None else tuple(llm_router.config_deployments())
held: Final = (db_held if isinstance(db_held, int) else 0) + count_heuristic_v2_routers(config_rows)
violation: Final = heuristic_v2_limit_violation(held=held + 1, limit=limit)
if violation is not None:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN, detail=f"{violation} {HEURISTIC_V2_LICENSE_REMEDY}"
)
yield tables.litellm_proxymodeltable
await publish_config_change(redis_cache=coordination_redis_cache(), object_type="litellm_proxymodeltable")
ENFORCE_RPM_TPM_ON_MODEL_ADD_SETTING: Final = "enforce_rpm_tpm_on_model_add"
_REQUIRED_RATE_LIMIT_FIELDS: Final = ("rpm", "tpm")
@ -720,22 +790,29 @@ async def patch_model(
)
requested_model_name: Final = patch_data.model_name
stored_model_name: str | None = None
async def write_row(update_data: PrismaCompatibleUpdateDBModel) -> _ProxyModelRow | None:
nonlocal stored_model_name
stored_model_name = update_data.get("model_name")
update_data["updated_by"] = user_api_key_dict.user_id or litellm_proxy_admin_name
update_data["updated_at"] = cast(str, get_utc_datetime())
async with _heuristic_v2_slot(
prisma_client,
effective_config=_effective_complexity_router_config(
patch_data.litellm_params, db_model.litellm_params
),
model_id=model_id,
) as table:
return await table.update(where={"model_id": model_id}, data=update_data)
# Handle team model updates with proper alias management
update_data: Final = await _update_team_model_in_db(
updated_model: Final = await _update_team_model_in_db(
db_model=db_model,
patch_data=patch_data,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
)
# Add metadata about update
update_data["updated_by"] = user_api_key_dict.user_id or litellm_proxy_admin_name
update_data["updated_at"] = cast(str, get_utc_datetime())
# Perform partial update
updated_model: Final = await _proxy_model_table(prisma_client).update(
where={"model_id": model_id},
data=update_data,
write_row=write_row,
)
if updated_model is None:
@ -746,7 +823,6 @@ async def patch_model(
param=None,
)
stored_model_name: Final = update_data.get("model_name")
if (
stored_model_name is not None
and stored_model_name == requested_model_name
@ -980,7 +1056,8 @@ async def _add_model_to_db(
prisma_client: PrismaClient,
new_encryption_key: str | None = None,
should_create_model_in_db: bool = True,
) -> "prisma_models.LiteLLM_ProxyModelTable | LiteLLM_ProxyModelTable | None":
slot: AbstractAsyncContextManager[_ProxyModelTable] | None = None,
) -> "_ProxyModelRow | LiteLLM_ProxyModelTable":
# encrypt litellm params #
_litellm_params_dict: Final = model_params.litellm_params.dict(exclude_none=True)
_original_litellm_model_name: Final = model_params.litellm_params.model
@ -998,18 +1075,20 @@ async def _add_model_to_db(
if model_params.model_info.id is not None:
_data["model_id"] = model_params.model_info.id
_create_data: Final = cast("Mapping[str, object]", _data) # cast-ok: str-keyed json payload built just above
if should_create_model_in_db:
model_response = await ModelRepository(prisma_client).table.create(data=_create_data)
else:
model_response = LiteLLM_ProxyModelTable(**_data)
return model_response
if not should_create_model_in_db:
return LiteLLM_ProxyModelTable(**_data)
if slot is None:
return await _proxy_model_table(prisma_client).create(data=_create_data)
async with slot as table:
return await table.create(data=_create_data)
async def _add_team_model_to_db(
model_params: Deployment,
user_api_key_dict: UserAPIKeyAuth,
prisma_client: PrismaClient,
) -> "prisma_models.LiteLLM_ProxyModelTable | LiteLLM_ProxyModelTable | None":
slot: AbstractAsyncContextManager[_ProxyModelTable] | None = None,
) -> "_ProxyModelRow | LiteLLM_ProxyModelTable":
"""
If 'team_id' is provided,
@ -1040,6 +1119,7 @@ async def _add_team_model_to_db(
model_params=model_params,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
slot=slot,
)
if original_model_name:
@ -1060,7 +1140,8 @@ async def _update_team_model_in_db(
patch_data: updateDeployment,
user_api_key_dict: UserAPIKeyAuth,
prisma_client: PrismaClient,
) -> PrismaCompatibleUpdateDBModel:
write_row: Callable[[PrismaCompatibleUpdateDBModel], Awaitable[_RowT]],
) -> _RowT:
"""
Handle team model updates with proper alias management.
@ -1068,6 +1149,9 @@ async def _update_team_model_in_db(
- Creates unique internal model_name and team alias
- Adds model to team object
- Preserves team_public_model_name for external reference
The row is written through ``write_row`` before the team's model list is touched, so a
refused or failed write leaves the team as it was (the create path orders itself the same way).
"""
# Validate team_id if present in patch_data
from litellm.proxy.proxy_server import premium_user
@ -1079,9 +1163,7 @@ async def _update_team_model_in_db(
premium_user=premium_user,
)
# Validated before any write, beside the premium check the create path already runs
# here. The team ACL is updated below and autocommits, so a validator that raises
# further down would leave the team mutated and the deployment row never written.
# Validated before the row write, beside the premium check the create path already runs here.
#
# The merged view is what gets stored, so that is what has to satisfy the invariants.
# Validating the patch alone rejected a partial edit of an already valid deployment:
@ -1101,7 +1183,7 @@ async def _update_team_model_in_db(
# No team_id in patch, proceed with standard update
if patch_team_id is None:
return update_db_model(db_model=db_model, updated_patch=patch_data)
return await write_row(update_db_model(db_model=db_model, updated_patch=patch_data))
# Determine public model name
public_model_name: Final = _get_public_model_name(
@ -1120,11 +1202,14 @@ async def _update_team_model_in_db(
db_team_id: Final = db_model.model_info.team_id if db_model.model_info else None
is_new_team_assignment: Final = db_team_id != patch_team_id
# Team rows keep their internal UUID-based model_name; the public name lives in model_info
patch_data.model_name = f"model_name_{patch_team_id}_{uuid.uuid4()}" if is_new_team_assignment else None
row: Final = await write_row(update_db_model(db_model=db_model, updated_patch=patch_data))
if is_new_team_assignment:
await _setup_new_team_model_assignment(
team_id=patch_team_id,
public_model_name=public_model_name,
patch_data=patch_data,
user_api_key_dict=user_api_key_dict,
)
else:
@ -1132,12 +1217,11 @@ async def _update_team_model_in_db(
team_id=patch_team_id,
public_model_name=public_model_name,
db_model=db_model,
patch_data=patch_data,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
)
return update_db_model(db_model=db_model, updated_patch=patch_data)
return row
def _get_public_model_name(
@ -1189,13 +1273,9 @@ def _get_public_model_name(
async def _setup_new_team_model_assignment(
team_id: str,
public_model_name: str,
patch_data: updateDeployment,
user_api_key_dict: UserAPIKeyAuth,
) -> None:
"""Set up a new team model with unique name and team membership."""
unique_model_name: Final = f"model_name_{team_id}_{uuid.uuid4()}"
patch_data.model_name = unique_model_name
"""Register a newly team-assigned model's public name on the team."""
await team_model_add(
data=TeamModelAddRequest(
team_id=team_id,
@ -1385,7 +1465,6 @@ async def _update_existing_team_model_assignment(
team_id: str,
public_model_name: str,
db_model: Deployment,
patch_data: updateDeployment,
user_api_key_dict: UserAPIKeyAuth,
prisma_client: PrismaClient | None,
) -> None:
@ -1409,9 +1488,6 @@ async def _update_existing_team_model_assignment(
old_public_name: Final = db_model.model_info.team_public_model_name if db_model.model_info else None
if old_public_name and public_model_name != old_public_name:
# Clear user-supplied public name from patch before any early return so the
# caller does not overwrite the internal UUID-based model_name in the DB.
patch_data.model_name = None
if prisma_client is None:
verbose_proxy_logger.warning(
"prisma_client not initialized; skipping public name update entirely to avoid orphaned entries"
@ -1459,10 +1535,6 @@ async def _update_existing_team_model_assignment(
# else: old_public_name == public_model_name (no rename needed)
# No team_model_add/delete calls required; public name is already registered
# Always clear patch_data.model_name to prevent caller from overwriting
# the internal UUID-based model_name in the DB with the user-supplied public name
patch_data.model_name = None
class ModelManagementAuthChecks:
"""
@ -1878,18 +1950,19 @@ async def add_new_model(
reload_outcome: ReconcileOutcome = ReconcileOutcome(still_desired=None, live_after=None)
try:
_original_litellm_model_name: Final = model_params.model_name
if model_params.model_info.team_id is None:
model_response = await _add_model_to_db(
model_params=priced_model_params,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
)
else:
model_response = await _add_team_model_to_db(
model_params=priced_model_params,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
)
add_model: Final = (
_add_model_to_db if model_params.model_info.team_id is None else _add_team_model_to_db
)
model_response = await add_model(
model_params=priced_model_params,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
slot=_heuristic_v2_slot(
prisma_client,
effective_config=priced_model_params.litellm_params.complexity_router_config,
model_id=priced_model_params.model_info.id,
),
)
reload_outcome = await proxy_config.add_deployment(
prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj
)
@ -1903,6 +1976,8 @@ async def add_new_model(
passed_model_info=priced_model_params.model_info,
)
except Exception as e:
if isinstance(e, HTTPException):
raise
verbose_proxy_logger.exception("Exception in add_new_model: %s", e)
else:
@ -2070,10 +2145,17 @@ async def update_model(
"updated_by": user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME,
**({} if renamed_to is None else {"model_name": renamed_to}),
}
model_response: Final = await _proxy_model_table(prisma_client).update(
where={"model_id": _model_id},
data=_data,
)
async with _heuristic_v2_slot(
prisma_client,
effective_config=_effective_complexity_router_config(
model_params.litellm_params, deployment.litellm_params
),
model_id=_model_id,
) as table:
model_response: Final = await table.update(
where={"model_id": _model_id},
data=_data,
)
if renamed_to is not None:
await sync_access_groups_for_renamed_model(
prisma_client=prisma_client,

View file

@ -29,6 +29,12 @@ from litellm.proxy._types import (
from litellm.proxy.auth.auth_utils import is_request_body_safe
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.path_utils import safe_filename
from litellm.proxy.prompts.prompt_registry import (
DEFAULT_PROMPT_ENVIRONMENT,
get_base_prompt_id,
get_version_number,
prompt_environment_or_default,
)
from litellm.repositories.table_repositories import PromptRepository
from litellm.types.prompts.init_prompts import (
ListPromptsResponse,
@ -102,165 +108,20 @@ def _prompt_table(prisma_client: "PrismaClient") -> _PromptTableActions:
return PromptRepository(prisma_client).table
def get_base_prompt_id(prompt_id: str) -> str:
"""
Extract the base prompt ID by stripping the version suffix if present.
Args:
prompt_id: Prompt ID that may include version suffix (e.g., "jack_success.v1" or "jack_success_v1")
Returns:
Base prompt ID without version suffix (e.g., "jack_success")
Examples:
>>> get_base_prompt_id("jack_success.v1")
"jack_success"
>>> get_base_prompt_id("jack_success_v1")
"jack_success"
>>> get_base_prompt_id("jack_success")
"jack_success"
"""
# Try dot separator first (.v)
if ".v" in prompt_id:
return prompt_id.split(".v")[0]
# Try underscore separator (_v)
if "_v" in prompt_id:
return prompt_id.split("_v")[0]
return prompt_id
def get_version_number(prompt_id: str) -> int:
"""
Extract the version number from a versioned prompt ID.
Args:
prompt_id: Prompt ID that may include version suffix (e.g., "jack_success.v2" or "jack_success_v2")
Returns:
Version number (defaults to 1 if no version suffix or invalid format)
Examples:
>>> get_version_number("jack_success.v2")
2
>>> get_version_number("jack_success_v2")
2
>>> get_version_number("jack_success")
1
"""
# Try dot separator first (.v)
if ".v" in prompt_id:
version_str = prompt_id.split(".v")[1]
try:
return int(version_str)
except ValueError:
pass
# Try underscore separator (_v)
if "_v" in prompt_id:
version_str = prompt_id.split("_v")[1]
try:
return int(version_str)
except ValueError:
pass
return 1
def construct_versioned_prompt_id(prompt_id: str, version: int | None = None) -> str:
"""
Construct a versioned prompt ID from a base prompt_id and version number.
Args:
prompt_id: Base prompt ID (e.g., "jack_success")
version: Version number (if None, returns the base prompt_id unchanged)
Returns:
Versioned prompt ID (e.g., "jack_success.v4")
Examples:
>>> construct_versioned_prompt_id("jack_success", 4)
"jack_success.v4"
>>> construct_versioned_prompt_id("jack_success", None)
"jack_success"
>>> construct_versioned_prompt_id("jack_success.v2", 4)
"jack_success.v4"
"""
if version is None:
return prompt_id
# Strip any existing version suffix first
base_id: Final = get_base_prompt_id(prompt_id)
return f"{base_id}.v{version}"
def get_latest_version_prompt_id(prompt_id: str, all_prompt_ids: Mapping[str, object]) -> str:
"""
Find the latest version of a prompt from available prompt IDs.
Args:
prompt_id: Base prompt ID or versioned prompt ID (e.g., "jack_success" or "jack_success.v2")
all_prompt_ids: Dictionary of all available prompt IDs (keys are prompt IDs)
Returns:
The prompt ID with the highest version number, or the original prompt_id if no versions exist
Examples:
>>> all_ids = {"jack.v1": {}, "jack.v2": {}, "jack.v3": {}}
>>> get_latest_version_prompt_id("jack", all_ids)
"jack.v3"
>>> get_latest_version_prompt_id("jack.v1", all_ids)
"jack.v3"
>>> all_ids = {"simple": {}}
>>> get_latest_version_prompt_id("simple", all_ids)
"simple"
"""
base_id: Final = get_base_prompt_id(prompt_id=prompt_id)
# Find all versions of this prompt
matching_versions: Final = []
for stored_prompt_id in all_prompt_ids:
if get_base_prompt_id(prompt_id=stored_prompt_id) == base_id:
version_num = get_version_number(prompt_id=stored_prompt_id)
matching_versions.append((version_num, stored_prompt_id))
# Use the highest version number
if matching_versions:
matching_versions.sort(reverse=True)
return matching_versions[0][1]
else:
# No versioned prompts found, use the base ID as-is
return prompt_id
def get_latest_prompt_versions(prompts: list[PromptSpec]) -> list[PromptSpec]:
"""
Filter a list of prompts to return only the latest version of each unique prompt.
Args:
prompts: List of PromptSpec objects
Returns:
List of PromptSpec objects with only the latest version of each prompt
Filter prompts down to the latest version per (base prompt id, environment).
"""
latest_prompts: Final[dict[str, PromptSpec]] = {}
for prompt in prompts:
base_id = get_base_prompt_id(prompt_id=prompt.prompt_id)
version = get_version_number(prompt_id=prompt.prompt_id)
# Keep the prompt with the highest version number
if base_id not in latest_prompts:
latest_prompts[base_id] = prompt
else:
existing_version = get_version_number(prompt_id=latest_prompts[base_id].prompt_id)
if version > existing_version:
latest_prompts[base_id] = prompt
sorted_prompts: Final = sorted(prompts, key=lambda prompt: get_version_number(prompt_id=prompt.prompt_id))
latest_prompts: Final = {
(get_base_prompt_id(prompt_id=prompt.prompt_id), prompt_environment_or_default(prompt.environment)): prompt
for prompt in sorted_prompts
}
return list(latest_prompts.values())
async def get_next_version_for_prompt(
prisma_client: "PrismaClient", prompt_id: str, environment: str = "development"
prisma_client: "PrismaClient", prompt_id: str, environment: str = DEFAULT_PROMPT_ENVIRONMENT
) -> int:
"""
Get the next version number for a prompt in a specific environment.
@ -403,11 +264,14 @@ async def list_prompts(
if key_metadata is not None:
prompts: Final = cast(list[str] | None, key_metadata.get("prompts", None))
if prompts is not None:
all_prompts = [
IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS[prompt_id]
for prompt_id in prompts
if prompt_id in IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS
allowed_prompt_ids: Final = frozenset(prompts)
allowed_prompts: Final = [
spec
for spec in IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS.values()
if spec.prompt_id in allowed_prompt_ids
or get_base_prompt_id(prompt_id=spec.prompt_id) in allowed_prompt_ids
]
all_prompts = get_latest_prompt_versions(prompts=allowed_prompts)
if environment:
all_prompts = [p for p in all_prompts if p.environment == environment]
prompt_list: Final = []
@ -576,7 +440,7 @@ def _get_prompt_template(prompt_spec: PromptSpec, base_prompt_id: str) -> Prompt
metadata=parsed.get("metadata"),
)
else:
prompt_callback: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id(prompt_spec.prompt_id)
prompt_callback: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_for_prompt(prompt=prompt_spec)
if prompt_callback is not None:
integration_name: Final = prompt_callback.integration_name
if integration_name == "dotprompt":
@ -690,15 +554,10 @@ async def get_prompt_info(
if env_prompts:
prompt_spec = create_versioned_prompt_spec(db_prompt=env_prompts[0])
# Fallback: use in-memory registry (no environment filter)
if prompt_spec is None and environment is None:
prompt_spec = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id)
if prompt_spec is None:
latest_prompt_id: Final = get_latest_version_prompt_id(
prompt_id=prompt_id,
all_prompt_ids=IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS,
)
prompt_spec = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(latest_prompt_id)
if prompt_spec is None:
prompt_spec = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec(
prompt_id, version=requested_version, environment=environment
)
if prompt_spec is None:
raise HTTPException(
@ -785,7 +644,7 @@ async def create_prompt(
environment: Final = (
request.prompt_info.environment
if request.prompt_info and request.prompt_info.environment
else "development"
else DEFAULT_PROMPT_ENVIRONMENT
)
# Get next version number
@ -885,7 +744,7 @@ async def update_prompt(
environment: Final = (
request.prompt_info.environment
if request.prompt_info and request.prompt_info.environment
else "development"
else DEFAULT_PROMPT_ENVIRONMENT
)
# Check if any version of this prompt exists (in any environment)
@ -897,9 +756,7 @@ async def update_prompt(
detail=f"Prompt with ID {base_prompt_id} not found",
)
# Check if it's a config prompt
existing_in_memory: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id)
if existing_in_memory and existing_in_memory.prompt_info.prompt_type == "config":
if IN_MEMORY_PROMPT_REGISTRY.has_config_prompt(base_prompt_id=base_prompt_id):
raise HTTPException(
status_code=400,
detail="Cannot update config prompts.",
@ -988,40 +845,26 @@ async def delete_prompt(
raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
try:
# Try to get prompt directly first
existing_prompt = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id)
# If not found, try to find the latest version
if existing_prompt is None:
latest_prompt_id: Final = get_latest_version_prompt_id(
prompt_id=prompt_id,
all_prompt_ids=IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS,
)
existing_prompt = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(latest_prompt_id)
# Use the resolved prompt_id for deletion
prompt_id = latest_prompt_id
base_prompt_id: Final = get_base_prompt_id(prompt_id=prompt_id)
existing_prompt: Final = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec(prompt_id, environment=environment)
if existing_prompt is None:
raise HTTPException(status_code=404, detail=f"Prompt with ID {prompt_id} not found")
if existing_prompt.prompt_info.prompt_type == "config":
if IN_MEMORY_PROMPT_REGISTRY.has_config_prompt(base_prompt_id=base_prompt_id):
raise HTTPException(
status_code=400,
detail="Cannot delete config prompts.",
)
# Get the base prompt ID (without version suffix) for database deletion
base_prompt_id: Final = get_base_prompt_id(prompt_id=prompt_id)
# Build delete filter; scope to environment if provided
delete_where: Final[dict[str, str]] = {"prompt_id": base_prompt_id}
if environment:
delete_where["environment"] = environment
# Delete versions from the database (scoped to environment if provided)
delete_where: Final[dict[str, str]] = {
"prompt_id": base_prompt_id,
**({"environment": environment} if environment else {}),
}
await _prompt_table(prisma_client).delete_many(where=delete_where)
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id(base_prompt_id, environment=environment or None)
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id(
base_prompt_id=base_prompt_id, environment=environment or None
)
env_msg: Final = f" from {environment}" if environment else ""
return {"message": f"Prompt {base_prompt_id} deleted successfully{env_msg}"}
@ -1093,7 +936,7 @@ async def patch_prompt(
try:
# Resolve the target row: find the latest version in the given environment
base_prompt_id: Final = get_base_prompt_id(prompt_id=prompt_id)
env: Final = environment or "development"
env: Final = prompt_environment_or_default(environment)
requested_version: Final = get_version_number(prompt_id=prompt_id) if prompt_id != base_prompt_id else None
# Build query to find the exact row by composite unique key
@ -1117,11 +960,7 @@ async def patch_prompt(
target_row: Final = db_rows[0]
# Check if prompt exists in memory
versioned_id: Final = f"{base_prompt_id}.v{target_row.version}"
existing_prompt: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(versioned_id)
if existing_prompt and existing_prompt.prompt_info.prompt_type == "config":
if IN_MEMORY_PROMPT_REGISTRY.has_config_prompt(base_prompt_id=base_prompt_id):
raise HTTPException(
status_code=400,
detail="Cannot update config prompts.",

View file

@ -1,6 +1,6 @@
import importlib
import os
from collections.abc import Callable
from collections.abc import Callable, Sequence
from pathlib import Path
from typing import Final
@ -14,6 +14,87 @@ from litellm.types.prompts.init_prompts import (
prompt_initializer_registry = {}
DEFAULT_PROMPT_ENVIRONMENT: Final = "development"
PROMPT_ENVIRONMENT_SERVE_PRECEDENCE: Final = ("production", "staging", "development")
def get_base_prompt_id(prompt_id: str) -> str:
"""
Extract the base prompt ID by stripping the version suffix if present.
Examples:
>>> get_base_prompt_id("jack_success.v1")
"jack_success"
>>> get_base_prompt_id("jack_success_v1")
"jack_success"
>>> get_base_prompt_id("jack_success")
"jack_success"
"""
if ".v" in prompt_id:
return prompt_id.split(".v")[0]
if "_v" in prompt_id:
return prompt_id.split("_v")[0]
return prompt_id
def get_version_number(prompt_id: str) -> int:
"""
Extract the version number from a versioned prompt ID (defaults to 1).
Examples:
>>> get_version_number("jack_success.v2")
2
>>> get_version_number("jack_success_v2")
2
>>> get_version_number("jack_success")
1
"""
if ".v" in prompt_id:
version_str = prompt_id.split(".v")[1]
try:
return int(version_str)
except ValueError:
pass
if "_v" in prompt_id:
version_str = prompt_id.split("_v")[1]
try:
return int(version_str)
except ValueError:
pass
return 1
def prompt_environment_or_default(environment: str | None) -> str:
return environment or DEFAULT_PROMPT_ENVIRONMENT
def registry_key_for_prompt(prompt: PromptSpec) -> str:
return f"{prompt.prompt_id}::{prompt_environment_or_default(prompt.environment)}"
def parse_prompt_version(raw_version: object) -> int | None:
if isinstance(raw_version, bool):
return None
if isinstance(raw_version, int):
return raw_version
if isinstance(raw_version, str) and raw_version.isdigit():
return int(raw_version)
return None
def _spec_version(prompt: PromptSpec) -> int:
return prompt.version if prompt.version is not None else get_version_number(prompt_id=prompt.prompt_id)
def _default_serve_environment(prompts: Sequence[PromptSpec]) -> str:
present: Final = frozenset(prompt_environment_or_default(prompt.environment) for prompt in prompts)
ladder_pick: Final = next((env for env in PROMPT_ENVIRONMENT_SERVE_PRECEDENCE if env in present), None)
if ladder_pick is not None:
return ladder_pick
return min(present) if present else DEFAULT_PROMPT_ENVIRONMENT
def get_prompt_initializer_from_integrations():
"""
@ -113,17 +194,16 @@ class InMemoryPromptRegistry:
"""
import litellm
prompt_id: Final = prompt.prompt_id
if prompt_id in self.IN_MEMORY_PROMPTS:
verbose_proxy_logger.debug("prompt_id already exists in IN_MEMORY_PROMPTS")
return self.IN_MEMORY_PROMPTS[prompt_id]
registry_key: Final = registry_key_for_prompt(prompt)
if registry_key in self.IN_MEMORY_PROMPTS:
verbose_proxy_logger.debug("prompt already exists in IN_MEMORY_PROMPTS")
return self.IN_MEMORY_PROMPTS[registry_key]
parsed_prompt, custom_prompt_callback = self._build_prompt_callback(prompt=prompt)
litellm.logging_callback_manager.add_litellm_callback(custom_prompt_callback)
# store references to the prompt in memory
self.IN_MEMORY_PROMPTS[prompt_id] = parsed_prompt
self.prompt_id_to_custom_prompt[prompt_id] = custom_prompt_callback
self.IN_MEMORY_PROMPTS[registry_key] = parsed_prompt
self.prompt_id_to_custom_prompt[registry_key] = custom_prompt_callback
return parsed_prompt
@ -166,68 +246,93 @@ class InMemoryPromptRegistry:
import litellm
parsed_prompt, new_callback = self._build_prompt_callback(prompt=prompt)
stale_callback: Final = self.prompt_id_to_custom_prompt.pop(prompt.prompt_id, None)
self.IN_MEMORY_PROMPTS.pop(prompt.prompt_id, None)
registry_key: Final = registry_key_for_prompt(parsed_prompt)
stale_callback: Final = self.prompt_id_to_custom_prompt.pop(registry_key, None)
self.IN_MEMORY_PROMPTS.pop(registry_key, None)
if stale_callback is not None:
litellm.logging_callback_manager.remove_callback_from_all_lists(stale_callback)
litellm.logging_callback_manager.add_litellm_callback(new_callback)
self.IN_MEMORY_PROMPTS[prompt.prompt_id] = parsed_prompt
self.prompt_id_to_custom_prompt[prompt.prompt_id] = new_callback
self.IN_MEMORY_PROMPTS[registry_key] = parsed_prompt
self.prompt_id_to_custom_prompt[registry_key] = new_callback
return parsed_prompt
def sync_prompt_from_db(self, prompt: PromptSpec) -> PromptSpec | None:
existing: Final = self.IN_MEMORY_PROMPTS.get(prompt.prompt_id)
existing: Final = self.IN_MEMORY_PROMPTS.get(registry_key_for_prompt(prompt))
if existing is None:
return self.initialize_prompt(prompt=prompt)
if existing.litellm_params == prompt.litellm_params and existing.prompt_info == prompt.prompt_info:
return existing
return self.reload_prompt(prompt=prompt)
def get_prompt_by_id(self, prompt_id: str) -> PromptSpec | None:
def resolve_prompt_spec(
self,
prompt_id: str,
version: int | None = None,
environment: str | None = None,
) -> PromptSpec | None:
"""
Get a prompt by its ID from memory
"""
return self.IN_MEMORY_PROMPTS.get(prompt_id)
Resolve a prompt spec by base prompt id, optional version, and optional environment.
def get_prompt_callback_by_id(self, prompt_id: str) -> CustomPromptManagement | None:
With no environment, resolves within the default serve environment
(production > staging > development > alphabetical first present).
With no version, resolves to the highest version in the chosen environment.
"""
Get a prompt callback by its ID from memory
"""
return self.prompt_id_to_custom_prompt.get(prompt_id)
base_prompt_id: Final = get_base_prompt_id(prompt_id=prompt_id)
base_matches: Final = tuple(
spec
for spec in self.IN_MEMORY_PROMPTS.values()
if get_base_prompt_id(prompt_id=spec.prompt_id) == base_prompt_id
)
if not base_matches:
return None
resolved_environment: Final = (
environment if environment is not None else _default_serve_environment(base_matches)
)
env_matches: Final = tuple(
spec for spec in base_matches if prompt_environment_or_default(spec.environment) == resolved_environment
)
if not env_matches:
return None
if version is not None:
return next((spec for spec in env_matches if _spec_version(spec) == version), None)
return max(env_matches, key=_spec_version)
def remove_prompt(self, prompt_id: str) -> None:
def get_prompt_callback_for_prompt(self, prompt: PromptSpec) -> CustomPromptManagement | None:
return self.prompt_id_to_custom_prompt.get(registry_key_for_prompt(prompt))
def has_config_prompt(self, base_prompt_id: str) -> bool:
return any(
spec.prompt_info.prompt_type == "config"
for spec in self.IN_MEMORY_PROMPTS.values()
if get_base_prompt_id(prompt_id=spec.prompt_id) == base_prompt_id
)
def remove_prompt(self, registry_key: str) -> None:
import litellm
self.IN_MEMORY_PROMPTS.pop(prompt_id, None)
stale_callback: Final = self.prompt_id_to_custom_prompt.pop(prompt_id, None)
self.IN_MEMORY_PROMPTS.pop(registry_key, None)
stale_callback: Final = self.prompt_id_to_custom_prompt.pop(registry_key, None)
if stale_callback is not None:
litellm.logging_callback_manager.remove_callback_from_all_lists(stale_callback)
def delete_prompts_by_base_id(self, base_prompt_id: str, environment: str | None = None) -> list[str]:
"""
Delete all prompts matching the given base prompt ID from memory, along with their
registered callbacks; scoped to one environment when given.
Delete matching prompts from memory, along with their registered callbacks,
scoped to one environment when given.
Args:
base_prompt_id: The base prompt ID (without version suffix)
environment: When set, only delete prompts deployed to this environment
Returns:
List of prompt IDs that were deleted
Returns the registry keys that were deleted.
"""
from litellm.proxy.prompts.prompt_endpoints import get_base_prompt_id
prompts_to_delete: Final = [
pid
for pid, prompt in self.IN_MEMORY_PROMPTS.items()
if get_base_prompt_id(prompt_id=pid) == base_prompt_id
and (environment is None or prompt.environment == environment)
keys_to_delete: Final = [
key
for key, spec in self.IN_MEMORY_PROMPTS.items()
if get_base_prompt_id(prompt_id=spec.prompt_id) == base_prompt_id
and (environment is None or prompt_environment_or_default(spec.environment) == environment)
]
for pid in prompts_to_delete:
self.remove_prompt(prompt_id=pid)
for key in keys_to_delete:
self.remove_prompt(registry_key=key)
return prompts_to_delete
return keys_to_delete
IN_MEMORY_PROMPT_REGISTRY: Final = InMemoryPromptRegistry()

View file

@ -120,6 +120,8 @@ from litellm.router_utils.add_retry_fallback_headers import (
from litellm.router_utils.auto_router_model_naming import (
STRATEGY_ROUTER_PARAM_FIELDS,
carries_complexity_router_settings,
count_heuristic_v2_routers,
heuristic_v2_limit_violation,
validate_complexity_router_config_placement,
)
from litellm.types.utils import (
@ -301,7 +303,7 @@ from litellm.proxy.auth.auth_utils import (
)
from litellm.proxy.auth.fallback_model_access import router_fallback_access_check
from litellm.proxy.auth.handle_jwt import JWTHandler
from litellm.proxy.auth.litellm_license import LicenseCheck
from litellm.proxy.auth.litellm_license import HEURISTIC_V2_LICENSE_REMEDY, LicenseCheck
from litellm.proxy.auth.model_checks import (
expand_wildcard_deployments_for_model_info,
get_all_fallbacks,
@ -4316,6 +4318,19 @@ def validate_deployment_complexity_router_placement(model: Mapping[str, object])
raise ValueError(f"model {model.get('model_name', '')!r}: {violation}")
def validate_heuristic_v2_router_limit(model_list: Sequence[Mapping[str, object]], *, limit: int | None) -> None:
"""
Refuse to start when config.yaml defines more heuristic_v2 auto-routers than the license allows.
Checked here rather than left to router registration for the same reason as the two
validators above: the proxy builds its router with `ignore_invalid_deployments=True`, so
the router's own refusal would turn the extra router into a silently missing model.
"""
violation: Final = heuristic_v2_limit_violation(held=count_heuristic_v2_routers(model_list), limit=limit)
if violation is not None:
raise ValueError(f"config.yaml model_list: {violation} {HEURISTIC_V2_LICENSE_REMEDY}")
def pin_complexity_router_model_id(model: dict) -> None: # mutable-ok: out-param, model_info is stamped in place
"""
Stamps `model_info.id` from the raw litellm_params before plugin resolution swaps
@ -5721,6 +5736,7 @@ class ProxyConfig:
model_list: Final = config.get("model_list", None)
if model_list:
router_params["model_list"] = model_list
validate_heuristic_v2_router_limit(model_list, limit=_license_check.heuristic_v2_router_limit())
print( # noqa: T201
"\033[32mLiteLLM: Proxy initialized with Config, Set models:\033[0m"
)
@ -5810,6 +5826,7 @@ class ProxyConfig:
),
ignore_invalid_deployments=True, # don't raise an error if a deployment is invalid
fallback_access_check=router_fallback_access_check,
heuristic_v2_router_limit=_license_check.heuristic_v2_router_limit,
)
if redis_usage_cache is not None and router.cache.redis_cache is None:
@ -6270,6 +6287,7 @@ class ProxyConfig:
search_tools=search_tools,
ignore_invalid_deployments=True,
fallback_access_check=router_fallback_access_check,
heuristic_v2_router_limit=_license_check.heuristic_v2_router_limit,
)
verbose_proxy_logger.debug("updated llm_router: %s", llm_router)
else:
@ -7571,7 +7589,7 @@ class ProxyConfig:
return create_versioned_prompt_spec(db_prompt=db_prompt)
async def _init_prompts_in_db(self, prisma_client: PrismaClient):
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY, registry_key_for_prompt
from litellm.types.prompts.init_prompts import PromptSpec
def parse_row(db_prompt: object) -> PromptSpec | None:
@ -7586,21 +7604,12 @@ class ProxyConfig:
return None
try:
prompt_ids_loaded_before_db_read: Final = frozenset(IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS)
registry_keys_loaded_before_db_read: Final = frozenset(IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS)
prompts_in_db: Final[Sequence[object]] = await PromptRepository(prisma_client).table.find_many()
parsed_specs: Final[tuple[PromptSpec, ...]] = tuple(
spec for row in prompts_in_db if (spec := parse_row(row)) is not None
)
newest_spec_per_id: Final[Mapping[str, PromptSpec]] = MappingProxyType(
{
spec.prompt_id: spec
for spec in sorted(
parsed_specs,
key=lambda s: s.updated_at.timestamp() if s.updated_at else float("-inf"),
)
}
)
for prompt_spec in newest_spec_per_id.values():
for prompt_spec in parsed_specs:
try:
IN_MEMORY_PROMPT_REGISTRY.sync_prompt_from_db(prompt=prompt_spec)
except Exception as prompt_sync_error: # noqa: BLE001 # one poisoned row must not block syncing the remaining prompts
@ -7612,15 +7621,16 @@ class ProxyConfig:
# An unparsable row still exists in the DB, so skip the sweep rather than unload its in-memory copy
every_row_parsed: Final = len(parsed_specs) == len(prompts_in_db)
if every_row_parsed:
deleted_db_prompt_ids: Final = tuple(
prompt_id
for prompt_id in prompt_ids_loaded_before_db_read
if (loaded_spec := IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS.get(prompt_id)) is not None
db_registry_keys: Final = frozenset(registry_key_for_prompt(spec) for spec in parsed_specs)
deleted_db_registry_keys: Final = tuple(
registry_key
for registry_key in registry_keys_loaded_before_db_read
if (loaded_spec := IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS.get(registry_key)) is not None
and loaded_spec.prompt_info.prompt_type == "db"
and prompt_id not in newest_spec_per_id
and registry_key not in db_registry_keys
)
for deleted_prompt_id in deleted_db_prompt_ids:
IN_MEMORY_PROMPT_REGISTRY.remove_prompt(prompt_id=deleted_prompt_id)
for deleted_registry_key in deleted_db_registry_keys:
IN_MEMORY_PROMPT_REGISTRY.remove_prompt(registry_key=deleted_registry_key)
except Exception as e:
verbose_proxy_logger.debug("litellm.proxy.proxy_server.py::ProxyConfig:_init_prompts_in_db - %s", e)

View file

@ -1478,28 +1478,27 @@ class ProxyLogging:
) -> None:
"""Process prompt template if applicable."""
from litellm.proxy.prompts.prompt_endpoints import (
construct_versioned_prompt_id,
get_latest_version_prompt_id,
)
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.utils import get_non_default_completion_params
if prompt_version is None:
lookup_prompt_id = get_latest_version_prompt_id(
prompt_id=prompt_id,
all_prompt_ids=IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS,
)
else:
lookup_prompt_id = construct_versioned_prompt_id(prompt_id=prompt_id, version=prompt_version)
custom_logger: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id(lookup_prompt_id)
prompt_spec: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(lookup_prompt_id)
raw_prompt_environment: Final = data.get("prompt_environment", None)
prompt_environment: Final = raw_prompt_environment if isinstance(raw_prompt_environment, str) else None
prompt_spec: Final = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec(
prompt_id,
version=prompt_version,
environment=prompt_environment,
)
custom_logger: Final = (
IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_for_prompt(prompt=prompt_spec)
if prompt_spec is not None
else None
)
litellm_prompt_id: str | None = None
if prompt_spec is not None:
litellm_prompt_id = prompt_spec.litellm_params.prompt_id
data.pop("prompt_id", None)
data.pop("prompt_environment", None)
if custom_logger and prompt_spec is not None:
is_responses_call: Final = call_type == "aresponses"
@ -1542,6 +1541,7 @@ class ProxyLogging:
data.pop("prompt_variables", None)
data.pop("prompt_label", None)
data.pop("prompt_version", None)
data.pop("prompt_environment", None)
def _process_guardrail_metadata(self, data: dict) -> None:
"""Process guardrails from metadata and add to applied_guardrails."""
@ -1750,7 +1750,6 @@ class ProxyLogging:
litellm_logging_obj: Final = cast(Optional["LiteLLMLoggingObj"], data.get("litellm_logging_obj", None))
prompt_id: Final[str | None] = data.get("prompt_id", None)
prompt_version: Final[int | None] = data.get("prompt_version", None)
## PROMPT TEMPLATE CHECK ##
@ -1760,11 +1759,13 @@ class ProxyLogging:
and prompt_id is not None
and (call_type == "completion" or call_type == "acompletion" or call_type == "aresponses")
):
from litellm.proxy.prompts.prompt_registry import parse_prompt_version
await self._process_prompt_template(
data=data,
litellm_logging_obj=litellm_logging_obj,
prompt_id=prompt_id,
prompt_version=prompt_version,
prompt_version=parse_prompt_version(data.get("prompt_version", None)),
call_type=call_type,
)

View file

@ -21,7 +21,7 @@ import time
import traceback
import weakref
from collections import defaultdict
from collections.abc import AsyncGenerator, AsyncIterator, Callable, Generator, Mapping, Sequence
from collections.abc import AsyncGenerator, AsyncIterator, Callable, Generator, Iterator, Mapping, Sequence
from functools import lru_cache, partial
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypeAlias, TypeVar, Union, cast
@ -117,6 +117,9 @@ from litellm.router_utils.add_retry_fallback_headers import (
from litellm.router_utils.auto_router_model_naming import (
AUTO_ROUTER_MODEL_PREFIX,
classify_strategy_router_model,
count_heuristic_v2_routers,
heuristic_v2_limit_violation,
uses_heuristic_v2_classifier,
)
from litellm.router_utils.batch_utils import (
_get_router_metadata_variable_name,
@ -211,6 +214,7 @@ from litellm.types.router import (
DeploymentTypedDict,
FallbackAccessCheck,
GuardrailTypedDict,
HeuristicV2RouterLimit,
LiteLLM_Params,
MockRouterTestingParams,
ModelGroupInfo,
@ -683,6 +687,7 @@ class Router:
background_health_check_model_groups: Sequence[str] | None = None,
enable_weighted_failover: bool = False,
fallback_access_check: FallbackAccessCheck | None = None,
heuristic_v2_router_limit: HeuristicV2RouterLimit | None = None,
) -> None:
"""
Initialize the Router class with the given parameters for caching, reliability, and routing strategy.
@ -759,6 +764,7 @@ class Router:
self.set_verbose = set_verbose
self.ignore_invalid_deployments = ignore_invalid_deployments
self.heuristic_v2_router_limit = heuristic_v2_router_limit
self.fallback_access_check: Final = fallback_access_check
self.debug_level = debug_level
self.enable_pre_call_checks = enable_pre_call_checks
@ -8796,6 +8802,30 @@ class Router:
"""
return classify_strategy_router_model(litellm_params.model) == "complexity"
def config_deployments(self) -> Iterator[Mapping[str, object]]:
"""The model_list rows that came from config.yaml rather than the DB (``model_info.db_model`` unset)."""
for deployment in self.model_list:
if not isinstance(deployment, Mapping):
continue
model_info = deployment.get("model_info")
if not (isinstance(model_info, Mapping) and model_info.get("db_model")):
yield deployment
def heuristic_v2_router_limit_violation(self) -> str | None:
"""
Why one more heuristic_v2 router cannot join this router, or None when it can.
Judged against every deployment currently on the model_list; an upsert pops the row being
edited first, so an edit of an existing heuristic_v2 router keeps its own slot. The limit is
resolved on every call through ``heuristic_v2_router_limit``; unset means unlimited, which
is the SDK default, and the proxy injects a resolver backed by its license.
"""
limit: Final = self.heuristic_v2_router_limit() if self.heuristic_v2_router_limit is not None else None
others: Final = count_heuristic_v2_routers(
deployment for deployment in self.model_list if isinstance(deployment, Mapping)
)
return heuristic_v2_limit_violation(held=others + 1, limit=limit)
def init_complexity_router_deployment(self, deployment: Deployment):
"""
Initialize the complexity-router deployment.
@ -8813,6 +8843,10 @@ class Router:
)
complexity_router_config: Final[dict | None] = deployment.litellm_params.complexity_router_config
if uses_heuristic_v2_classifier(complexity_router_config):
limit_violation: Final = self.heuristic_v2_router_limit_violation()
if limit_violation is not None:
raise ValueError(limit_violation)
default_model: str | None = deployment.litellm_params.complexity_router_default_model
@ -9636,8 +9670,16 @@ class Router:
raise e
def _restore_deployment_after_failed_upsert(self, previous_deployment: Deployment | None, model_id: str) -> None:
"""Put a deployment back the way it was before a failed upsert popped it.
A rollback re-admits state that was already serving, so it does not go through the
heuristic_v2 ceiling a newcomer gets: with the ceiling tightened since the deployment first
registered, judging the rollback would drop a serving router over an unrelated failed edit.
"""
if previous_deployment is None or self.has_model_id(model_id):
return
limit_resolver: Final = self.heuristic_v2_router_limit
self.heuristic_v2_router_limit = None
try:
self.add_deployment(deployment=previous_deployment)
verbose_router_logger.info(
@ -9652,6 +9694,8 @@ class Router:
model_id,
restore_error,
)
finally:
self.heuristic_v2_router_limit = limit_resolver
@staticmethod
def _backend_cost_map_keys(model: str, custom_llm_provider: str | None) -> tuple[str, ...]:

View file

@ -2,23 +2,41 @@
Auto-Routing Strategy that works with a Semantic Router Config
"""
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final, Optional
from pydantic import BaseModel, ConfigDict
from litellm._logging import verbose_router_logger
from litellm.constants import DEFAULT_AUTO_ROUTER_MAX_INPUT_CHARS
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.internal_call_metadata import (
effective_turn_off_message_logging,
forwarded_internal_call_metadata,
parent_session_kwargs,
)
from litellm.types.utils import AUTOROUTER_CLASSIFIER_CALL_ORIGIN
if TYPE_CHECKING:
from semantic_router.routers import SemanticRouter
from semantic_router.routers.base import Route
from litellm.router import Router
from litellm.router_strategy.auto_router.litellm_encoder import LiteLLMRouterEncoder
from litellm.types.router import PreRoutingHookResponse
else:
Router = Any
PreRoutingHookResponse = Any
Route = Any
SemanticRouter = Any
LiteLLMRouterEncoder = Any
class _CallerMetadata(BaseModel):
model_config = ConfigDict(frozen=True, extra="ignore")
metadata: Mapping[str, object] | None = None
litellm_metadata: Mapping[str, object] | None = None
class AutoRouter(CustomLogger):
@ -50,6 +68,8 @@ class AutoRouter(CustomLogger):
"""
from semantic_router.routers import SemanticRouter
from litellm.router_strategy.auto_router.litellm_encoder import LiteLLMRouterEncoder
self.auto_router_config_path: str | None = auto_router_config_path
self.auto_router_config: str | None = auto_router_config
self.auto_sync_value = self.DEFAULT_AUTO_SYNC_VALUE
@ -59,6 +79,11 @@ class AutoRouter(CustomLogger):
self.embedding_model: str = embedding_model
self.max_input_chars: int = max_input_chars
self.litellm_router_instance: Router = litellm_router_instance
self.encoder: LiteLLMRouterEncoder = LiteLLMRouterEncoder(
litellm_router_instance=litellm_router_instance,
model_name=embedding_model,
max_input_chars=max_input_chars,
)
def _load_semantic_routing_routes(self) -> list[Route]:
from semantic_router.routers import SemanticRouter
@ -129,9 +154,6 @@ class AutoRouter(CustomLogger):
from semantic_router.routers import SemanticRouter
from litellm.litellm_core_utils.prompt_templates.factory import resolve_structured_messages
from litellm.router_strategy.auto_router.litellm_encoder import (
LiteLLMRouterEncoder,
)
from litellm.types.router import PreRoutingHookResponse
resolved_messages: Final = (
@ -149,34 +171,47 @@ class AutoRouter(CustomLogger):
#######################
routelayer = SemanticRouter(
routes=self.loaded_routes,
encoder=LiteLLMRouterEncoder(
litellm_router_instance=self.litellm_router_instance,
model_name=self.embedding_model,
max_input_chars=self.max_input_chars,
),
encoder=self.encoder,
auto_sync=self.auto_sync_value,
)
self.routelayer = routelayer
message_content: Final = self._extract_text_from_messages(resolved_messages)
route_name: Final = self._matched_route_name(routelayer, message_content)
route_name: Final = await self._matched_route_name(routelayer, message_content, request_kwargs)
return PreRoutingHookResponse(
model=route_name or self.default_model,
messages=messages,
)
def _matched_route_name(self, routelayer: "SemanticRouter", text: str) -> str | None:
async def _matched_route_name(
self, routelayer: "SemanticRouter", text: str, request_kwargs: Mapping[str, object]
) -> str | None:
"""Name of the route `text` matches, or None when nothing matched or the match failed.
The route layer embeds `text` to compare it against the routes, and that embedding call can
`text` is embedded here rather than by `routelayer(text=...)` so the caller's metadata reaches
`aembedding()` and the embedding's spend lands on the key/team that sent the request;
SemanticRouter has no way to pass kwargs through to its encoder. That embedding call can
fail (context limit, timeout, provider error). Choosing a model is a routing decision, so a
failure here falls back to the default model rather than failing the user's request.
"""
from semantic_router.schema import RouteChoice
try:
route_choice: Final = routelayer(text=text)
caller: Final = _CallerMetadata.model_validate(request_kwargs)
query_vector: Final = (
await self.encoder.aencode_queries(
[text],
metadata=forwarded_internal_call_metadata(caller.metadata, AUTOROUTER_CLASSIFIER_CALL_ORIGIN),
litellm_metadata=forwarded_internal_call_metadata(
caller.litellm_metadata, AUTOROUTER_CLASSIFIER_CALL_ORIGIN
),
proxy_server_request={"body": {"model": self.embedding_model, "input": [text]}},
turn_off_message_logging=effective_turn_off_message_logging(request_kwargs),
**parent_session_kwargs(request_kwargs),
)
)[0]
route_choice: Final = await routelayer.acall(vector=query_vector)
except Exception as e: # noqa: BLE001 -- the embedding call behind the route layer can fail many ways (context limit, timeout, provider/network error); none of them may fail the request
verbose_router_logger.warning(
"AutoRouter: semantic routing failed (%s), falling back to default model %s", e, self.default_model

View file

@ -10,7 +10,7 @@ the router silently dropping the deployment at load time under
``ignore_invalid_deployments``.
"""
from collections.abc import Mapping, Sequence
from collections.abc import Iterable, Mapping, Sequence
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final, Literal, TypeAlias
@ -163,6 +163,38 @@ def strategy_router_dependencies(
)
def uses_heuristic_v2_classifier(complexity_router_config: object) -> bool:
"""Whether this complexity config classifies with the bundled heuristic_v2 model."""
return _mapping(complexity_router_config).get("classifier_type") == "heuristic_v2"
def is_heuristic_v2_router(litellm_params: Mapping[str, object]) -> bool:
"""Whether this deployment is a complexity router that classifies with heuristic_v2."""
return classify_strategy_router_model(str(litellm_params.get("model") or "")) == "complexity" and (
uses_heuristic_v2_classifier(litellm_params.get("complexity_router_config"))
)
def count_heuristic_v2_routers(deployments: Iterable[Mapping[str, object]]) -> int:
"""How many of ``deployments`` (router model_list entries or config.yaml rows) are heuristic_v2 routers."""
return sum(1 for deployment in deployments if is_heuristic_v2_router(_mapping(deployment.get("litellm_params"))))
def heuristic_v2_limit_violation(*, held: int, limit: int | None) -> str | None:
"""Why holding ``held`` heuristic_v2 routers exceeds ``limit``, or None when it fits.
``limit`` None means unlimited. The message is shared by every enforcement point (config
load, model writes, router registration) and stays SDK-neutral: it names the cap and what
the caller can change; the proxy appends how its license lifts the cap.
"""
if limit is None or held <= limit:
return None
return (
f"At most {limit} auto-router(s) with classifier_type 'heuristic_v2' can be registered but this would make "
f"{held}. Use classifier_type 'heuristic' for this router or remove an existing heuristic_v2 router."
)
def validate_complexity_router_config_write(complexity_router_config: Mapping[str, object] | None) -> str | None:
"""Reject a complexity config the router would refuse to build a deployment from.

View file

@ -1,6 +1,8 @@
from typing import Any, Final
from collections.abc import Mapping
from typing import Any, Final, Literal
from pydantic import BaseModel, field_validator
from typing_extensions import ReadOnly, TypedDict
from litellm.proxy._types import (
LiteLLM_UserTableWithKeyCount,
@ -9,6 +11,17 @@ from litellm.proxy._types import (
)
class InsensitiveContains(TypedDict):
contains: ReadOnly[str]
mode: ReadOnly[Literal["insensitive"]]
class UserSearchWhere(TypedDict):
"""Prisma filter behind `/user/list?search=`: user_id or user_email contains the term, case-insensitive."""
OR: ReadOnly[tuple[Mapping[Literal["user_id", "user_email"], InsensitiveContains], ...]]
class UserListResponse(BaseModel):
"""
Response model for the user list endpoint

View file

@ -885,6 +885,18 @@ class FallbackAccessCheck(Protocol):
async def __call__(self, *, model: str, request_kwargs: Mapping[str, object], llm_router: "Router") -> bool: ...
class HeuristicV2RouterLimit(Protocol):
"""
Resolves how many heuristic_v2 complexity routers the Router may hold right now; None means unlimited.
The Router calls it on every registration and limit query instead of caching the answer, so the
proxy can keep the limit on its license object (re-verified on config load) rather than hand
over a snapshot.
"""
def __call__(self) -> int | None: ...
class LiteLLM_RouterFileObject(TypedDict, total=False):
"""
Tracking the litellm params hash, used for mapping the file id to the right model

View file

@ -3612,6 +3612,7 @@ all_litellm_params = (
"litellm_system_prompt",
"provider_specific_header",
"prompt_version",
"prompt_environment",
"api_base",
"force_timeout",
"logger_fn",

View file

@ -1,6 +1,13 @@
import time
from collections.abc import Iterator
from typing import Final
import httpx
from openai import OpenAI, BadRequestError, NotFoundError, APIStatusError
import pytest
from openai import APIStatusError, BadRequestError, NotFoundError, OpenAI, Stream
from openai.types.responses import ResponseStreamEvent
BACKGROUND_STREAM_ADMISSION_DEADLINE_SECONDS: Final = 90
def generate_key():
@ -153,43 +160,48 @@ def test_cancel_response():
raise e
def admitted_response_id(chunk: ResponseStreamEvent) -> str | None:
response: Final = getattr(chunk, "response", None)
return None if response is None else response.id
def events_until_admission(stream: Stream[ResponseStreamEvent], started: float) -> Iterator[ResponseStreamEvent]:
for chunk in stream:
print("stream chunk=", chunk)
yield chunk
if admitted_response_id(chunk) is not None:
return
if time.monotonic() - started > BACKGROUND_STREAM_ADMISSION_DEADLINE_SECONDS:
return
def test_cancel_streaming_response():
try:
client = get_test_client()
from litellm.types.llms.openai import ResponsesAPIResponse
client: Final = get_test_client()
started: Final = time.monotonic()
stream: Final = client.responses.create(
model="gpt-5.5",
input="count from 1 to 500, one number per line",
stream=True,
background=True,
timeout=BACKGROUND_STREAM_ADMISSION_DEADLINE_SECONDS,
)
stream = client.responses.create(
model="gpt-5.5",
input="just respond with the word 'ping'",
stream=True,
background=True,
with stream:
events: Final = tuple(events_until_admission(stream, started))
elapsed: Final = time.monotonic() - started
keepalive_events: Final = sum(1 for chunk in events if chunk.type == "keepalive")
response_id: Final = next((rid for rid in map(admitted_response_id, events) if rid is not None), None)
if response_id is None and keepalive_events:
pytest.skip(
f"OpenAI held the background stream in keepalive for {elapsed:.0f}s "
f"({keepalive_events} keepalive events) without creating the response"
)
assert response_id is not None, f"no response event within {elapsed:.0f}s of streaming a background response"
collected_chunks = []
response_id = None
for chunk in stream:
print("stream chunk=", chunk)
collected_chunks.append(chunk)
# Extract response ID from the first chunk that has it
if (
response_id is None
and hasattr(chunk, "response")
and hasattr(chunk.response, "id")
):
response_id = chunk.response.id
assert len(collected_chunks) > 0
# cancel the response if we got a response ID
if response_id:
cancel_response = client.responses.cancel(response_id)
print("CANCEL streaming response=", cancel_response)
assert hasattr(cancel_response, "id")
except Exception as e:
if "Cannot cancel a completed response" in str(e):
pass
else:
raise e
cancel_response: Final = client.responses.cancel(response_id)
print("CANCEL streaming response=", cancel_response)
assert cancel_response.status == "cancelled"
def test_cancel_invalid_response_id():

View file

@ -336,3 +336,15 @@ class TestAzureResolvesTheDeclaredDefaultEffort:
drop_params=True,
)
assert ("temperature" in mapped) is temperature_survives
def test_azure_gpt_6_astra_takes_the_reasoning_series_request_shape():
params = litellm.get_optional_params(
model="gpt-6-astra",
custom_llm_provider="azure",
max_tokens=100,
reasoning_effort="max",
)
assert params["max_completion_tokens"] == 100
assert "max_tokens" not in params
assert params["reasoning_effort"] == "max"

View file

@ -1718,6 +1718,8 @@ class TestResponsesSurfaceSharesTheEffortRule:
("gpt-5.6-sol", None, False),
("gpt-5.6-terra", "none", True),
("gpt-5.6-terra", "medium", False),
("gpt-6-astra", None, False),
("gpt-6-astra", "low", False),
],
)
def test_temperature_follows_the_resolved_effort(

View file

@ -1505,3 +1505,17 @@ class TestACatalogueOlderThanTheCodeDoesNotStripTemperature:
drop_params=True,
)
assert "temperature" not in mapped
def test_gpt_6_astra_takes_the_reasoning_series_request_shape():
params = litellm.get_optional_params(
model="gpt-6-astra",
custom_llm_provider="openai",
max_tokens=100,
reasoning_effort="max",
verbosity="low",
)
assert params["max_completion_tokens"] == 100
assert "max_tokens" not in params
assert params["reasoning_effort"] == "max"
assert params["verbosity"] == "low"

View file

@ -41,6 +41,8 @@ from litellm.llms.azure.chat.gpt_5_transformation import AzureOpenAIGPT5Config
# Models that MUST be classified as GPT-5 (routed through GPT-5 reasoning path)
GPT5_MODELS = [
"gpt-6-astra",
"openai/gpt-6-astra",
"gpt-5",
"gpt-5.1",
"gpt-5.2",
@ -120,6 +122,8 @@ class TestOpenAIGPT5ConfigIsModelGpt5Model:
# /v1/responses bridge (when reasoning_effort is set and tools are passed) on
# is_model_gpt_5_4_plus_model, so the gpt-5.6 family must land on the True side.
GPT5_4_PLUS_MODELS = [
"gpt-6-astra",
"openai/gpt-6-astra",
"gpt-5.4",
"gpt-5.5",
"gpt-5.5-pro",

View file

@ -964,6 +964,48 @@ def test_construct_target_url_with_version_prefix():
assert str(target_url) == expected_url
@pytest.mark.parametrize(
("requested_route", "expected_url"),
[
(
"/projects/test-project/locations/global/publishers/anthropic/models/claude-sonnet-4-6:streamRawPredict",
"https://aiplatform.googleapis.com/v1/projects/test-project/locations/global/publishers/anthropic/models/claude-sonnet-4-6:streamRawPredict",
),
(
"/projects/test-project/locations/global/publishers/anthropic/models/count-tokens:rawPredict",
"https://aiplatform.googleapis.com/v1/projects/test-project/locations/global/publishers/anthropic/models/count-tokens:rawPredict",
),
(
"/projects/other-project/locations/us-east5/publishers/anthropic/models/claude-sonnet-4-6:rawPredict",
"https://aiplatform.googleapis.com/v1/projects/test-project/locations/global/publishers/anthropic/models/claude-sonnet-4-6:rawPredict",
),
(
"/projects/test-project/locations/global/cachedContents",
"https://aiplatform.googleapis.com/v1beta1/projects/test-project/locations/global/cachedContents",
),
(
"/v1/projects/test-project/locations/global/publishers/anthropic/models/claude-sonnet-4-6:streamRawPredict",
"https://aiplatform.googleapis.com/v1/projects/test-project/locations/global/publishers/anthropic/models/claude-sonnet-4-6:streamRawPredict",
),
(
"/v1beta1/projects/test-project/locations/global/cachedContents",
"https://aiplatform.googleapis.com/v1beta1/projects/test-project/locations/global/cachedContents",
),
],
)
def test_construct_target_url_versionless_project_route_gets_api_version(requested_route: str, expected_url: str) -> None:
from litellm.llms.vertex_ai.common_utils import construct_target_url
target_url = construct_target_url(
base_url="https://aiplatform.googleapis.com",
requested_route=requested_route,
vertex_project="test-project",
vertex_location="global",
)
assert str(target_url) == expected_url
def test_fix_enum_types():
"""
Test _fix_enum_types function removes enum fields when type is not string.

View file

@ -2,6 +2,8 @@ import asyncio
import json
from unittest.mock import AsyncMock, MagicMock, patch
from cryptography.hazmat.primitives.asymmetric.rsa import RSAPublicKey
from litellm.proxy.auth.litellm_license import LicenseCheck
@ -30,3 +32,70 @@ def test_is_over_limit():
assert license_check.is_over_limit(101) is False
assert license_check.is_over_limit(100) is False
assert license_check.is_over_limit(99) is False
def test_heuristic_v2_router_limit() -> None:
"""Only the signed license's auto_router feature lifts the one-router limit; an API-verified
license (no airgapped data) and an airgapped license without the feature keep it."""
license_check = LicenseCheck()
license_check.airgapped_license_data = {"expiration_date": "2999-01-01", "allowed_features": ["auto_router"]}
assert license_check.heuristic_v2_router_limit() is None
license_check.airgapped_license_data = {
"expiration_date": "2999-01-01",
"allowed_features": ["sso", "auto_router", "audit_logs"],
}
assert license_check.heuristic_v2_router_limit() is None
license_check.airgapped_license_data = {"expiration_date": "2999-01-01", "allowed_features": ["sso"]}
assert license_check.heuristic_v2_router_limit() == 1
license_check.airgapped_license_data = {"expiration_date": "2999-01-01"}
assert license_check.heuristic_v2_router_limit() == 1
license_check.airgapped_license_data = None
assert license_check.heuristic_v2_router_limit() == 1
def _signed_license(expiration_date: str) -> tuple[RSAPublicKey, str]:
import base64
from cryptography.hazmat.primitives import hashes
from cryptography.hazmat.primitives.asymmetric import padding, rsa
private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048)
message = json.dumps(
{"expiration_date": expiration_date, "user_id": "u", "allowed_features": ["auto_router"]}
).encode()
signature = private_key.sign(
message,
padding.PSS(mgf=padding.MGF1(hashes.SHA256()), salt_length=padding.PSS.MAX_LENGTH),
hashes.SHA256(),
)
return private_key.public_key(), base64.b64encode(message + b"." + signature).decode()
def test_expired_or_unreadable_license_grants_no_features() -> None:
"""The verifier stores the signed payload only after the expiry check passes and clears it when a
later verify rejects the license, so a stale payload cannot keep lifting the heuristic_v2 limit."""
license_check = LicenseCheck()
public_key, valid_key = _signed_license("2999-01-01")
assert license_check.verify_license_without_api_request(public_key=public_key, license_key=valid_key) is True
assert license_check.heuristic_v2_router_limit() is None
_, expired_key = _signed_license("2000-01-01")
assert license_check.verify_license_without_api_request(public_key=public_key, license_key=expired_key) is not True
assert license_check.airgapped_license_data is None
assert license_check.heuristic_v2_router_limit() == 1
assert license_check.verify_license_without_api_request(public_key=public_key, license_key=valid_key) is True
assert license_check.verify_license_without_api_request(public_key=public_key, license_key="not-a-license") is not True
assert license_check.airgapped_license_data is None
def test_valid_signed_license_with_auto_router_lifts_the_limit() -> None:
license_check = LicenseCheck()
public_key, license_key = _signed_license("2999-01-01")
assert license_check.verify_license_without_api_request(public_key=public_key, license_key=license_key) is True
assert license_check.heuristic_v2_router_limit() is None

View file

@ -2014,6 +2014,67 @@ async def test_get_users_user_id_partial_match(mocker):
assert captured_where_conditions["user_id"]["in"] == ["user1", "user2", "user3"]
def test_get_users_search_matches_user_id_or_email(mocker):
"""
`search` ORs a case-insensitive contains match over user_id and user_email on both the rows
query and the count, while the legacy `user_email` param keeps filtering only user_email.
"""
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
searched_user_id = "a6f5c02b-0163-45ce-815f-f88d10e95686"
mock_user_row = mocker.MagicMock()
mock_user_row.user_id = searched_user_id
mock_user_row.model_dump.return_value = {
"user_id": searched_user_id,
"user_email": "search@example.com",
"user_role": "internal_user",
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
}
find_many_wheres = []
count_wheres = []
async def mock_find_many(*args, **kwargs):
find_many_wheres.append(kwargs["where"])
return [mock_user_row]
async def mock_count(*args, **kwargs):
count_wheres.append(kwargs["where"])
return 1
async def mock_key_count(*args, **kwargs):
return 0
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.db.litellm_usertable.find_many = mock_find_many
mock_prisma_client.db.litellm_usertable.count = mock_count
mock_prisma_client.db.litellm_verificationtoken.count = mock_key_count
mocker.patch( # test-quality-ok: /user/list reads prisma_client off proxy_server at call time
"litellm.proxy.proxy_server.prisma_client", mock_prisma_client
)
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN
)
try:
search_response = client.get("/user/list", params={"search": "A6F5C02B-0163"})
assert search_response.status_code == 200, search_response.text
expected_or = (
{"user_id": {"contains": "A6F5C02B-0163", "mode": "insensitive"}},
{"user_email": {"contains": "A6F5C02B-0163", "mode": "insensitive"}},
)
assert find_many_wheres == [{"OR": expected_or}]
assert count_wheres == [{"OR": expected_or}]
assert [user["user_id"] for user in search_response.json()["users"]] == [searched_user_id]
assert search_response.json()["total"] == 1
legacy_response = client.get("/user/list", params={"user_email": "search@example.com"})
assert legacy_response.status_code == 200, legacy_response.text
assert find_many_wheres[-1] == {"user_email": {"contains": "search@example.com", "mode": "insensitive"}}
finally:
app.dependency_overrides.pop(user_api_key_auth, None)
def test_update_internal_user_params_reset_max_budget_with_none():
"""
Test that _update_internal_user_params allows setting max_budget to None.

View file

@ -28,9 +28,18 @@ from litellm.proxy.management_endpoints.model_management_endpoints import (
delete_team_models,
)
from litellm.proxy.utils import PrismaClient
from litellm.router import Router
from litellm.types.router import Deployment, LiteLLM_Params, updateDeployment
async def _passthrough_row(update_data):
return update_data
async def _write_empty_row(**kwargs):
return await kwargs["write_row"]({})
class MockPrismaClient:
def __init__(
self,
@ -1191,7 +1200,7 @@ class TestTeamModelSiblingRouting:
team_id = "team_no_alias"
public_name = "gpt-4.1-mini"
async def mock_add_model_to_db(model_params, user_api_key_dict, prisma_client):
async def mock_add_model_to_db(model_params, user_api_key_dict, prisma_client, slot=None):
return MagicMock(model_id=str(uuid.uuid4()))
mock_team_model_add = AsyncMock()
@ -1372,7 +1381,8 @@ class TestTeamModelUpdate:
db_model=db_model,
patch_data=patch_data,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client, # type: ignore
prisma_client=prisma_client, # type: ignore,
write_row=_passthrough_row,
)
assert result.get("model_name", "").startswith("model_name_test_team_123_")
@ -1435,7 +1445,6 @@ class TestTeamModelUpdate:
team_id="team_123",
public_model_name="new-public-name",
db_model=db_model,
patch_data=patch_data,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client, # type: ignore
)
@ -1481,7 +1490,6 @@ class TestTeamModelUpdate:
team_id="team_123",
public_model_name="new-public-name",
db_model=db_model,
patch_data=patch_data,
user_api_key_dict=user_api_key_dict,
prisma_client=None,
)
@ -1490,39 +1498,72 @@ class TestTeamModelUpdate:
mock_delete.assert_not_called()
@pytest.mark.asyncio
async def test_rename_with_prisma_none_clears_patch_model_name(self):
"""Rename path must clear patch_data.model_name even when prisma is unavailable (P1)."""
async def test_a_refused_row_write_leaves_the_team_untouched(self):
"""The team's model list autocommits, so it is written only after the row write succeeded: a
refused write (the heuristic_v2 slot 403, a DB error) must not leave the team listing a name
whose row never changed."""
from fastapi import HTTPException
from litellm.proxy.management_endpoints.model_management_endpoints import (
_update_existing_team_model_assignment,
_update_team_model_in_db,
)
from litellm.types.router import ModelInfo
db_model = Deployment(
model_name="model_name_team_123_uuid1",
model_name="gpt-4o",
litellm_params=LiteLLM_Params(model="azure/gpt-4o-mini"),
model_info=ModelInfo(
team_id="team_123", team_public_model_name="old-public-name"
model_info=ModelInfo(),
)
user_api_key_dict = UserAPIKeyAuth(user_id="test_user", user_role=LitellmUserRoles.PROXY_ADMIN)
events: list[str] = []
written: dict[str, object] = {}
def patch_data() -> updateDeployment:
return updateDeployment(model_name="team-public", model_info=ModelInfo(team_id="team_123"))
async def refuse_row(update_data):
events.append("row")
raise HTTPException(status_code=403, detail="slot held")
async def accept_row(update_data):
events.append("row")
written.update(update_data)
return update_data
async def team_add(**_):
events.append("team_model_add")
with (
patch( # test-quality-ok: the team auth check needs a live DB; the write order is what is under test
"litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.allow_team_model_action",
AsyncMock(return_value=True),
),
)
patch_data = updateDeployment(
model_name="new-public-name",
model_info=ModelInfo(team_id="team_123"),
)
user_api_key_dict = UserAPIKeyAuth(
user_id="test_user",
user_role=LitellmUserRoles.PROXY_ADMIN,
)
patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: team models are premium-gated through a proxy global with no injection seam
patch( # test-quality-ok: the team list write is the collaborator whose ordering is asserted
"litellm.proxy.management_endpoints.model_management_endpoints.team_model_add",
side_effect=team_add,
),
):
with pytest.raises(HTTPException):
await _update_team_model_in_db(
db_model=db_model,
patch_data=patch_data(),
user_api_key_dict=user_api_key_dict,
prisma_client=MockPrismaClient(team_exists=True), # type: ignore
write_row=refuse_row,
)
assert events == ["row"]
await _update_existing_team_model_assignment(
team_id="team_123",
public_model_name="new-public-name",
db_model=db_model,
patch_data=patch_data,
user_api_key_dict=user_api_key_dict,
prisma_client=None,
)
assert patch_data.model_name is None
await _update_team_model_in_db(
db_model=db_model,
patch_data=patch_data(),
user_api_key_dict=user_api_key_dict,
prisma_client=MockPrismaClient(team_exists=True), # type: ignore
write_row=accept_row,
)
assert events == ["row", "row", "team_model_add"]
assert str(written["model_name"]).startswith("model_name_team_123_")
assert "team-public" in str(written["model_info"])
@pytest.mark.asyncio
async def test_rename_handles_legacy_string_model_info(self):
@ -1574,7 +1615,6 @@ class TestTeamModelUpdate:
team_id="team_123",
public_model_name="new-public-name",
db_model=db_model,
patch_data=patch_data,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client, # type: ignore
)
@ -1614,7 +1654,8 @@ class TestTeamModelUpdate:
db_model=db_model,
patch_data=patch_data,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client, # type: ignore
prisma_client=prisma_client, # type: ignore,
write_row=_passthrough_row,
)
assert "403" in str(exc_info.value)
@ -1900,7 +1941,8 @@ class TestTeamModelUpdate:
db_model=db_model,
patch_data=patch_data,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client, # type: ignore
prisma_client=prisma_client, # type: ignore,
write_row=_passthrough_row,
)
# team ACL must not be touched on a no-op edit
@ -4311,6 +4353,321 @@ class TestStrategyRouterWriteValidation:
is None
)
@staticmethod
def _live_router_holding_one_heuristic_v2(limit: int | None) -> Router:
return Router(
model_list=[
{"model_name": "gpt-4o-mini", "litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "k"}},
{
"model_name": "held-v2",
"litellm_params": {
"model": "auto_router/complexity_router",
"complexity_router_config": {"classifier_type": "heuristic_v2", "tiers": {"SIMPLE": "gpt-4o-mini"}},
},
"model_info": {"id": "held-id"},
},
],
heuristic_v2_router_limit=lambda: limit,
)
class _FakeTx:
"""Stands in for a prisma transaction: records the raw statements and exposes the model table."""
def __init__(self, db_held: int) -> None:
self.db_held = db_held
self.raw_calls: list[tuple[str, tuple[object, ...]]] = []
self.litellm_proxymodeltable = MagicMock(create=AsyncMock(), update=AsyncMock())
async def query_raw(self, sql: str, *args: object) -> list[dict[str, object]]:
self.raw_calls.append((sql, args))
return [{"held": self.db_held}] if "count(*)" in sql else []
async def __aenter__(self) -> "TestStrategyRouterWriteValidation._FakeTx":
return self
async def __aexit__(self, *exc: object) -> None:
return None
class _FakeDb:
"""Stands in for prisma_client: the plain client and the transaction it opens are told apart by identity."""
def __init__(self, db_held: int, existing_row: object = None) -> None:
self.db = self
self.tx_obj = TestStrategyRouterWriteValidation._FakeTx(db_held)
self.litellm_proxymodeltable = MagicMock(
create=AsyncMock(), update=AsyncMock(), find_unique=AsyncMock(return_value=existing_row)
)
def tx(self) -> "TestStrategyRouterWriteValidation._FakeTx":
return self.tx_obj
_V2 = {"classifier_type": "heuristic_v2", "tiers": {"SIMPLE": "gpt-4o-mini"}}
_V1 = {"classifier_type": "heuristic", "tiers": {"SIMPLE": "gpt-4o-mini"}}
@pytest.mark.parametrize(
"incoming,existing,expected",
[
(_V2, None, _V2),
(_V2, _V1, _V2),
(None, _V1, _V1),
(None, None, None),
("no-config", _V2, _V2),
],
)
def test_effective_complexity_router_config(
self, incoming: object, existing: object, expected: object
) -> None:
"""A write is judged on the config it leaves on the row: the incoming one when it carries one, else the stored one."""
from litellm.proxy.management_endpoints.model_management_endpoints import (
_effective_complexity_router_config,
)
from litellm.types.router import updateLiteLLMParams
incoming_params = None if incoming is None else updateLiteLLMParams(
complexity_router_config=None if incoming == "no-config" else incoming
)
existing_params = None if existing is None else updateLiteLLMParams(complexity_router_config=existing)
assert _effective_complexity_router_config(incoming_params, existing_params) == expected
@pytest.mark.asyncio
@pytest.mark.parametrize(
"limit,effective_config,db_held,config_holds_one,model_id,expected",
[
(1, _V2, 1, False, None, "refused"),
(1, _V2, 0, True, None, "refused"),
(1, _V2, 0, False, None, "reserved"),
(1, _V2, 0, False, "held-id", "reserved"),
(2, _V2, 1, False, None, "reserved"),
(1, _V1, 5, True, None, "plain"),
(1, None, 5, True, None, "plain"),
(None, _V2, 5, True, None, "plain"),
],
)
async def test_heuristic_v2_slot_matrix(
self,
limit: int | None,
effective_config: object,
db_held: int,
config_holds_one: bool,
model_id: str | None,
expected: str,
) -> None:
"""The slot is claimed inside a locked transaction only for a heuristic_v2 write under a limit; the DB rows
(other pods included) plus config.yaml routers decide, the row being edited is excluded through the SQL
parameter, and every other write runs on the plain client with no lock."""
from fastapi import HTTPException
from litellm.proxy.management_endpoints.model_management_endpoints import (
HEURISTIC_V2_SLOT_LOCK_KEY,
_heuristic_v2_slot,
)
fake = self._FakeDb(db_held)
live_router = self._live_router_holding_one_heuristic_v2(limit) if config_holds_one else None
with (
patch("litellm.proxy.proxy_server._license_check.heuristic_v2_router_limit", lambda: limit), # test-quality-ok: the guard reads the proxy license singleton with no injection seam
patch("litellm.proxy.proxy_server.llm_router", live_router), # test-quality-ok: the guard reads the proxy router global with no injection seam
patch( # test-quality-ok: the cross-pod publish is the side effect under test; redis is not configured here
"litellm.proxy.management_endpoints.model_management_endpoints.publish_config_change",
new=AsyncMock(),
) as published,
):
if expected == "refused":
with pytest.raises(HTTPException) as exc_info:
async with _heuristic_v2_slot(fake, effective_config=effective_config, model_id=model_id):
pass
assert exc_info.value.status_code == 403
assert "At most 1 auto-router" in str(exc_info.value.detail)
assert "'auto_router' feature lifts the limit" in str(exc_info.value.detail)
return
async with _heuristic_v2_slot(fake, effective_config=effective_config, model_id=model_id) as tables:
handle = tables
if expected == "plain":
await handle.create(data={})
fake.litellm_proxymodeltable.create.assert_awaited_once_with(data={})
assert fake.tx_obj.raw_calls == []
return
assert handle is fake.tx_obj.litellm_proxymodeltable
published.assert_awaited_once_with(redis_cache=None, object_type="litellm_proxymodeltable")
(lock_sql, lock_params), (_count_sql, count_params) = fake.tx_obj.raw_calls
assert "pg_advisory_xact_lock($1)" in lock_sql and "count" not in lock_sql
assert lock_params == (HEURISTIC_V2_SLOT_LOCK_KEY,)
assert count_params == (model_id or "",)
@pytest.mark.asyncio
async def test_team_model_bookkeeping_runs_after_the_slot_is_released(self) -> None:
"""team_model_add needs a second pool connection, so it must run only after the slot transaction
(and its advisory lock) has closed; a pool-sized burst of team creates would otherwise stall on the
lock holder waiting for a connection the waiters are occupying."""
from contextlib import asynccontextmanager
from litellm.proxy.management_endpoints.model_management_endpoints import _add_team_model_to_db
from litellm.types.router import ModelInfo
events: list[str] = []
created = MagicMock(model_id="row-1")
@asynccontextmanager
async def slot():
events.append("slot-enter")
yield MagicMock(create=AsyncMock(return_value=created))
events.append("slot-exit")
async def team_model_add(**_: object) -> None:
events.append("team_model_add")
deployment = Deployment(
model_name="public-v2",
litellm_params=LiteLLM_Params(model="auto_router/complexity_router", complexity_router_config=self._V2),
model_info=ModelInfo(id="row-1", team_id="team-1"),
)
with (
patch( # test-quality-ok: params are encrypted with the proxy master key, which this test does not configure
"litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper",
lambda value, new_encryption_key=None: value,
),
patch( # test-quality-ok: the team list write is the collaborator whose ordering is asserted
"litellm.proxy.management_endpoints.model_management_endpoints.team_model_add",
side_effect=team_model_add,
),
):
result = await _add_team_model_to_db(
model_params=deployment,
user_api_key_dict=UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN),
prisma_client=MagicMock(),
slot=slot(),
)
assert result is created
assert events == ["slot-enter", "slot-exit", "team_model_add"]
@pytest.mark.asyncio
async def test_add_new_model_refuses_a_second_heuristic_v2_router_before_the_db_write(self) -> None:
from litellm.proxy._types import ProxyException
from litellm.proxy.management_endpoints.model_management_endpoints import (
add_new_model,
)
admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
fake = self._FakeDb(db_held=1)
with (
patch("litellm.proxy.proxy_server.prisma_client", fake), # test-quality-ok: endpoint reads proxy server globals with no injection seam
patch("litellm.proxy.proxy_server.llm_router", None), # test-quality-ok: endpoint reads proxy server globals with no injection seam
patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: endpoint reads proxy server globals with no injection seam
patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: endpoint reads proxy server globals with no injection seam
patch("litellm.proxy.proxy_server._license_check.heuristic_v2_router_limit", lambda: 1), # test-quality-ok: endpoint reads proxy server globals with no injection seam
patch( # test-quality-ok: prior auth check needs a live DB; only the license limit is under test
"litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call",
new=AsyncMock(return_value=None),
),
patch( # test-quality-ok: params are encrypted before the slot is entered; no master key in this test
"litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper",
lambda value, new_encryption_key=None: value,
),
):
with pytest.raises(ProxyException) as exc_info:
await add_new_model(
model_params=Deployment(
model_name="second-v2",
litellm_params=LiteLLM_Params(model="auto_router/complexity_router", complexity_router_config=self._V2),
),
user_api_key_dict=admin,
)
assert exc_info.value.code == "403"
assert "At most 1 auto-router" in str(exc_info.value.message)
fake.tx_obj.litellm_proxymodeltable.create.assert_not_awaited()
fake.litellm_proxymodeltable.create.assert_not_awaited()
@pytest.mark.asyncio
async def test_patch_model_refuses_switching_another_router_to_heuristic_v2(self) -> None:
"""patch_model relays HTTPException as-is, so the license refusal reaches the client as a plain 403."""
from fastapi import HTTPException
from litellm.proxy.management_endpoints.model_management_endpoints import (
patch_model,
)
from litellm.types.router import updateLiteLLMParams
model_id = "other-id"
admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
fake = self._FakeDb(db_held=1)
with (
patch("litellm.proxy.proxy_server.prisma_client", fake), # test-quality-ok: endpoint reads proxy server globals with no injection seam
patch("litellm.proxy.proxy_server.llm_router", None), # test-quality-ok: endpoint reads proxy server globals with no injection seam
patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: endpoint reads proxy server globals with no injection seam
patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: endpoint reads proxy server globals with no injection seam
patch("litellm.proxy.proxy_server._license_check.heuristic_v2_router_limit", lambda: 1), # test-quality-ok: endpoint reads proxy server globals with no injection seam
patch( # test-quality-ok: the write must be refused before this DB step runs
"litellm.proxy.management_endpoints.model_management_endpoints.get_db_model",
new=AsyncMock(return_value=self._db_complexity_router(model_id)),
),
patch( # test-quality-ok: prior auth check needs a live DB; only the license limit is under test
"litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call",
new=AsyncMock(return_value=None),
),
patch( # test-quality-ok: the helper's team bookkeeping needs a live DB; the row writer it is handed is what is under test
"litellm.proxy.management_endpoints.model_management_endpoints._update_team_model_in_db",
new=AsyncMock(side_effect=_write_empty_row),
),
):
with pytest.raises(HTTPException) as exc_info:
await patch_model(
model_id=model_id,
patch_data=updateDeployment(litellm_params=updateLiteLLMParams(complexity_router_config=self._V2)),
user_api_key_dict=admin,
)
assert exc_info.value.status_code == 403
fake.tx_obj.litellm_proxymodeltable.update.assert_not_awaited()
fake.litellm_proxymodeltable.update.assert_not_awaited()
@pytest.mark.asyncio
async def test_update_model_refuses_switching_another_router_to_heuristic_v2(self) -> None:
from litellm.proxy._types import ProxyException
from litellm.proxy.management_endpoints.model_management_endpoints import (
update_model,
)
from litellm.types.router import ModelInfo, updateLiteLLMParams
model_id = "other-id"
admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
existing_row = MagicMock()
existing_row.model_dump.return_value = {
"model_name": "my-auto-router",
"litellm_params": {
"model": "auto_router/complexity_router",
"complexity_router_config": {"tiers": {"SIMPLE": "gpt-4o-mini"}},
},
"model_info": {"id": model_id},
}
existing_row.litellm_params = existing_row.model_dump.return_value["litellm_params"]
fake = self._FakeDb(db_held=1, existing_row=existing_row)
with (
patch("litellm.proxy.proxy_server.prisma_client", fake), # test-quality-ok: endpoint reads proxy server globals with no injection seam
patch("litellm.proxy.proxy_server.llm_router", None), # test-quality-ok: endpoint reads proxy server globals with no injection seam
patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: endpoint reads proxy server globals with no injection seam
patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: endpoint reads proxy server globals with no injection seam
patch("litellm.proxy.proxy_server._license_check.heuristic_v2_router_limit", lambda: 1), # test-quality-ok: endpoint reads proxy server globals with no injection seam
patch( # test-quality-ok: prior auth check needs a live DB; only the license limit is under test
"litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call",
new=AsyncMock(return_value=None),
),
):
with pytest.raises(ProxyException) as exc_info:
await update_model(
model_params=updateDeployment(
litellm_params=updateLiteLLMParams(complexity_router_config=self._V2),
model_info=ModelInfo(id=model_id),
),
user_api_key_dict=admin,
)
assert exc_info.value.code == "403"
fake.tx_obj.litellm_proxymodeltable.update.assert_not_awaited()
fake.litellm_proxymodeltable.update.assert_not_awaited()
@pytest.mark.asyncio
async def test_update_model_rejects_prefix_strip(self):
from litellm.proxy._types import ProxyException

View file

@ -45,6 +45,10 @@ from litellm.types.router import (
from litellm.types.utils import Usage
async def _passthrough_row(update_data):
return update_data
def test_model_info_accepts_valid_ptu_fields():
info = ModelInfo(
id="x",
@ -385,6 +389,7 @@ class TestTeamModelUpdateValidatesBeforeWriting:
patch_data=patch_data,
user_api_key_dict=MagicMock(),
prisma_client=MagicMock(),
write_row=_passthrough_row,
)
return result, touched
@ -914,6 +919,7 @@ class TestPtuDeploymentsAreNotBilledPerToken:
patch_data=patch,
user_api_key_dict=UserAPIKeyAuth(user_id="a", user_role=LitellmUserRoles.PROXY_ADMIN),
prisma_client=MagicMock(),
write_row=_passthrough_row,
)
assert exc.value.status_code == 400

View file

@ -104,97 +104,6 @@ class TestPromptVersioning:
assert get_base_prompt_id(prompt_id="jack") == "jack"
assert get_base_prompt_id(prompt_id="my_prompt.v10") == "my_prompt"
def test_get_latest_version_prompt_id(self):
"""
Test that get_latest_version_prompt_id returns the highest version
"""
from litellm.proxy.prompts.prompt_endpoints import get_latest_version_prompt_id
# Mock prompt IDs dictionary
all_prompt_ids = {
"jack.v1": {},
"jack.v2": {},
"jack.v3": {},
"jane.v1": {},
"simple_prompt": {},
}
# Test with base prompt ID - should return latest version
assert (
get_latest_version_prompt_id(
prompt_id="jack", all_prompt_ids=all_prompt_ids
)
== "jack.v3"
)
# Test with versioned prompt ID - should still return latest version
assert (
get_latest_version_prompt_id(
prompt_id="jack.v1", all_prompt_ids=all_prompt_ids
)
== "jack.v3"
)
# Test with single version
assert (
get_latest_version_prompt_id(
prompt_id="jane", all_prompt_ids=all_prompt_ids
)
== "jane.v1"
)
# Test with non-versioned prompt
assert (
get_latest_version_prompt_id(
prompt_id="simple_prompt", all_prompt_ids=all_prompt_ids
)
== "simple_prompt"
)
# Test with non-existent prompt
assert (
get_latest_version_prompt_id(
prompt_id="nonexistent", all_prompt_ids=all_prompt_ids
)
== "nonexistent"
)
def test_construct_versioned_prompt_id(self):
"""
Test that construct_versioned_prompt_id correctly builds versioned IDs
"""
from litellm.proxy.prompts.prompt_endpoints import construct_versioned_prompt_id
# Test with base prompt ID and version
assert (
construct_versioned_prompt_id(prompt_id="jack_success", version=4)
== "jack_success.v4"
)
# Test with None version - should return base ID unchanged
assert (
construct_versioned_prompt_id(prompt_id="jack_success", version=None)
== "jack_success"
)
# Test with existing versioned ID - should replace version
assert (
construct_versioned_prompt_id(prompt_id="jack_success.v2", version=4)
== "jack_success.v4"
)
# Test with hyphenated prompt ID
assert (
construct_versioned_prompt_id(prompt_id="my-prompt", version=1)
== "my-prompt.v1"
)
# Test with double-digit version
assert (
construct_versioned_prompt_id(prompt_id="test_prompt", version=10)
== "test_prompt.v10"
)
class TestPromptVersionsEndpoint:
"""
@ -444,7 +353,7 @@ class TestAdminViewerReadAccess:
"litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY"
) as mock_registry,
):
mock_registry.get_prompt_by_id.return_value = PromptSpec(
mock_registry.resolve_prompt_spec.return_value = PromptSpec(
prompt_id="jack.v2",
litellm_params=PromptLiteLLMParams(
prompt_id="jack",
@ -453,10 +362,101 @@ class TestAdminViewerReadAccess:
),
prompt_info=PromptInfo(prompt_type="db"),
)
mock_registry.IN_MEMORY_PROMPTS = {"jack.v1": {}, "jack.v2": {}}
mock_registry.get_prompt_callback_by_id.return_value = None
mock_registry.get_prompt_callback_for_prompt.return_value = None
response = await get_prompt_info(prompt_id="jack", user_api_key_dict=viewer)
assert response.prompt_spec.prompt_id == "jack"
assert response.prompt_spec.version == 2
class TestConfigPromptInfoWithEnvironment:
"""
Regression: /prompts/{id}/info with an environment param must still resolve
config-file (in-memory) prompts on a DB-backed proxy instead of 400ing.
"""
def _registry_with_config_prompt(self):
from litellm.proxy.prompts.prompt_registry import InMemoryPromptRegistry
registry = InMemoryPromptRegistry()
registry.IN_MEMORY_PROMPTS["envgreet::development"] = PromptSpec(
prompt_id="envgreet",
litellm_params=PromptLiteLLMParams(
prompt_id="envgreet",
prompt_integration="dotprompt",
dotprompt_content="AHOY {{user_message}}",
),
prompt_info=PromptInfo(prompt_type="config"),
)
return registry
def _prisma_client_with_empty_prompt_table(self):
from unittest.mock import AsyncMock
mock_prisma = MagicMock()
mock_prisma.db.litellm_prompttable.find_many = AsyncMock(return_value=[])
return mock_prisma
@pytest.mark.asyncio
async def test_get_prompt_info_with_environment_falls_back_to_registry(self):
from unittest.mock import patch
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.prompts.prompt_endpoints import get_prompt_info
admin = UserAPIKeyAuth(
api_key="test_key", user_role=LitellmUserRoles.PROXY_ADMIN
)
with (
patch( # test-quality-ok: endpoint reads prisma_client and the registry from module globals at call time; no injection seam
"litellm.proxy.proxy_server.prisma_client",
self._prisma_client_with_empty_prompt_table(),
),
patch( # test-quality-ok: endpoint reads prisma_client and the registry from module globals at call time; no injection seam
"litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY",
self._registry_with_config_prompt(),
),
):
response = await get_prompt_info(
prompt_id="envgreet",
environment="development",
user_api_key_dict=admin,
)
assert response.prompt_spec.prompt_id == "envgreet"
assert response.prompt_spec.litellm_params.dotprompt_content == "AHOY {{user_message}}"
@pytest.mark.asyncio
async def test_get_prompt_info_with_wrong_environment_still_400s(self):
from unittest.mock import patch
from fastapi import HTTPException
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.prompts.prompt_endpoints import get_prompt_info
admin = UserAPIKeyAuth(
api_key="test_key", user_role=LitellmUserRoles.PROXY_ADMIN
)
with (
patch( # test-quality-ok: endpoint reads prisma_client and the registry from module globals at call time; no injection seam
"litellm.proxy.proxy_server.prisma_client",
self._prisma_client_with_empty_prompt_table(),
),
patch( # test-quality-ok: endpoint reads prisma_client and the registry from module globals at call time; no injection seam
"litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY",
self._registry_with_config_prompt(),
),
):
with pytest.raises(HTTPException) as exc_info:
await get_prompt_info(
prompt_id="envgreet",
environment="production",
user_api_key_dict=admin,
)
assert exc_info.value.status_code == 400
assert "environment production" in exc_info.value.detail

View file

@ -52,8 +52,6 @@ async def test_delete_prompt_success():
with patch(
"litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY"
) as mock_registry:
# User passes "test_prompt.v2"
# We simulate that get_prompt_by_id returns the prompt spec for v2
prompt_spec = PromptSpec(
prompt_id="test_prompt.v2",
litellm_params=PromptLiteLLMParams(
@ -61,7 +59,8 @@ async def test_delete_prompt_success():
),
prompt_info=PromptInfo(prompt_type="db"),
)
mock_registry.get_prompt_by_id.return_value = prompt_spec
mock_registry.resolve_prompt_spec.return_value = prompt_spec
mock_registry.has_config_prompt.return_value = False
# Patch the prisma client in the endpoint module
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client):
@ -79,7 +78,7 @@ async def test_delete_prompt_success():
# 2. Memory deletion should use base ID
mock_registry.delete_prompts_by_base_id.assert_called_once_with(
expected_base_id, environment=None
base_prompt_id=expected_base_id, environment=None
)
assert response == {
@ -108,31 +107,14 @@ async def test_delete_prompt_by_base_id_success():
with patch(
"litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY"
) as mock_registry:
# User passes "test_prompt" (base ID)
# 1. get_prompt_by_id("test_prompt") -> None (if it's not registered as base)
# 2. It calls get_latest_version_prompt_id -> returns "test_prompt.v3"
# 3. get_prompt_by_id("test_prompt.v3") -> returns Spec
# Setup mocks behavior
def get_prompt_side_effect(prompt_id):
if prompt_id == "test_prompt":
return None
if prompt_id == "test_prompt.v3":
return PromptSpec(
prompt_id="test_prompt.v3",
litellm_params=PromptLiteLLMParams(
prompt_id="test_prompt", prompt_integration="dotprompt"
),
prompt_info=PromptInfo(prompt_type="db"),
)
return None
mock_registry.get_prompt_by_id.side_effect = get_prompt_side_effect
mock_registry.IN_MEMORY_PROMPTS = {
"test_prompt.v1": {},
"test_prompt.v2": {},
"test_prompt.v3": {},
}
mock_registry.resolve_prompt_spec.return_value = PromptSpec(
prompt_id="test_prompt.v3",
litellm_params=PromptLiteLLMParams(
prompt_id="test_prompt", prompt_integration="dotprompt"
),
prompt_info=PromptInfo(prompt_type="db"),
)
mock_registry.has_config_prompt.return_value = False
# Patch the prisma client in the endpoint module
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client):
@ -150,7 +132,7 @@ async def test_delete_prompt_by_base_id_success():
# 2. Memory deletion should use base ID
mock_registry.delete_prompts_by_base_id.assert_called_once_with(
expected_base_id, environment=None
base_prompt_id=expected_base_id, environment=None
)
assert response == {
@ -169,11 +151,12 @@ async def test_delete_prompt_environment_scope_reaches_db_and_registry():
with patch( # test-quality-ok: stubs the collaborator so the test pins what the endpoint deletes
"litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY"
) as mock_registry:
mock_registry.get_prompt_by_id.return_value = PromptSpec(
mock_registry.resolve_prompt_spec.return_value = PromptSpec(
prompt_id="test_prompt.v2",
litellm_params=PromptLiteLLMParams(prompt_id="test_prompt", prompt_integration="dotprompt"),
prompt_info=PromptInfo(prompt_type="db"),
prompt_info=PromptInfo(prompt_type="db", environment="production"),
)
mock_registry.has_config_prompt.return_value = False
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client): # test-quality-ok: proxy_server module global is the endpoint's only injection point
response = await delete_prompt(
@ -185,7 +168,7 @@ async def test_delete_prompt_environment_scope_reaches_db_and_registry():
mock_prisma_client.db.litellm_prompttable.delete_many.assert_called_once_with(
where={"prompt_id": "test_prompt", "environment": "production"}
)
mock_registry.delete_prompts_by_base_id.assert_called_once_with("test_prompt", environment="production")
mock_registry.delete_prompts_by_base_id.assert_called_once_with(base_prompt_id="test_prompt", environment="production")
assert response == {"message": "Prompt test_prompt deleted successfully from production"}
@ -218,24 +201,8 @@ async def test_get_prompt_info_by_base_id():
prompt_info=PromptInfo(prompt_type="db"),
)
# When get_prompt_by_id is called with "test_prompt", return None (so it searches versions)
# When called with "test_prompt.v3", return the spec
def get_prompt_side_effect(prompt_id):
if prompt_id == "test_prompt":
return None
if prompt_id == "test_prompt.v3":
return prompt_spec_v3
return None
mock_registry.get_prompt_by_id.side_effect = get_prompt_side_effect
mock_registry.IN_MEMORY_PROMPTS = {
"test_prompt.v1": {},
"test_prompt.v2": {},
"test_prompt.v3": {},
}
# We also need to mock get_prompt_callback_by_id to avoid content extraction errors/logic
mock_registry.get_prompt_callback_by_id.return_value = None
mock_registry.resolve_prompt_spec.return_value = prompt_spec_v3
mock_registry.get_prompt_callback_for_prompt.return_value = None
response = await get_prompt_info(
prompt_id="test_prompt", user_api_key_dict=mock_user_auth
@ -284,7 +251,7 @@ async def test_patch_prompt_row_deleted_mid_update_returns_404():
"litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY"
) as mock_registry,
):
mock_registry.get_prompt_by_id.return_value = existing_prompt
mock_registry.has_config_prompt.return_value = False
with pytest.raises(HTTPException) as exc_info:
await patch_prompt(
@ -325,7 +292,7 @@ async def test_patch_prompt_merges_unsent_fields_from_db_row_not_stale_memory():
"litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY"
) as mock_registry,
):
mock_registry.get_prompt_by_id.return_value = stale_in_memory
mock_registry.has_config_prompt.return_value = False
mock_registry.reload_prompt.side_effect = lambda prompt: prompt
response = await patch_prompt(
@ -487,7 +454,7 @@ async def test_patch_prompt_info_only_keeps_legacy_keyed_row_patchable():
"litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY"
) as mock_registry,
):
mock_registry.get_prompt_by_id.return_value = existing_prompt
mock_registry.has_config_prompt.return_value = False
await patch_prompt(
prompt_id="agent-prompt",

View file

@ -191,11 +191,7 @@ async def test_update_prompt_stores_environment_and_created_by():
with patch(
"litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY"
) as mock_registry:
mock_registry.get_prompt_by_id.return_value = PromptSpec(
prompt_id="my_prompt.v1",
litellm_params=request.litellm_params,
prompt_info=PromptInfo(prompt_type="db"),
)
mock_registry.has_config_prompt.return_value = False
mock_registry.initialize_prompt.return_value = PromptSpec(
prompt_id="my_prompt.v2",
litellm_params=request.litellm_params,
@ -239,7 +235,8 @@ async def test_delete_prompt_scoped_to_environment():
prompt_info=PromptInfo(prompt_type="db"),
environment="staging",
)
mock_registry.get_prompt_by_id.return_value = prompt_spec
mock_registry.resolve_prompt_spec.return_value = prompt_spec
mock_registry.has_config_prompt.return_value = False
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client):
await delete_prompt(
@ -251,3 +248,6 @@ async def test_delete_prompt_scoped_to_environment():
mock_prisma_client.db.litellm_prompttable.delete_many.assert_called_once_with(
where={"prompt_id": "test_prompt", "environment": "staging"}
)
mock_registry.delete_prompts_by_base_id.assert_called_once_with(
base_prompt_id="test_prompt", environment="staging"
)

View file

@ -1,26 +1,35 @@
import pytest
import litellm
from litellm.proxy.prompts.prompt_registry import InMemoryPromptRegistry
from litellm.integrations.custom_prompt_management import CustomPromptManagement
from litellm.proxy.prompts.prompt_registry import InMemoryPromptRegistry, parse_prompt_version
from litellm.types.prompts.init_prompts import PromptInfo, PromptLiteLLMParams, PromptSpec
def _db_prompt_spec(content: str) -> PromptSpec:
def _db_prompt_spec(content: str, environment: str = "development", version: int = 1) -> PromptSpec:
return PromptSpec(
prompt_id="greeting.v1",
prompt_id=f"greeting.v{version}",
litellm_params=PromptLiteLLMParams(
prompt_id="greeting",
prompt_integration="dotprompt",
prompt_data={"content": content, "metadata": {}},
),
prompt_info=PromptInfo(prompt_type="db"),
version=version,
environment=environment,
)
def _served_content(registry: InMemoryPromptRegistry) -> str:
callback = registry.get_prompt_callback_by_id("greeting.v1")
def _resolved_callback(registry: InMemoryPromptRegistry, environment: str | None = None) -> CustomPromptManagement:
spec = registry.resolve_prompt_spec("greeting", environment=environment)
assert spec is not None
callback = registry.get_prompt_callback_for_prompt(prompt=spec)
assert callback is not None
return callback.prompt_manager.get_prompt("greeting").content
return callback
def _served_content(registry: InMemoryPromptRegistry, environment: str | None = None) -> str:
return _resolved_callback(registry, environment=environment).prompt_manager.get_prompt("greeting").content
@pytest.fixture
@ -32,32 +41,34 @@ def isolated_callbacks(monkeypatch: pytest.MonkeyPatch) -> list:
def test_sync_prompt_from_db_reloads_row_edited_elsewhere(isolated_callbacks: list) -> None:
registry = InMemoryPromptRegistry()
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with AHOY"))
stale_callback = registry.get_prompt_callback_by_id("greeting.v1")
stale_callback = _resolved_callback(registry)
assert _served_content(registry) == "begin every reply with AHOY"
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with HOWDY"))
assert _served_content(registry) == "begin every reply with HOWDY"
assert registry.get_prompt_by_id("greeting.v1").litellm_params.prompt_data["content"] == "begin every reply with HOWDY"
reloaded_spec = registry.resolve_prompt_spec("greeting", environment="development")
assert reloaded_spec is not None
assert reloaded_spec.litellm_params.prompt_data["content"] == "begin every reply with HOWDY"
assert stale_callback not in isolated_callbacks
assert isolated_callbacks == [registry.get_prompt_callback_by_id("greeting.v1")]
assert isolated_callbacks == [_resolved_callback(registry)]
def test_sync_prompt_from_db_keeps_unchanged_row_in_place(isolated_callbacks: list) -> None:
registry = InMemoryPromptRegistry()
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with AHOY"))
first_callback = registry.get_prompt_callback_by_id("greeting.v1")
first_callback = _resolved_callback(registry)
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with AHOY"))
assert registry.get_prompt_callback_by_id("greeting.v1") is first_callback
assert _resolved_callback(registry) is first_callback
assert isolated_callbacks == [first_callback]
def test_reload_prompt_replaces_callback_without_leaking_the_old_one(isolated_callbacks: list) -> None:
registry = InMemoryPromptRegistry()
registry.initialize_prompt(prompt=_db_prompt_spec("begin every reply with AHOY"))
stale_callback = registry.get_prompt_callback_by_id("greeting.v1")
stale_callback = _resolved_callback(registry)
reloaded = registry.reload_prompt(prompt=_db_prompt_spec("begin every reply with HOWDY"))
@ -70,7 +81,7 @@ def test_reload_prompt_replaces_callback_without_leaking_the_old_one(isolated_ca
def test_reload_prompt_keeps_the_old_template_when_the_replacement_fails(isolated_callbacks: list) -> None:
registry = InMemoryPromptRegistry()
registry.initialize_prompt(prompt=_db_prompt_spec("begin every reply with AHOY"))
old_callback = registry.get_prompt_callback_by_id("greeting.v1")
old_callback = _resolved_callback(registry)
broken = PromptSpec(
prompt_id="greeting.v1",
@ -80,63 +91,134 @@ def test_reload_prompt_keeps_the_old_template_when_the_replacement_fails(isolate
prompt_data={"content": "begin every reply with HOWDY", "metadata": {}},
),
prompt_info=PromptInfo(prompt_type="db"),
version=1,
environment="development",
)
with pytest.raises(ValueError, match="Unsupported prompt"):
registry.reload_prompt(prompt=broken)
assert registry.get_prompt_callback_by_id("greeting.v1") is old_callback
assert _resolved_callback(registry) is old_callback
assert _served_content(registry) == "begin every reply with AHOY"
assert isolated_callbacks == [old_callback]
def _versioned_prompt_spec(version: int, environment: str) -> PromptSpec:
return PromptSpec(
prompt_id=f"greeting.v{version}",
def test_environments_sharing_a_prompt_id_keep_separate_templates(isolated_callbacks: list) -> None:
registry = InMemoryPromptRegistry()
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with AHOY", environment="development"))
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with HOWDY", environment="production"))
assert _served_content(registry, environment="development") == "begin every reply with AHOY"
assert _served_content(registry, environment="production") == "begin every reply with HOWDY"
assert _resolved_callback(registry, environment="development") is not _resolved_callback(
registry, environment="production"
)
def test_default_resolution_prefers_production(isolated_callbacks: list) -> None:
registry = InMemoryPromptRegistry()
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with AHOY", environment="development"))
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with HOWDY", environment="production"))
assert _served_content(registry) == "begin every reply with HOWDY"
@pytest.mark.parametrize("environment", ["staging", "qa"])
def test_default_resolution_serves_the_only_environment_present(isolated_callbacks: list, environment: str) -> None:
registry = InMemoryPromptRegistry()
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with AHOY", environment=environment))
assert _served_content(registry) == "begin every reply with AHOY"
def test_resolution_picks_exact_version_and_latest_within_an_environment(isolated_callbacks: list) -> None:
registry = InMemoryPromptRegistry()
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with AHOY", environment="development", version=1))
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with YO", environment="development", version=2))
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with HOWDY", environment="production", version=1))
exact = registry.resolve_prompt_spec("greeting", version=1, environment="development")
assert exact is not None
assert exact.litellm_params.prompt_data["content"] == "begin every reply with AHOY"
latest = registry.resolve_prompt_spec("greeting", environment="development")
assert latest is not None
assert latest.litellm_params.prompt_data["content"] == "begin every reply with YO"
assert registry.resolve_prompt_spec("greeting", version=3, environment="development") is None
def test_resolution_returns_none_for_unknown_environment_or_prompt(isolated_callbacks: list) -> None:
registry = InMemoryPromptRegistry()
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with AHOY", environment="development"))
assert registry.resolve_prompt_spec("greeting", environment="production") is None
assert registry.resolve_prompt_spec("no_such_prompt") is None
def test_delete_prompts_by_base_id_scoped_to_one_environment(isolated_callbacks: list) -> None:
registry = InMemoryPromptRegistry()
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with AHOY", environment="development"))
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with HOWDY", environment="production"))
production_callback = _resolved_callback(registry, environment="production")
deleted = registry.delete_prompts_by_base_id(base_prompt_id="greeting", environment="development")
assert deleted == ["greeting.v1::development"]
assert registry.resolve_prompt_spec("greeting", environment="development") is None
assert _resolved_callback(registry, environment="production") is production_callback
assert _served_content(registry, environment="production") == "begin every reply with HOWDY"
deleted_rest = registry.delete_prompts_by_base_id(base_prompt_id="greeting")
assert deleted_rest == ["greeting.v1::production"]
assert registry.resolve_prompt_spec("greeting") is None
def test_has_config_prompt_matches_any_version_of_the_base_id(isolated_callbacks: list) -> None:
registry = InMemoryPromptRegistry()
config_spec = PromptSpec(
prompt_id="greeting",
litellm_params=PromptLiteLLMParams(
prompt_id="greeting",
prompt_integration="dotprompt",
prompt_data={"content": f"begin every reply with AHOY v{version}", "metadata": {}},
prompt_data={"content": "begin every reply with AHOY", "metadata": {}},
),
prompt_info=PromptInfo(prompt_type="db", environment=environment),
version=version,
environment=environment,
prompt_info=PromptInfo(prompt_type="config"),
)
registry.initialize_prompt(prompt=config_spec)
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with HOWDY", environment="production"))
assert registry.has_config_prompt(base_prompt_id="greeting") is True
assert registry.has_config_prompt(base_prompt_id="other_prompt") is False
def test_delete_prompts_by_base_id_removes_the_callbacks_from_litellm_callbacks(isolated_callbacks: list) -> None:
registry = InMemoryPromptRegistry()
registry.initialize_prompt(prompt=_versioned_prompt_spec(1, "development"))
registry.initialize_prompt(prompt=_versioned_prompt_spec(2, "development"))
registry.initialize_prompt(prompt=_db_prompt_spec("begin every reply with AHOY", version=1))
registry.initialize_prompt(prompt=_db_prompt_spec("begin every reply with YO", version=2))
assert len(isolated_callbacks) == 1
deleted = registry.delete_prompts_by_base_id("greeting")
deleted = registry.delete_prompts_by_base_id(base_prompt_id="greeting")
assert sorted(deleted) == ["greeting.v1", "greeting.v2"]
assert registry.get_prompt_by_id("greeting.v1") is None
assert registry.get_prompt_callback_by_id("greeting.v2") is None
assert sorted(deleted) == ["greeting.v1::development", "greeting.v2::development"]
assert registry.resolve_prompt_spec("greeting") is None
assert isolated_callbacks == []
def test_delete_prompts_by_base_id_environment_scope_keeps_other_environments(isolated_callbacks: list) -> None:
def test_remove_prompt_is_a_no_op_for_an_unknown_registry_key(isolated_callbacks: list) -> None:
registry = InMemoryPromptRegistry()
registry.initialize_prompt(prompt=_versioned_prompt_spec(1, "development"))
registry.initialize_prompt(prompt=_versioned_prompt_spec(2, "production"))
production_callback = registry.get_prompt_callback_by_id("greeting.v2")
registry.initialize_prompt(prompt=_db_prompt_spec("begin every reply with AHOY"))
deleted = registry.delete_prompts_by_base_id("greeting", environment="development")
registry.remove_prompt(registry_key="not_there.v1::development")
assert deleted == ["greeting.v1"]
assert registry.get_prompt_by_id("greeting.v1") is None
assert registry.get_prompt_by_id("greeting.v2") is not None
assert registry.get_prompt_callback_by_id("greeting.v2") is production_callback
def test_remove_prompt_is_a_no_op_for_an_unknown_id(isolated_callbacks: list) -> None:
registry = InMemoryPromptRegistry()
registry.initialize_prompt(prompt=_versioned_prompt_spec(1, "development"))
registry.remove_prompt(prompt_id="not_there.v1")
assert registry.get_prompt_by_id("greeting.v1") is not None
assert registry.resolve_prompt_spec("greeting") is not None
assert len(isolated_callbacks) == 1
@pytest.mark.parametrize(
("raw_version", "expected"),
[(2, 2), ("2", 2), (None, None), ("v2", None), (True, None), (2.0, None)],
)
def test_parse_prompt_version_accepts_integers_and_json_strings(raw_version: object, expected: int | None) -> None:
assert parse_prompt_version(raw_version) == expected

View file

@ -28,6 +28,7 @@ from litellm.proxy.proxy_server import (
resolve_routing_plugins,
validate_deployment_complexity_router_placement,
validate_deployment_max_agentic_loops,
validate_heuristic_v2_router_limit,
)
from .conftest import normalize
@ -193,6 +194,120 @@ def test_validate_deployment_complexity_router_placement_leaves_valid_deployment
assert model["litellm_params"] == litellm_params
def _heuristic_v2_row(model_name: str, classifier_type: str = "heuristic_v2") -> dict[str, object]:
return {
"model_name": model_name,
"litellm_params": {
"model": "auto_router/complexity_router",
"complexity_router_config": {"classifier_type": classifier_type, "tiers": {"SIMPLE": "gpt-4o-mini"}},
},
}
def test_validate_heuristic_v2_router_limit_refuses_to_start_over_the_limit() -> None:
"""Same reason as the two validators above: the proxy router swallows registration errors, so
an over-limit config.yaml must fail here instead of booting with a silently missing router."""
with pytest.raises(ValueError, match=re.escape("At most 1 auto-router")) as exc_info:
validate_heuristic_v2_router_limit(
[_heuristic_v2_row("a"), _heuristic_v2_row("b"), _heuristic_v2_row("c", "heuristic")], limit=1
)
assert "'auto_router' feature lifts the limit" in str(exc_info.value)
@pytest.mark.parametrize(
"model_list,limit",
[
([_heuristic_v2_row("a"), _heuristic_v2_row("b")], None),
([_heuristic_v2_row("a"), _heuristic_v2_row("c", "heuristic")], 1),
([{"model_name": "gpt-4o", "litellm_params": {"model": "gpt-4o"}}], 1),
],
)
def test_validate_heuristic_v2_router_limit_leaves_configs_within_the_limit_alone(
model_list: list[dict[str, object]], limit: int | None
) -> None:
assert validate_heuristic_v2_router_limit(model_list, limit=limit) is None
_TWO_HEURISTIC_V2_ROUTERS_YAML = (
"model_list:\n"
" - model_name: gpt-4o-mini\n"
" litellm_params:\n"
" model: openai/gpt-4o-mini\n"
" api_key: k\n"
" - model_name: v2-a\n"
" litellm_params:\n"
" model: auto_router/complexity_router\n"
" complexity_router_config:\n"
" classifier_type: heuristic_v2\n"
" tiers: {SIMPLE: gpt-4o-mini}\n"
" - model_name: v2-b\n"
" litellm_params:\n"
" model: auto_router/complexity_router\n"
" complexity_router_config:\n"
" classifier_type: heuristic_v2\n"
" tiers: {SIMPLE: gpt-4o-mini}\n"
"router_settings:\n"
" heuristic_v2_router_limit: 99\n"
)
@pytest.mark.asyncio
@pytest.mark.parametrize("license_limit", [1, None])
async def test_ProxyConfig_load_config_takes_the_heuristic_v2_limit_from_the_license_only(
tmp_path, monkeypatch, license_limit: int | None
) -> None:
"""`router_settings.heuristic_v2_router_limit` is managed outside config.yaml: an operator
cannot grant the entitlement by editing the config, and a licensed proxy boots both routers."""
f = tmp_path / "c.yaml"
f.write_text(_TWO_HEURISTIC_V2_ROUTERS_YAML)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False)
monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False)
monkeypatch.setattr(
"litellm.proxy.proxy_server._license_check.heuristic_v2_router_limit", lambda: license_limit
)
if license_limit is None:
router, _model_list, _general_settings = await ProxyConfig().load_config(
router=None, config_file_path=str(f)
)
assert router.heuristic_v2_router_limit is not None
assert router.heuristic_v2_router_limit() is None
assert sorted(router.complexity_routers) == ["v2-a", "v2-b"]
return
with pytest.raises(ValueError, match=re.escape("config.yaml model_list: At most 1 auto-router")):
await ProxyConfig().load_config(router=None, config_file_path=str(f))
@pytest.mark.asyncio
async def test_ProxyConfig_load_config_router_refuses_a_db_heuristic_v2_router_beyond_the_license(
tmp_path, monkeypatch
) -> None:
"""config.yaml holds the one allowed heuristic_v2 router; a second one arriving later from the DB
is refused at registration because the router was built with the license's ceiling."""
from litellm.types.router import Deployment
f = tmp_path / "c.yaml"
f.write_text(_TWO_HEURISTIC_V2_ROUTERS_YAML.replace(" - model_name: v2-b\n", " - model_name: v1-b\n", 1).replace(
"classifier_type: heuristic_v2\n tiers: {SIMPLE: gpt-4o-mini}\nrouter_settings",
"classifier_type: heuristic\n tiers: {SIMPLE: gpt-4o-mini}\nrouter_settings",
))
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False)
monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False)
monkeypatch.setattr("litellm.proxy.proxy_server._license_check.heuristic_v2_router_limit", lambda: 1)
router, _model_list, _general_settings = await ProxyConfig().load_config(router=None, config_file_path=str(f))
assert router.heuristic_v2_router_limit is not None
assert router.heuristic_v2_router_limit() == 1
assert sorted(router.complexity_routers) == ["v1-b", "v2-a"]
db_row = Deployment(**_heuristic_v2_row("v2-from-db"), model_info={"id": "db-id"})
assert router.upsert_deployment(db_row) is None
assert sorted(router.complexity_routers) == ["v1-b", "v2-a"]
def test_validate_deployment_max_agentic_loops_allows_a_deployment_without_the_key():
model = {"model_name": "gpt-4o", "litellm_params": {"model": "gpt-4o"}}

View file

@ -12065,10 +12065,15 @@ async def test_init_prompts_in_db_reloads_rows_patched_on_another_worker(monkeyp
}
return row
def served_content() -> str:
callback = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_sync.v1")
def served_callback():
spec = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec("greeting_sync")
assert spec is not None
callback = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_for_prompt(prompt=spec)
assert callback is not None
return callback.prompt_manager.get_prompt("greeting_sync").content
return callback
def served_content() -> str:
return served_callback().prompt_manager.get_prompt("greeting_sync").content
prisma_client = MagicMock()
try:
@ -12080,7 +12085,7 @@ async def test_init_prompts_in_db_reloads_rows_patched_on_another_worker(monkeyp
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
assert served_content() == "Begin every reply with HOWDY"
assert litellm.callbacks == [IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_sync.v1")]
assert litellm.callbacks == [served_callback()]
finally:
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id("greeting_sync")
@ -12119,22 +12124,25 @@ async def test_init_prompts_in_db_syncs_remaining_rows_when_one_row_fails(monkey
)
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id("broken_sync.v1") is None
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("healthy_sync.v1") is not None
assert litellm.callbacks == [IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("healthy_sync.v1")]
assert IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec("broken_sync") is None
healthy_spec = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec("healthy_sync")
assert healthy_spec is not None
healthy_callback = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_for_prompt(prompt=healthy_spec)
assert healthy_callback is not None
assert litellm.callbacks == [healthy_callback]
finally:
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id("healthy_sync")
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id("broken_sync")
@pytest.mark.asyncio
async def test_init_prompts_in_db_serves_the_newest_row_when_environments_collide_on_a_versioned_id(monkeypatch):
async def test_init_prompts_in_db_syncs_every_environment_sharing_a_versioned_id(monkeypatch):
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
from litellm.proxy.proxy_server import ProxyConfig
monkeypatch.setattr(litellm, "callbacks", [])
def db_row(environment: str, content: str, updated_at: datetime) -> MagicMock:
def db_row(environment: str, content: str) -> MagicMock:
row = MagicMock()
row.model_dump.return_value = {
"prompt_id": "greeting_env",
@ -12150,30 +12158,41 @@ async def test_init_prompts_in_db_serves_the_newest_row_when_environments_collid
),
"prompt_info": json.dumps({"prompt_type": "db"}),
"created_at": None,
"updated_at": updated_at,
"updated_at": None,
}
return row
freshly_patched = db_row(
"production", "Begin every reply with HOWDY", datetime(2026, 8, 26, 12, 0, tzinfo=timezone.utc)
)
stale_sibling = db_row(
"development", "Begin every reply with AHOY", datetime(2026, 8, 26, 11, 0, tzinfo=timezone.utc)
)
def served_content(environment: str | None) -> str:
spec = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec("greeting_env", environment=environment)
assert spec is not None
callback = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_for_prompt(prompt=spec)
assert callback is not None
return callback.prompt_manager.get_prompt("greeting_env").content
prisma_client = MagicMock()
try:
prisma_client.db.litellm_prompttable.find_many = AsyncMock(return_value=[freshly_patched, stale_sibling])
prisma_client.db.litellm_prompttable.find_many = AsyncMock(
return_value=[
db_row("development", "Begin every reply with AHOY"),
db_row("production", "Begin every reply with HOWDY"),
]
)
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
first_callback = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_env.v1")
assert first_callback is not None
assert first_callback.prompt_manager.get_prompt("greeting_env").content == "Begin every reply with HOWDY"
assert served_content("development") == "Begin every reply with AHOY"
assert served_content("production") == "Begin every reply with HOWDY"
assert served_content(None) == "Begin every reply with HOWDY"
prisma_client.db.litellm_prompttable.find_many = AsyncMock(
return_value=[
db_row("development", "Begin every reply with YO"),
db_row("production", "Begin every reply with HOWDY"),
]
)
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_env.v1") is first_callback
assert litellm.callbacks == [first_callback]
assert served_content("development") == "Begin every reply with YO"
assert served_content("production") == "Begin every reply with HOWDY"
finally:
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id("greeting_env")
@ -12216,13 +12235,14 @@ async def test_init_prompts_in_db_unloads_rows_deleted_on_another_worker(monkeyp
return_value=[_prompt_db_row("greeting_del", _dotprompt_params("greeting_del"))]
)
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_del.v1") is not None
loaded_spec = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec("greeting_del")
assert loaded_spec is not None
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_for_prompt(prompt=loaded_spec) is not None
prisma_client.db.litellm_prompttable.find_many = AsyncMock(return_value=[])
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id("greeting_del.v1") is None
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_del.v1") is None
assert IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec("greeting_del") is None
assert litellm.callbacks == []
finally:
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id("greeting_del")
@ -12253,10 +12273,12 @@ async def test_init_prompts_in_db_keeps_config_prompts_when_their_id_has_no_db_r
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_cfg") is not None
surviving_spec = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec("greeting_cfg")
assert surviving_spec is not None
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_for_prompt(prompt=surviving_spec) is not None
assert len(litellm.callbacks) == 1
finally:
IN_MEMORY_PROMPT_REGISTRY.remove_prompt(prompt_id="greeting_cfg")
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id("greeting_cfg")
@pytest.mark.asyncio
@ -12272,7 +12294,9 @@ async def test_init_prompts_in_db_keeps_the_in_memory_copy_when_a_row_fails_to_p
return_value=[_prompt_db_row("greeting_broken", _dotprompt_params("greeting_broken"))]
)
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
loaded_callback = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_broken.v1")
loaded_spec = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec("greeting_broken")
assert loaded_spec is not None
loaded_callback = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_for_prompt(prompt=loaded_spec)
assert loaded_callback is not None
prisma_client.db.litellm_prompttable.find_many = AsyncMock(
@ -12280,7 +12304,9 @@ async def test_init_prompts_in_db_keeps_the_in_memory_copy_when_a_row_fails_to_p
)
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_broken.v1") is loaded_callback
kept_spec = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec("greeting_broken")
assert kept_spec is not None
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_for_prompt(prompt=kept_spec) is loaded_callback
assert litellm.callbacks == [loaded_callback]
finally:
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id("greeting_broken")
@ -12314,8 +12340,9 @@ async def test_init_prompts_in_db_keeps_a_prompt_created_while_the_sync_was_read
prisma_client.db.litellm_prompttable.find_many = AsyncMock(side_effect=create_prompt_behind_the_select)
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
surviving_callback = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_race.v1")
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id("greeting_race.v1") is not None
surviving_spec = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec("greeting_race")
assert surviving_spec is not None
surviving_callback = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_for_prompt(prompt=surviving_spec)
assert surviving_callback is not None
assert litellm.callbacks == [surviving_callback]
finally:

View file

@ -726,10 +726,7 @@ async def test_process_prompt_template_no_op_when_no_prompt_spec(proxy_logging,
from litellm.proxy.prompts import prompt_registry
monkeypatch.setattr(
prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "get_prompt_callback_by_id", lambda *a, **kw: None
)
monkeypatch.setattr(
prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "get_prompt_by_id", lambda *a, **kw: None
prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "resolve_prompt_spec", lambda *a, **kw: None
)
data: Dict[str, Any] = {"messages": [{"role": "user"}], "model": "m", "temperature": 0.1}
await proxy_logging._process_prompt_template(
@ -752,11 +749,11 @@ async def test_process_prompt_template_applies_when_spec_resolves(proxy_logging,
monkeypatch.setattr(
prompt_registry.IN_MEMORY_PROMPT_REGISTRY,
"get_prompt_callback_by_id",
"get_prompt_callback_for_prompt",
lambda *a, **kw: custom_logger,
)
monkeypatch.setattr(
prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "get_prompt_by_id", lambda *a, **kw: prompt_spec
prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "resolve_prompt_spec", lambda *a, **kw: prompt_spec
)
logging_obj = MagicMock()
@ -802,11 +799,11 @@ async def test_process_prompt_template_async_get_prompt_error_raises(proxy_loggi
prompt_spec.litellm_params = MagicMock(prompt_id="x")
monkeypatch.setattr(
prompt_registry.IN_MEMORY_PROMPT_REGISTRY,
"get_prompt_callback_by_id",
"get_prompt_callback_for_prompt",
lambda *a, **kw: custom_logger,
)
monkeypatch.setattr(
prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "get_prompt_by_id", lambda *a, **kw: prompt_spec
prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "resolve_prompt_spec", lambda *a, **kw: prompt_spec
)
logging_obj = MagicMock()
logging_obj.async_get_chat_completion_prompt = AsyncMock(side_effect=RuntimeError("bad prompt"))
@ -820,6 +817,82 @@ async def test_process_prompt_template_async_get_prompt_error_raises(proxy_loggi
)
@pytest.mark.asyncio
async def test_process_prompt_template_resolves_the_requested_environment(proxy_logging, monkeypatch):
from litellm.proxy.prompts import prompt_registry
prompt_spec = MagicMock()
prompt_spec.litellm_params = MagicMock(prompt_id="greeting")
resolve_calls: list[dict] = []
def fake_resolve(prompt_id, version=None, environment=None):
resolve_calls.append({"prompt_id": prompt_id, "version": version, "environment": environment})
return prompt_spec
monkeypatch.setattr(prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "resolve_prompt_spec", fake_resolve)
monkeypatch.setattr(
prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "get_prompt_callback_for_prompt", lambda *a, **kw: MagicMock()
)
logging_obj = MagicMock()
logging_obj.async_get_chat_completion_prompt = AsyncMock(
return_value=("m", [{"role": "user", "content": "rendered"}], {})
)
data: Dict[str, Any] = {
"messages": [{"role": "user", "content": "orig"}],
"model": "m",
"prompt_id": "greeting",
"prompt_version": 1,
"prompt_environment": "development",
}
await proxy_logging._process_prompt_template(
data=data,
litellm_logging_obj=logging_obj,
prompt_id="greeting",
prompt_version=1,
call_type="completion",
)
assert resolve_calls == [{"prompt_id": "greeting", "version": 1, "environment": "development"}]
assert "prompt_environment" not in data
assert "prompt_id" not in data
assert data["messages"] == [{"role": "user", "content": "rendered"}]
@pytest.mark.asyncio
async def test_pre_call_hook_matches_a_prompt_version_sent_as_a_json_string(proxy_logging, monkeypatch):
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.prompts import prompt_registry
prompt_spec = MagicMock()
prompt_spec.litellm_params = MagicMock(prompt_id="greeting")
resolve_calls: list[dict] = []
def fake_resolve(prompt_id, version=None, environment=None):
resolve_calls.append({"prompt_id": prompt_id, "version": version, "environment": environment})
return prompt_spec
monkeypatch.setattr(prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "resolve_prompt_spec", fake_resolve)
monkeypatch.setattr(
prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "get_prompt_callback_for_prompt", lambda *a, **kw: MagicMock()
)
logging_obj = MagicMock()
logging_obj.async_get_chat_completion_prompt = AsyncMock(
return_value=("m", [{"role": "user", "content": "rendered"}], {})
)
data: Dict[str, Any] = {
"messages": [{"role": "user", "content": "orig"}],
"model": "m",
"prompt_id": "greeting",
"prompt_version": "2",
"litellm_logging_obj": logging_obj,
}
result = await proxy_logging.pre_call_hook(user_api_key_dict=UserAPIKeyAuth(), data=data, call_type="completion")
assert resolve_calls == [{"prompt_id": "greeting", "version": 2, "environment": None}]
assert result["messages"] == [{"role": "user", "content": "rendered"}]
@pytest.mark.asyncio
async def test_process_prompt_template_aresponses_swaps_model_and_merges_input(proxy_logging, monkeypatch):
from litellm.proxy.prompts import prompt_registry
@ -829,11 +902,11 @@ async def test_process_prompt_template_aresponses_swaps_model_and_merges_input(p
prompt_spec.litellm_params = MagicMock(prompt_id="resolved-id")
monkeypatch.setattr(
prompt_registry.IN_MEMORY_PROMPT_REGISTRY,
"get_prompt_callback_by_id",
"get_prompt_callback_for_prompt",
lambda *a, **kw: custom_logger,
)
monkeypatch.setattr(
prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "get_prompt_by_id", lambda *a, **kw: prompt_spec
prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "resolve_prompt_spec", lambda *a, **kw: prompt_spec
)
logging_obj = MagicMock()

View file

@ -330,36 +330,46 @@ ROUTER_CONFIG: Final = json.dumps(
)
class FailingRouteLayer:
"""Route layer whose embedding call fails, as it does when the prompt exceeds the encoder's window."""
def __call__(self, text: str) -> Any:
raise ValueError(
"Internal_litellm_router API call failed. Error: litellm.InternalServerError: "
"input is too large to process. increase the physical batch size"
)
class FixedRouteLayer:
"""Route layer that returns whatever the test tells it to, recording the text it was asked about."""
"""Route layer that returns whatever the test tells it to for the query vector it is handed."""
def __init__(self, route_choice: Any) -> None:
self.route_choice = route_choice
self.seen_text: str | None = None
def __call__(self, text: str) -> Any:
self.seen_text = text
async def acall(self, vector: Any) -> Any:
return self.route_choice
def _embedding_response(input: List[str]) -> Any:
import litellm
return litellm.EmbeddingResponse(
data=[{"embedding": [0.1, 0.2], "index": i, "object": "embedding"} for i in range(len(input))]
)
class StubEmbeddingRouter:
"""Stands in for the LiteLLM Router when the route index has to be built for real."""
"""Stands in for the LiteLLM Router, recording the text and kwargs each query embedding was made with."""
def __init__(self) -> None:
self.seen_text: str | None = None
self.aembedding_kwargs: Dict[str, Any] | None = None
def embedding(self, input: List[str], model: str, **kwargs: Any) -> Any:
import litellm
return _embedding_response(input)
return litellm.EmbeddingResponse(
data=[{"embedding": [0.1, 0.2], "index": i, "object": "embedding"} for i in range(len(input))]
async def aembedding(self, input: List[str], model: str, **kwargs: Any) -> Any:
self.seen_text = input[0]
self.aembedding_kwargs = kwargs
return _embedding_response(input)
class FailingEmbeddingRouter(StubEmbeddingRouter):
"""Router whose query embedding fails, as it does when the prompt exceeds the encoder's window."""
async def aembedding(self, input: List[str], model: str, **kwargs: Any) -> Any:
raise ValueError(
"litellm.InternalServerError: input is too large to process. increase the physical batch size"
)
@ -369,7 +379,7 @@ def _auto_router(routelayer: Any, litellm_router_instance: Any = None, **kwargs:
auto_router_config=ROUTER_CONFIG,
default_model="fallback-model",
embedding_model="text-embedding-3-small",
litellm_router_instance=litellm_router_instance or MagicMock(),
litellm_router_instance=litellm_router_instance or StubEmbeddingRouter(),
**kwargs,
)
auto_router.routelayer = routelayer
@ -381,7 +391,7 @@ class TestAutoRouterAlwaysResolvesARoutableModel:
@pytest.mark.asyncio
async def test_should_fall_back_to_default_model_when_the_embedding_call_fails(self):
auto_router: Final = _auto_router(FailingRouteLayer())
auto_router: Final = _auto_router(FixedRouteLayer(None), litellm_router_instance=FailingEmbeddingRouter())
result: Final = await auto_router.async_pre_routing_hook(
model="my-auto-router",
@ -440,8 +450,8 @@ class TestAutoRouterAlwaysResolvesARoutableModel:
async def test_should_still_route_to_the_matched_route_when_one_matches(self):
from semantic_router.schema import RouteChoice
layer: Final = FixedRouteLayer(RouteChoice(name="code-model"))
auto_router: Final = _auto_router(layer)
router: Final = StubEmbeddingRouter()
auto_router: Final = _auto_router(FixedRouteLayer(RouteChoice(name="code-model")), litellm_router_instance=router)
result: Final = await auto_router.async_pre_routing_hook(
model="my-auto-router",
@ -451,7 +461,7 @@ class TestAutoRouterAlwaysResolvesARoutableModel:
assert result is not None
assert result.model == "code-model"
assert layer.seen_text == "fix this stack trace"
assert router.seen_text == "fix this stack trace"
class TestAutoRouterEmbeddingInputCap:
@ -483,8 +493,8 @@ class TestAutoRouterRoutesResponsesApiInput:
async def test_should_route_a_string_input_when_messages_is_none(self):
from semantic_router.schema import RouteChoice
layer: Final = FixedRouteLayer(RouteChoice(name="code-model"))
auto_router: Final = _auto_router(layer)
router: Final = StubEmbeddingRouter()
auto_router: Final = _auto_router(FixedRouteLayer(RouteChoice(name="code-model")), litellm_router_instance=router)
result: Final = await auto_router.async_pre_routing_hook(
model="my-auto-router",
@ -498,14 +508,14 @@ class TestAutoRouterRoutesResponsesApiInput:
assert result is not None
assert result.model == "code-model"
assert result.messages is None
assert layer.seen_text == "fix this stack trace"
assert router.seen_text == "fix this stack trace"
@pytest.mark.asyncio
async def test_should_route_a_list_input_with_instructions_when_messages_is_none(self):
from semantic_router.schema import RouteChoice
layer: Final = FixedRouteLayer(RouteChoice(name="code-model"))
auto_router: Final = _auto_router(layer)
router: Final = StubEmbeddingRouter()
auto_router: Final = _auto_router(FixedRouteLayer(RouteChoice(name="code-model")), litellm_router_instance=router)
result: Final = await auto_router.async_pre_routing_hook(
model="my-auto-router",
@ -525,13 +535,13 @@ class TestAutoRouterRoutesResponsesApiInput:
assert result is not None
assert result.model == "code-model"
assert layer.seen_text is not None
assert "fix this stack trace" in layer.seen_text
assert router.seen_text is not None
assert "fix this stack trace" in router.seen_text
@pytest.mark.asyncio
async def test_should_skip_routing_when_neither_messages_nor_input_is_present(self):
layer: Final = FixedRouteLayer(None)
auto_router: Final = _auto_router(layer)
router: Final = StubEmbeddingRouter()
auto_router: Final = _auto_router(FixedRouteLayer(None), litellm_router_instance=router)
result: Final = await auto_router.async_pre_routing_hook(
model="my-auto-router",
@ -540,12 +550,12 @@ class TestAutoRouterRoutesResponsesApiInput:
)
assert result is None
assert layer.seen_text is None
assert router.seen_text is None
@pytest.mark.asyncio
async def test_should_keep_routing_an_empty_messages_list_to_the_default_model(self):
layer: Final = FixedRouteLayer(None)
auto_router: Final = _auto_router(layer)
router: Final = StubEmbeddingRouter()
auto_router: Final = _auto_router(FixedRouteLayer(None), litellm_router_instance=router)
result: Final = await auto_router.async_pre_routing_hook(
model="my-auto-router",
@ -555,4 +565,42 @@ class TestAutoRouterRoutesResponsesApiInput:
assert result is not None
assert result.model == "fallback-model"
assert layer.seen_text == ""
assert router.seen_text == ""
class TestAutoRouterAttributesItsEmbeddingSpend:
"""The query embedding is billed to the key that sent the request, like any other call it made."""
@pytest.mark.asyncio
async def test_should_forward_the_callers_identity_to_the_query_embedding_minus_its_budget_reservation(self):
from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY
router: Final = StubEmbeddingRouter()
auto_router: Final = _auto_router(None, litellm_router_instance=router)
request_kwargs: Final = {
"metadata": {
"user_api_key": "hashed-key",
"user_api_key_team_id": "team-1",
"user_api_key_budget_reservation": {"reservation_id": "r-1"},
},
"litellm_session_id": "session-1",
}
result: Final = await auto_router.async_pre_routing_hook(
model="my-auto-router",
request_kwargs=request_kwargs,
messages=[{"role": "user", "content": "fix this stack trace"}],
)
assert result is not None
assert router.seen_text == "fix this stack trace"
assert router.aembedding_kwargs is not None
forwarded: Final = router.aembedding_kwargs["metadata"]
assert forwarded["user_api_key"] == "hashed-key"
assert forwarded["user_api_key_team_id"] == "team-1"
assert forwarded[INTERNAL_CALL_ORIGIN_METADATA_KEY] == "autorouter_classifier"
assert "user_api_key_budget_reservation" not in forwarded
assert router.aembedding_kwargs["litellm_session_id"] == "session-1"
assert router.aembedding_kwargs["proxy_server_request"] == {
"body": {"model": "text-embedding-3-small", "input": ["fix this stack trace"]}
}

View file

@ -14,6 +14,7 @@ from pydantic import ValidationError
import litellm
from litellm import Router
from litellm.router_utils.auto_router_model_naming import count_heuristic_v2_routers
from litellm._logging import verbose_router_logger
from litellm.caching.dual_cache import DualCache
from litellm.constants import RETURN_RAW_MODEL_NAME_METADATA_KEY, SESSION_ID_GENERATED_METADATA_KEY
@ -1085,6 +1086,165 @@ class TestRouterComplexityDeploymentMethods:
router.init_complexity_router_deployment(deployment)
assert "auto_router/complexity_router/test-router" in router.complexity_routers
@staticmethod
def _router_row(model_name: str, model_id: str, classifier_type: str) -> dict[str, object]:
return {
"model_name": model_name,
"litellm_params": {
"model": "auto_router/complexity_router",
"complexity_router_config": {
"classifier_type": classifier_type,
"tiers": {"SIMPLE": "gpt-4o-mini", "MEDIUM": "gpt-4o"},
},
},
"model_info": {"id": model_id},
}
_POOL: dict[str, object] = {
"model_name": "gpt-4o-mini",
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "k"},
}
def test_heuristic_v2_ceiling_keeps_the_first_router_and_drops_the_rest(self) -> None:
"""The proxy runs with ignore_invalid_deployments, so the second heuristic_v2 router is dropped
at registration while a heuristic (v1) sibling and the first v2 router stay routable."""
router = Router(
model_list=[
self._POOL,
self._router_row("v2-a", "id-a", "heuristic_v2"),
self._router_row("v2-b", "id-b", "heuristic_v2"),
self._router_row("v1-c", "id-c", "heuristic"),
],
heuristic_v2_router_limit=lambda: 1,
ignore_invalid_deployments=True,
)
assert sorted(router.complexity_routers) == ["v1-c", "v2-a"]
assert router.get_deployment(model_id="id-b") is None
def test_heuristic_v2_ceiling_raises_without_ignore_invalid_deployments(self) -> None:
with pytest.raises(ValueError, match="At most 1 auto-router"):
Router(
model_list=[
self._POOL,
self._router_row("v2-a", "id-a", "heuristic_v2"),
self._router_row("v2-b", "id-b", "heuristic_v2"),
],
heuristic_v2_router_limit=lambda: 1,
)
def test_heuristic_v2_limit_is_resolved_on_every_registration(self) -> None:
"""The Router never caches the limit: when the resolver's answer moves (the proxy re-verified
its license), the next registration and the next limit query see the new value."""
limits = {"value": None}
router = Router(
model_list=[
self._POOL,
self._router_row("v2-a", "id-a", "heuristic_v2"),
self._router_row("v2-b", "id-b", "heuristic_v2"),
],
heuristic_v2_router_limit=lambda: limits["value"],
ignore_invalid_deployments=True,
)
assert sorted(router.complexity_routers) == ["v2-a", "v2-b"]
assert router.heuristic_v2_router_limit_violation() is None
limits["value"] = 1
assert router.heuristic_v2_router_limit_violation() is not None
assert router.upsert_deployment(Deployment(**self._router_row("v2-c", "id-c", "heuristic_v2"))) is None
assert sorted(router.complexity_routers) == ["v2-a", "v2-b"]
def test_heuristic_v2_ceiling_tightening_refuses_the_edit_and_keeps_the_live_router(self) -> None:
"""Two heuristic_v2 routers registered under an unlimited ceiling, then the ceiling drops to one:
an edit to either must be refused before its live row is popped, or the failed re-add and
the failed restore would drop a serving router while the write reports success."""
limits = {"value": None}
router = Router(
model_list=[
self._POOL,
self._router_row("v2-a", "id-a", "heuristic_v2"),
self._router_row("v2-b", "id-b", "heuristic_v2"),
],
heuristic_v2_router_limit=lambda: limits["value"],
ignore_invalid_deployments=True,
)
limits["value"] = 1
assert router.upsert_deployment(Deployment(**self._router_row("v2-a-renamed", "id-a", "heuristic_v2"))) is None
assert sorted(router.complexity_routers) == ["v2-a", "v2-b"]
assert router.get_deployment(model_id="id-a") is not None
assert router.upsert_deployment(Deployment(**self._router_row("v1-a", "id-a", "heuristic"))) is not None
assert sorted(router.complexity_routers) == ["v1-a", "v2-b"]
def test_config_deployments_excludes_db_rows(self) -> None:
"""The proxy counts config.yaml routers from here and DB rows from the database, so a DB-loaded
row (``model_info.db_model``) must not show up twice."""
router = Router(model_list=[self._POOL, self._router_row("v2-a", "id-a", "heuristic_v2")])
db_row = self._router_row("v2-db", "id-db", "heuristic_v2")
db_row["model_info"] = {"id": "id-db", "db_model": True}
assert router.upsert_deployment(Deployment(**db_row)) is not None
assert sorted(str(row["model_name"]) for row in router.config_deployments()) == ["gpt-4o-mini", "v2-a"]
assert count_heuristic_v2_routers(router.config_deployments()) == 1
def test_failed_edit_of_a_live_v2_router_rolls_back_without_the_ceiling(self) -> None:
"""A rollback after a failed upsert re-admits state that was already serving, so it must not be
judged by a ceiling that tightened since: converting one of two live heuristic_v2 routers to a
config whose registration fails must leave it serving its previous v2 configuration."""
limits = {"value": None}
router = Router(
model_list=[
self._POOL,
self._router_row("v2-a", "id-a", "heuristic_v2"),
self._router_row("v2-b", "id-b", "heuristic_v2"),
],
heuristic_v2_router_limit=lambda: limits["value"],
ignore_invalid_deployments=True,
)
limits["value"] = 1
broken = self._router_row("v1-a", "id-a", "heuristic")
broken["litellm_params"]["complexity_router_config"]["tiers"] = {}
assert router.upsert_deployment(Deployment(**broken)) is None
assert sorted(router.complexity_routers) == ["v2-a", "v2-b"]
live = router.get_deployment(model_id="id-a")
assert live is not None and live.litellm_params.complexity_router_config["classifier_type"] == "heuristic_v2"
assert router.heuristic_v2_router_limit_violation() is not None
def test_heuristic_v2_routers_are_unlimited_by_default(self) -> None:
router = Router(
model_list=[
self._POOL,
self._router_row("v2-a", "id-a", "heuristic_v2"),
self._router_row("v2-b", "id-b", "heuristic_v2"),
]
)
assert sorted(router.complexity_routers) == ["v2-a", "v2-b"]
assert router.heuristic_v2_router_limit_violation() is None
def test_heuristic_v2_router_limit_violation_frees_the_slot_of_the_router_being_edited(self) -> None:
"""A DB reload upserts the existing heuristic_v2 router again; that edit must keep its own slot
while a different deployment switching to heuristic_v2 is refused."""
router = Router(
model_list=[self._POOL, self._router_row("v2-a", "id-a", "heuristic_v2")],
heuristic_v2_router_limit=lambda: 1,
ignore_invalid_deployments=True,
)
assert router.heuristic_v2_router_limit_violation() is not None
edited = self._router_row("v2-a-renamed", "id-a", "heuristic_v2")
assert router.upsert_deployment(Deployment(**edited)) is not None
assert sorted(router.complexity_routers) == ["v2-a-renamed"]
assert router.upsert_deployment(Deployment(**self._router_row("v2-b", "id-b", "heuristic_v2"))) is None
assert sorted(router.complexity_routers) == ["v2-a-renamed"]
assert router.upsert_deployment(Deployment(**self._router_row("v1-c", "id-c", "heuristic"))) is not None
assert sorted(router.complexity_routers) == ["v1-c", "v2-a-renamed"]
def test_hybrid_initialization_waits_for_later_pool_deployments(self):
router = Router(
model_list=[

View file

@ -1,8 +1,13 @@
from collections.abc import Mapping
import pytest
from litellm.router_utils.auto_router_model_naming import (
carries_complexity_router_settings,
classify_strategy_router_model,
count_heuristic_v2_routers,
heuristic_v2_limit_violation,
is_heuristic_v2_router,
strategy_router_dependencies,
validate_complexity_router_config_placement,
validate_complexity_router_config_write,
@ -369,3 +374,55 @@ def test_placement_is_scoped_to_complexity_router_deployments(model, present_fie
flat param on an s3_vectors vector store, so an unscoped gate would reject a valid deployment.
Either complexity field names one on its own, which is what the load itself requires."""
assert carries_complexity_router_settings(model, present_fields) is scoped
@pytest.mark.parametrize(
"litellm_params,expected",
[
({"model": "auto_router/complexity_router", "complexity_router_config": {"classifier_type": "heuristic_v2"}}, True),
({"model": "auto_router/complexity_router-eu", "complexity_router_config": {"classifier_type": "heuristic_v2"}}, True),
({"model": "auto_router/complexity_router", "complexity_router_config": {"classifier_type": "heuristic"}}, False),
({"model": "auto_router/complexity_router", "complexity_router_config": {"tiers": {"SIMPLE": "a"}}}, False),
({"model": "auto_router/complexity_router"}, False),
({"model": "auto_router/quality_router", "complexity_router_config": {"classifier_type": "heuristic_v2"}}, False),
({"model": "openai/gpt-4o", "complexity_router_config": {"classifier_type": "heuristic_v2"}}, False),
({"model": "auto_router/complexity_router", "complexity_router_config": "heuristic_v2"}, False),
({}, False),
],
)
def test_is_heuristic_v2_router(litellm_params: Mapping[str, object], expected: bool) -> None:
"""Only a complexity router whose config selects heuristic_v2 counts toward the license limit."""
assert is_heuristic_v2_router(litellm_params) is expected
def test_count_heuristic_v2_routers_reads_model_list_rows_and_ignores_malformed_ones() -> None:
v2 = {"model": "auto_router/complexity_router", "complexity_router_config": {"classifier_type": "heuristic_v2"}}
rows: list[Mapping[str, object]] = [
{"model_name": "a", "litellm_params": v2},
{"model_name": "b", "litellm_params": {"model": "openai/gpt-4o"}},
{"model_name": "c", "litellm_params": v2},
{"model_name": "d"},
{"model_name": "e", "litellm_params": "not a mapping"},
]
assert count_heuristic_v2_routers(rows) == 2
assert count_heuristic_v2_routers(()) == 0
@pytest.mark.parametrize(
"held,limit,violates",
[
(1, 1, False),
(2, 1, True),
(0, 1, False),
(5, None, False),
(3, 3, False),
(4, 3, True),
],
)
def test_heuristic_v2_limit_violation(held: int, limit: int | None, violates: bool) -> None:
violation = heuristic_v2_limit_violation(held=held, limit=limit)
assert (violation is not None) is violates
if violation is not None:
assert f"At most {limit} auto-router" in violation
assert f"would make {held}" in violation
assert "license" not in violation

View file

@ -784,6 +784,19 @@ def test_responses_api_bridge_check_gpt_5_4_tools_plus_reasoning_routes_to_respo
assert model_info.get("mode") == "responses"
def test_responses_api_bridge_check_gpt_6_astra_tools_with_default_reasoning_routes_to_responses():
from litellm.main import responses_api_bridge_check
model_info, model = responses_api_bridge_check(
model="gpt-6-astra",
custom_llm_provider="openai",
tools=[{"type": "function", "function": {"name": "get_capital"}}],
)
assert model == "gpt-6-astra"
assert model_info.get("mode") == "responses"
def test_responses_api_bridge_check_gpt_5_5_tools_plus_reasoning_routes_to_responses():
"""gpt-5.5+ with both tools and reasoning_effort should route to Responses API."""
from litellm.main import responses_api_bridge_check

View file

@ -12553,3 +12553,33 @@ async def test_prompt_management_factory_marks_injection_for_every_deployment(mo
bucket = captured.get("litellm_metadata") or captured["metadata"]
assert captured["model_info"]["id"] == "provisional-dep"
assert bucket["litellm_gateway_injected_cache"] == ""
def test_get_configured_mode_reads_deployment_model_info():
router = Router(
model_list=[
{
"model_name": "my-tts",
"litellm_params": {"model": "openai/some-unmapped-mode-model"},
"model_info": {"mode": "audio_speech"},
}
]
)
assert router.get_configured_mode("my-tts") == "audio_speech"
@pytest.mark.parametrize("model_info", [{}, {"mode": ""}, {"mode": " "}, {"mode": 123}])
def test_get_configured_mode_returns_none_for_unset_blank_or_unknown(model_info):
router = Router(
model_list=[
{
"model_name": "plain-model",
"litellm_params": {"model": "openai/some-unmapped-mode-model"},
"model_info": model_info,
}
]
)
assert router.get_configured_mode("plain-model") is None
assert router.get_configured_mode("unknown-model") is None

View file

@ -27,7 +27,7 @@
"limit": 0
},
"LIT010": {
"limit": 16477
"limit": 16476
},
"LIT011": {
"limit": 5519

View file

@ -65,11 +65,13 @@ describe("PromptTable", () => {
expect(within(rows[1]).getByText("prompt-older")).toBeInTheDocument();
});
it("should call onPromptClick when the prompt ID is clicked", async () => {
it("should call onPromptClick with the row's environment, defaulting to development", async () => {
const user = userEvent.setup();
render(<PromptTable {...defaultProps} />);
await user.click(screen.getByRole("button", { name: "prompt-newer" }));
expect(mockOnPromptClick).toHaveBeenCalledWith("prompt-newer");
expect(mockOnPromptClick).toHaveBeenCalledWith("prompt-newer", "production");
await user.click(screen.getByRole("button", { name: "prompt-older" }));
expect(mockOnPromptClick).toHaveBeenCalledWith("prompt-older", "development");
});
it("should label the environment and default missing environments to development", () => {
@ -83,7 +85,7 @@ describe("PromptTable", () => {
render(<PromptTable {...defaultProps} />);
await user.click(screen.getByTestId("prompt-actions-prompt-newer"));
await user.click(await screen.findByTestId("prompt-action-delete"));
expect(mockOnDeleteClick).toHaveBeenCalledWith("prompt-newer", "prompt-newer");
expect(mockOnDeleteClick).toHaveBeenCalledWith("prompt-newer", "prompt-newer", "production");
});
it("should copy the prompt ID through the actions menu", async () => {

View file

@ -13,8 +13,8 @@ import { ModelGroupInfo } from "./prompt_utils";
interface PromptTableProps {
promptsList: PromptSpec[];
isLoading: boolean;
onPromptClick?: (id: string) => void;
onDeleteClick?: (id: string, name: string) => void;
onPromptClick?: (id: string, environment: string) => void;
onDeleteClick?: (id: string, name: string, environment: string) => void;
accessToken: string | null;
isAdmin: boolean;
}
@ -74,7 +74,9 @@ const PromptTable: React.FC<PromptTableProps> = ({
<DataTable
data={promptsList}
columns={columns}
getRowId={(prompt, index) => prompt.prompt_id || String(index)}
getRowId={(prompt, index) =>
prompt.prompt_id ? `${prompt.prompt_id}::${prompt.environment || "development"}` : String(index)
}
sortingMode="client"
sorting={sorting}
onSortingChange={setSorting}

View file

@ -64,7 +64,7 @@ function PromptModelCell({ prompt, modelHubData }: { prompt: PromptSpec; modelHu
interface PromptRowActionsProps {
prompt: PromptSpec;
isAdmin: boolean;
onDeleteClick?: (id: string, name: string) => void;
onDeleteClick?: (id: string, name: string, environment: string) => void;
}
function PromptRowActions({ prompt, isAdmin, onDeleteClick }: PromptRowActionsProps) {
@ -91,7 +91,13 @@ function PromptRowActions({ prompt, isAdmin, onDeleteClick }: PromptRowActionsPr
<DropdownMenuItem
variant="destructive"
data-testid="prompt-action-delete"
onClick={() => onDeleteClick?.(prompt.prompt_id, prompt.prompt_id || "Unknown Prompt")}
onClick={() =>
onDeleteClick?.(
prompt.prompt_id,
prompt.prompt_id || "Unknown Prompt",
prompt.environment || "development",
)
}
>
<Trash2 />
Delete
@ -106,8 +112,8 @@ function PromptRowActions({ prompt, isAdmin, onDeleteClick }: PromptRowActionsPr
interface PromptTableColumnsDeps {
modelHubData: Map<string, ModelGroupInfo>;
isAdmin: boolean;
onPromptClick?: (id: string) => void;
onDeleteClick?: (id: string, name: string) => void;
onPromptClick?: (id: string, environment: string) => void;
onDeleteClick?: (id: string, name: string, environment: string) => void;
}
export const getPromptTableColumns = ({
@ -128,7 +134,11 @@ export const getPromptTableColumns = ({
title={row.original.prompt_id}
titleClassName="font-mono text-xs font-normal"
className="max-w-60"
onClick={onPromptClick ? () => onPromptClick(row.original.prompt_id) : undefined}
onClick={
onPromptClick
? () => onPromptClick(row.original.prompt_id, row.original.environment || "development")
: undefined
}
/>
),
},

View file

@ -16,21 +16,31 @@ vi.mock("./PromptTable", () => ({
__esModule: true,
default: ({
isLoading,
onPromptClick,
onDeleteClick,
}: {
isLoading: boolean;
onDeleteClick: (id: string, name: string) => void;
onPromptClick: (id: string, environment: string) => void;
onDeleteClick: (id: string, name: string, environment: string) => void;
}) => (
<div data-testid="prompt-table">
{isLoading ? "table-loading" : "table-loaded"}
<button type="button" onClick={() => onDeleteClick("prompt-1", "my-prompt")}>
<button type="button" onClick={() => onPromptClick("prompt-1", "staging")}>
row-open
</button>
<button type="button" onClick={() => onDeleteClick("prompt-1", "my-prompt", "staging")}>
row-delete
</button>
</div>
),
}));
vi.mock("./prompt_info", () => ({ __esModule: true, default: () => <div>prompt-info-view</div> }));
vi.mock("./prompt_info", () => ({
__esModule: true,
default: ({ initialEnvironment }: { initialEnvironment?: string }) => (
<div>prompt-info-view:{initialEnvironment ?? "none"}</div>
),
}));
vi.mock("./add_prompt_form", () => ({
__esModule: true,
default: ({ visible }: { visible: boolean }) => (visible ? <div>add-prompt-form</div> : null),
@ -141,6 +151,22 @@ describe("PromptsPanel toolbar", () => {
});
});
describe("PromptsPanel row navigation", () => {
beforeEach(() => {
vi.clearAllMocks();
mockGetPromptsList.mockResolvedValue({ prompts: [] } as never);
});
it("should open the info view preselected to the clicked row's environment", async () => {
const user = userEvent.setup();
renderPanel("Admin");
await user.click(await screen.findByRole("button", { name: "row-open" }));
expect(screen.getByText("prompt-info-view:staging")).toBeInTheDocument();
});
});
describe("PromptsPanel delete confirmation", () => {
beforeEach(() => {
vi.clearAllMocks();
@ -154,13 +180,13 @@ describe("PromptsPanel delete confirmation", () => {
await user.click(await screen.findByRole("button", { name: "row-delete" }));
expect(await screen.findByText(/delete prompt: my-prompt/i)).toBeInTheDocument();
expect(await screen.findByText(/the staging copy of prompt: my-prompt/i)).toBeInTheDocument();
expect(screen.getByText(/cannot be undone/i)).toBeInTheDocument();
expect(mockDeletePromptCall).not.toHaveBeenCalled();
await user.click(screen.getByRole("button", { name: /^delete$/i }));
await waitFor(() => expect(mockDeletePromptCall).toHaveBeenCalledWith("sk-test", "prompt-1"));
await waitFor(() => expect(mockDeletePromptCall).toHaveBeenCalledWith("sk-test", "prompt-1", "staging"));
});
it("should abandon the delete when the confirmation is dismissed", async () => {
@ -168,11 +194,11 @@ describe("PromptsPanel delete confirmation", () => {
renderPanel("Admin");
await user.click(await screen.findByRole("button", { name: "row-delete" }));
await screen.findByText(/delete prompt: my-prompt/i);
await screen.findByText(/the staging copy of prompt: my-prompt/i);
await user.click(screen.getByRole("button", { name: /cancel/i }));
await waitFor(() => expect(screen.queryByText(/delete prompt: my-prompt/i)).not.toBeInTheDocument());
await waitFor(() => expect(screen.queryByText(/the staging copy of prompt: my-prompt/i)).not.toBeInTheDocument());
expect(mockDeletePromptCall).not.toHaveBeenCalled();
});
@ -187,14 +213,14 @@ describe("PromptsPanel delete confirmation", () => {
renderPanel("Admin");
await user.click(await screen.findByRole("button", { name: "row-delete" }));
await screen.findByText(/delete prompt: my-prompt/i);
await screen.findByText(/the staging copy of prompt: my-prompt/i);
await user.click(screen.getByRole("button", { name: /^delete$/i }));
await waitFor(() => expect(mockDeletePromptCall).toHaveBeenCalledWith("sk-test", "prompt-1"));
await waitFor(() => expect(mockDeletePromptCall).toHaveBeenCalledWith("sk-test", "prompt-1", "staging"));
await user.keyboard("{Escape}");
expect(screen.getByText(/delete prompt: my-prompt/i)).toBeInTheDocument();
expect(screen.getByText(/the staging copy of prompt: my-prompt/i)).toBeInTheDocument();
finishDelete();
await waitFor(() => expect(screen.queryByText(/delete prompt: my-prompt/i)).not.toBeInTheDocument());
await waitFor(() => expect(screen.queryByText(/the staging copy of prompt: my-prompt/i)).not.toBeInTheDocument());
});
});

View file

@ -41,11 +41,12 @@ const PromptsPanel: React.FC<PromptsProps> = ({ accessToken, userRole }) => {
const [isLoading, setIsLoading] = useState(true);
const [selectedEnvironment, setSelectedEnvironment] = useState<string | undefined>(undefined);
const [selectedPromptId, setSelectedPromptId] = useState<string | null>(null);
const [selectedPromptEnvironment, setSelectedPromptEnvironment] = useState<string | undefined>(undefined);
const [isAddModalVisible, setIsAddModalVisible] = useState(false);
const [showEditorView, setShowEditorView] = useState(false);
const [editPromptData, setEditPromptData] = useState<any>(null);
const [isDeleting, setIsDeleting] = useState(false);
const [promptToDelete, setPromptToDelete] = useState<{ id: string; name: string } | null>(null);
const [promptToDelete, setPromptToDelete] = useState<{ id: string; name: string; environment: string } | null>(null);
// Admin Viewer follows the read-parity rule: see prompts, no writes.
const canModify = userRole ? isProxyAdminRole(userRole) : false;
@ -71,8 +72,9 @@ const PromptsPanel: React.FC<PromptsProps> = ({ accessToken, userRole }) => {
fetchPrompts();
}, [accessToken, selectedEnvironment]);
const handlePromptClick = (promptId: string) => {
const handlePromptClick = (promptId: string, environment: string) => {
setSelectedPromptId(promptId);
setSelectedPromptEnvironment(environment);
};
const handleAddPrompt = () => {
@ -111,8 +113,8 @@ const PromptsPanel: React.FC<PromptsProps> = ({ accessToken, userRole }) => {
setSelectedPromptId(null);
};
const handleDeleteClick = (promptId: string, promptName: string) => {
setPromptToDelete({ id: promptId, name: promptName });
const handleDeleteClick = (promptId: string, promptName: string, environment: string) => {
setPromptToDelete({ id: promptId, name: promptName, environment });
};
const handleDeleteConfirm = async () => {
@ -120,8 +122,8 @@ const PromptsPanel: React.FC<PromptsProps> = ({ accessToken, userRole }) => {
setIsDeleting(true);
try {
await deletePromptCall(accessToken, promptToDelete.id);
toast.success(`Prompt "${promptToDelete.name}" deleted successfully`);
await deletePromptCall(accessToken, promptToDelete.id, promptToDelete.environment);
toast.success(`Prompt "${promptToDelete.name}" deleted successfully from ${promptToDelete.environment}`);
fetchPrompts(); // Refresh the list
} catch (error) {
console.error("Error deleting prompt:", error);
@ -148,6 +150,7 @@ const PromptsPanel: React.FC<PromptsProps> = ({ accessToken, userRole }) => {
) : selectedPromptId ? (
<PromptInfoView
promptId={selectedPromptId}
initialEnvironment={selectedPromptEnvironment}
onClose={() => setSelectedPromptId(null)}
accessToken={accessToken}
isAdmin={canModify}
@ -219,7 +222,8 @@ const PromptsPanel: React.FC<PromptsProps> = ({ accessToken, userRole }) => {
<AlertDialogHeader>
<AlertDialogTitle>Delete Prompt</AlertDialogTitle>
<AlertDialogDescription>
Are you sure you want to delete prompt: {promptToDelete.name} ? This action cannot be undone.
Are you sure you want to delete the {promptToDelete.environment} copy of prompt: {promptToDelete.name}?
This action cannot be undone.
</AlertDialogDescription>
</AlertDialogHeader>
<AlertDialogFooter>

View file

@ -44,4 +44,28 @@ describe("PromptCodeSnippets", () => {
expect(screen.getByRole("combobox", { name: "Language" })).toHaveTextContent("Python (OpenAI SDK)");
});
it("includes the viewed environment in every generated request", async () => {
const user = userEvent.setup({ pointerEventsCheck: PointerEventsCheckLevel.Never });
render(
<PromptCodeSnippets
promptId="welcome"
model="gpt-4o"
accessToken="token"
version="2"
environment="development"
/>,
);
await user.click(screen.getByRole("button", { name: /get code/i }));
await screen.findByText("Generated Code");
await user.click(screen.getByRole("button", { name: /copy to clipboard/i }));
expect(await navigator.clipboard.readText()).toContain('"prompt_environment": "development"');
await user.click(screen.getByRole("tab", { name: "With Version" }));
await user.click(screen.getByRole("button", { name: /copy to clipboard/i }));
const versionSnippet = await navigator.clipboard.readText();
expect(versionSnippet).toContain('"prompt_environment": "development"');
expect(versionSnippet).toContain('"prompt_version": 2');
});
});

View file

@ -22,6 +22,7 @@ interface PromptCodeSnippetsProps {
promptVariables?: Record<string, string>;
accessToken: string | null;
version?: string;
environment?: string;
proxySettings?: {
PROXY_BASE_URL?: string;
LITELLM_UI_API_DOC_BASE_URL?: string | null;
@ -34,6 +35,7 @@ const PromptCodeSnippets: React.FC<PromptCodeSnippetsProps> = ({
promptVariables = {},
accessToken,
version = "1",
environment,
proxySettings,
}) => {
const syntaxTheme = useSyntaxTheme(coy);
@ -64,6 +66,9 @@ const PromptCodeSnippets: React.FC<PromptCodeSnippetsProps> = ({
// Generate code based on selected language and tab
const generateCode = () => {
const hasVariables = Object.keys(promptVariables).length > 0;
const curlEnvironment = environment ? `,\n "prompt_environment": "${environment}"` : "";
const pythonEnvironment = environment ? `,\n "prompt_environment": "${environment}"` : "";
const jsEnvironment = environment ? `,\n prompt_environment: "${environment}"` : "";
if (selectedLanguage === "curl") {
if (selectedTab === "basic") {
@ -72,7 +77,7 @@ const PromptCodeSnippets: React.FC<PromptCodeSnippetsProps> = ({
-H 'Authorization: Bearer ${effectiveApiKey}' \\
-d '{
"model": "${model}",
"prompt_id": "${promptId}"${
"prompt_id": "${promptId}"${curlEnvironment}${
hasVariables
? `,
"prompt_variables": ${JSON.stringify(promptVariables, null, 6).replace(/\n/g, "\n ")}`
@ -85,7 +90,7 @@ const PromptCodeSnippets: React.FC<PromptCodeSnippetsProps> = ({
-H 'Authorization: Bearer ${effectiveApiKey}' \\
-d '{
"model": "${model}",
"prompt_id": "${promptId}"${
"prompt_id": "${promptId}"${curlEnvironment}${
hasVariables
? `,
"prompt_variables": ${JSON.stringify(promptVariables, null, 6).replace(/\n/g, "\n ")}`
@ -104,7 +109,7 @@ const PromptCodeSnippets: React.FC<PromptCodeSnippetsProps> = ({
-H 'Authorization: Bearer ${effectiveApiKey}' \\
-d '{
"model": "${model}",
"prompt_id": "${promptId}",
"prompt_id": "${promptId}"${curlEnvironment},
"prompt_version": ${version},
"messages": [
{
@ -127,7 +132,7 @@ client = openai.OpenAI(
response = client.chat.completions.create(
model="${model}",
extra_body={
"prompt_id": "${promptId}"${
"prompt_id": "${promptId}"${pythonEnvironment}${
hasVariables
? `,
"prompt_variables": ${JSON.stringify(promptVariables, null, 8).replace(/\n/g, "\n ")}`
@ -145,7 +150,7 @@ response = client.chat.completions.create(
{"role": "user", "content": "hi"}
],
extra_body={
"prompt_id": "${promptId}"${
"prompt_id": "${promptId}"${pythonEnvironment}${
hasVariables
? `,
"prompt_variables": ${JSON.stringify(promptVariables, null, 8).replace(/\n/g, "\n ")}`
@ -163,7 +168,7 @@ response = client.chat.completions.create(
{"role": "user", "content": "Who are u"}
],
extra_body={
"prompt_id": "${promptId}",
"prompt_id": "${promptId}"${pythonEnvironment},
"prompt_version": ${version}
}
)
@ -186,9 +191,9 @@ async function main() {
model: "${model}",
${
hasVariables
? `prompt_id: "${promptId}",
? `prompt_id: "${promptId}"${jsEnvironment},
prompt_variables: ${JSON.stringify(promptVariables, null, 8).replace(/\n/g, "\n ")}`
: `prompt_id: "${promptId}"`
: `prompt_id: "${promptId}"${jsEnvironment}`
}
});
@ -206,9 +211,9 @@ async function main() {
],
${
hasVariables
? `prompt_id: "${promptId}",
? `prompt_id: "${promptId}"${jsEnvironment},
prompt_variables: ${JSON.stringify(promptVariables, null, 8).replace(/\n/g, "\n ")}`
: `prompt_id: "${promptId}"`
: `prompt_id: "${promptId}"${jsEnvironment}`
}
});
@ -224,7 +229,7 @@ async function main() {
messages: [
{ role: "user", content: "Who are u" }
],
prompt_id: "${promptId}",
prompt_id: "${promptId}"${jsEnvironment},
prompt_version: ${version}
});
@ -241,7 +246,7 @@ main();`;
if (isModalVisible) {
setGeneratedCode(generateCode());
}
}, [isModalVisible, selectedLanguage, selectedTab, promptId, model, promptVariables]);
}, [isModalVisible, selectedLanguage, selectedTab, promptId, model, promptVariables, version, environment]);
return (
<>

View file

@ -2,7 +2,9 @@ import { fireEvent, render, screen } from "@testing-library/react";
import { describe, expect, it, vi } from "vitest";
import PromptEditorHeader from "./PromptEditorHeader";
vi.mock("./PromptCodeSnippets", () => ({ default: () => <button>Get Code</button> }));
vi.mock("./PromptCodeSnippets", () => ({
default: ({ environment }: { environment?: string }) => <button data-environment={environment}>Get Code</button>,
}));
describe("PromptEditorHeader", () => {
it("preserves navigation, naming, and save actions", () => {
@ -48,5 +50,6 @@ describe("PromptEditorHeader", () => {
);
expect(screen.getByRole("combobox", { name: "Environment" })).toHaveTextContent(label);
expect(screen.getByRole("button", { name: "Get Code" })).toHaveAttribute("data-environment", environment);
});
});

View file

@ -89,6 +89,7 @@ const PromptEditorHeader: React.FC<PromptEditorHeaderProps> = ({
promptVariables={promptVariables}
accessToken={accessToken}
version={version?.replace("v", "") || "1"}
environment={environment}
proxySettings={proxySettings}
/>
{editMode && onShowHistory && (

View file

@ -12,7 +12,9 @@ vi.mock("@/components/networking", () => ({
}));
vi.mock("./prompt_editor_view/PromptCodeSnippets", () => ({
default: () => <div data-testid="prompt-code-snippets" />,
default: ({ environment }: { environment?: string }) => (
<div data-testid="prompt-code-snippets" data-environment={environment} />
),
}));
const promptWithoutTemplate = {
@ -29,6 +31,67 @@ const promptWithoutTemplate = {
environments: [],
};
describe("PromptInfoView environment scoping", () => {
beforeEach(() => {
vi.mocked(networking.getPromptInfo).mockReset().mockResolvedValue(promptWithoutTemplate);
vi.mocked(networking.getPromptVersions).mockReset().mockResolvedValue({ prompts: [] });
});
it("fetches the initial environment it was opened with", async () => {
render(
<PromptInfoView
promptId="support-reply"
initialEnvironment="staging"
onClose={vi.fn()}
accessToken="sk-test"
isAdmin={true}
/>,
);
await screen.findByRole("tab", { name: "Raw JSON" });
expect(networking.getPromptInfo).toHaveBeenCalledWith("sk-test", "support-reply", "staging");
});
it("fetches the serve default when opened without an environment", async () => {
render(<PromptInfoView promptId="support-reply" onClose={vi.fn()} accessToken="sk-test" isAdmin={true} />);
await screen.findByRole("tab", { name: "Raw JSON" });
expect(networking.getPromptInfo).toHaveBeenCalledWith("sk-test", "support-reply", undefined);
});
});
describe("PromptInfoView code snippets", () => {
beforeEach(() => {
vi.mocked(networking.getPromptVersions).mockReset().mockResolvedValue({ prompts: [] });
});
it.each([
["a prompt with several environments", "staging", ["development", "staging"]],
["a config prompt with no environment list", "development", []],
])("hands the viewed environment of %s to the code snippets", async (_label, environment, environments) => {
vi.mocked(networking.getPromptInfo)
.mockReset()
.mockResolvedValue({
...promptWithoutTemplate,
prompt_spec: { ...promptWithoutTemplate.prompt_spec, environment },
environments,
});
render(
<PromptInfoView
promptId="support-reply"
initialEnvironment={environment}
onClose={vi.fn()}
accessToken="sk-test"
isAdmin={true}
/>,
);
await screen.findByRole("tab", { name: "Raw JSON" });
expect(screen.getByTestId("prompt-code-snippets")).toHaveAttribute("data-environment", environment);
});
});
describe("PromptInfoView tabs", () => {
beforeEach(() => {
vi.mocked(networking.getPromptInfo).mockReset().mockResolvedValue(promptWithoutTemplate);

View file

@ -20,6 +20,7 @@ import { Dialog, DialogContent, DialogFooter, DialogHeader, DialogTitle } from "
export interface PromptInfoProps {
promptId: string;
initialEnvironment?: string;
onClose: () => void;
accessToken: string | null;
isAdmin: boolean;
@ -27,7 +28,15 @@ export interface PromptInfoProps {
onEdit?: (promptData: any) => void;
}
const PromptInfoView: React.FC<PromptInfoProps> = ({ promptId, onClose, accessToken, isAdmin, onDelete, onEdit }) => {
const PromptInfoView: React.FC<PromptInfoProps> = ({
promptId,
initialEnvironment,
onClose,
accessToken,
isAdmin,
onDelete,
onEdit,
}) => {
const [promptData, setPromptData] = useState<PromptSpec | null>(null);
const [promptTemplate, setPromptTemplate] = useState<PromptTemplateBase | null>(null);
const [rawApiResponse, setRawApiResponse] = useState<any>(null);
@ -43,7 +52,7 @@ const PromptInfoView: React.FC<PromptInfoProps> = ({ promptId, onClose, accessTo
const [selectedVersion, setSelectedVersion] = useState<number | null>(null);
const [loadingVersions, setLoadingVersions] = useState(false);
// Initial fetch — no environment filter, gets default + all environments list
// Fetches the requested environment (or the serve-time default when omitted) plus the environments list
const fetchPromptInfo = async (environment?: string) => {
try {
setLoading(true);
@ -89,7 +98,7 @@ const PromptInfoView: React.FC<PromptInfoProps> = ({ promptId, onClose, accessTo
setSelectedEnv(null);
setEnvironments([]);
setVersionHistory([]);
fetchPromptInfo();
fetchPromptInfo(initialEnvironment);
}, [promptId, accessToken]);
// When environment changes (user clicks tab), re-fetch — skip initial mount
@ -212,6 +221,7 @@ const PromptInfoView: React.FC<PromptInfoProps> = ({ promptId, onClose, accessTo
promptVariables={extractTemplateVariables(promptTemplate?.content)}
accessToken={accessToken}
version={currentVersion}
environment={selectedEnv ?? promptData.environment}
/>
<Button onClick={() => onEdit?.(rawApiResponse)} className="flex items-center">
<Pencil />
@ -493,7 +503,7 @@ const PromptInfoView: React.FC<PromptInfoProps> = ({ promptId, onClose, accessTo
<DialogTitle>Delete Prompt</DialogTitle>
</DialogHeader>
<p>
Are you sure you want to delete prompt: <strong>{basePromptId}</strong>?
Are you sure you want to delete prompt: <strong>{basePromptId}</strong> from every environment?
</p>
<p>This action cannot be undone.</p>
<DialogFooter>

View file

@ -323,5 +323,26 @@ describe("ViewUserDashboard", () => {
expect(latest[2]).toBe(1);
});
});
it("sends the toolbar search as the combined search param instead of user_email", async () => {
const user = userEvent.setup();
renderDashboard();
await waitFor(() => {
expect(screen.getByText("test@example.com")).toBeInTheDocument();
});
const searchedUserId = "a6f5c02b-0163-45ce-815f-f88d10e95686";
await user.type(screen.getByPlaceholderText("Search by email or ID…"), searchedUserId);
await waitFor(() => {
const latest = userListCall.mock.calls[userListCall.mock.calls.length - 1];
expect(latest[11]).toBe(searchedUserId);
});
const latest = userListCall.mock.calls[userListCall.mock.calls.length - 1];
expect(latest[1]).toBeNull();
expect(latest[4]).toBeNull();
expect(latest[2]).toBe(1);
});
});
});

View file

@ -65,7 +65,7 @@ const ViewUserDashboard: React.FC<ViewUserDashboardProps> = ({
const [sorting, setSorting] = useState<SortingState>(DEFAULT_SORTING);
const [columnFilters, setColumnFilters] = useState<ColumnFiltersState>([]);
const [searchInput, setSearchInput] = useState("");
const [searchEmail] = useDebouncedValue(searchInput, { wait: DEBOUNCE_WAIT_MS });
const [searchQuery] = useDebouncedValue(searchInput, { wait: DEBOUNCE_WAIT_MS });
const [rowSelection, setRowSelection] = useState<RowSelectionState>({});
const [selectionMode, setSelectionMode] = useState(false);
@ -222,12 +222,12 @@ const ViewUserDashboard: React.FC<ViewUserDashboardProps> = ({
const ssoUserIdFilter = getFilterValue("sso_user_id");
const userRoleFilter = getFilterValue("user_role");
const teamFilter = getFilterValue("team");
const emailFilter = searchEmail.trim() || null;
const searchFilter = searchQuery.trim() || null;
const userListQueryFilters = {
page: pagination.pageIndex + 1,
pageSize: pagination.pageSize,
email: emailFilter,
search: searchFilter,
userId: userIdFilter,
ssoUserId: ssoUserIdFilter,
role: userRoleFilter,
@ -247,13 +247,14 @@ const ViewUserDashboard: React.FC<ViewUserDashboardProps> = ({
userIdFilter ? [userIdFilter] : null,
pagination.pageIndex + 1,
pagination.pageSize,
emailFilter,
null,
userRoleFilter ?? null,
teamFilter ?? null,
ssoUserIdFilter ?? null,
sortBy,
sortOrder,
orgAdminOrgIds ? orgAdminOrgIds.map((o) => o.organization_id) : null,
searchFilter,
);
},
enabled: Boolean(accessToken && token && userRole && userID),

View file

@ -158,7 +158,7 @@ export function UsersTable({
table={table}
searchValue={searchValue}
onSearchChange={onSearchChange}
searchPlaceholder="Search by email…"
searchPlaceholder="Search by email or ID…"
onOpenFilters={() => setFiltersOpen(true)}
filterLabels={FILTER_LABELS}
formatFilterValue={formatFilterValue}

View file

@ -815,3 +815,42 @@ describe("daily activity api_key filter", () => {
expect(requestedUrl(mockFetch)).toContain("user_id=");
});
});
describe("userListCall search serialization", () => {
const originalFetch = global.fetch;
afterEach(() => {
global.fetch = originalFetch;
});
const mockOkFetch = () => {
const emptyPage = { users: [], total: 0, page: 1, page_size: 25, total_pages: 0 };
const body = JSON.stringify(emptyPage);
const mockFetch = vi.fn().mockResolvedValue({ ok: true, text: vi.fn().mockResolvedValue(body) } as any);
global.fetch = mockFetch as any;
return mockFetch;
};
const lastParams = (mockFetch: ReturnType<typeof vi.fn>) => {
const [url] = mockFetch.mock.calls.at(-1) ?? [];
return new URL(url as string, "http://example.com").searchParams;
};
it("sends the combined search term as search, not user_email", async () => {
const mockFetch = mockOkFetch();
await Networking.userListCall("token", null, 1, 25, null, null, null, null, null, null, null, "a6f5c02b");
expect(lastParams(mockFetch).get("search")).toBe("a6f5c02b");
expect(lastParams(mockFetch).has("user_email")).toBe(false);
});
it("omits search when no search term is given and keeps user_email as before", async () => {
const mockFetch = mockOkFetch();
await Networking.userListCall("token", null, 1, 25, "ada@example.com");
expect(lastParams(mockFetch).has("search")).toBe(false);
expect(lastParams(mockFetch).get("user_email")).toBe("ada@example.com");
});
});

View file

@ -1033,6 +1033,7 @@ export const userListCall = async (
sortBy: string | null = null,
sortOrder: "asc" | "desc" | null = null,
organizationIds: string[] | null = null,
search: string | null = null,
) => {
/**
* Get all available teams on proxy
@ -1051,6 +1052,7 @@ export const userListCall = async (
sort_by: sortBy || undefined,
sort_order: sortOrder || undefined,
organization_ids: organizationIds && organizationIds.length > 0 ? organizationIds.join(",") : undefined,
search: search || undefined,
},
})) as UserListResponse;
return data;
@ -4694,9 +4696,12 @@ export const updatePromptCall = async (accessToken: string, promptId: string, pr
}
};
export const deletePromptCall = async (accessToken: string, promptId: string) => {
export const deletePromptCall = async (accessToken: string, promptId: string, environment?: string) => {
try {
const data = await apiClient.delete(`/prompts/${promptId}`, { accessToken });
const data = await apiClient.delete(`/prompts/${promptId}`, {
accessToken,
query: { environment: environment || undefined },
});
return data;
} catch (error) {
console.error("Failed to delete prompt:", error);

View file

@ -334,7 +334,7 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey, autoOp
const fetchPrompts = async () => {
try {
const response = await getPromptsList(accessToken);
setPromptsList(response.prompts.map((prompt) => prompt.prompt_id));
setPromptsList(Array.from(new Set(response.prompts.map((prompt) => prompt.prompt_id))));
} catch (error) {
console.error("Failed to fetch prompts:", error);
}

View file

@ -1888,6 +1888,8 @@ describe("TeamInfoView - the exact bytes the update call sends", () => {
mcp_access_groups: [],
mcp_tool_permissions: {},
mcp_toolsets: [],
agents: [],
agent_access_groups: [],
vector_stores: ["vs-1"],
};
@ -1908,6 +1910,49 @@ describe("TeamInfoView - the exact bytes the update call sends", () => {
});
});
const openEditorWithAgents = async (user: ReturnType<typeof userEvent.setup>) => {
vi.mocked(networking.teamInfoCall).mockResolvedValue(
createMockTeamData({
models: ["gpt-4"],
object_permission: { agents: ["agent-1"], agent_access_groups: ["group-a"] },
}),
);
vi.mocked(networking.teamUpdateCall).mockResolvedValue({ data: {}, team_id: "123" } as any);
renderWithProviders(<TeamInfoView {...props} />);
await waitFor(() => expect(screen.queryAllByText("Test Team").length).toBeGreaterThan(0));
await user.click(screen.getByRole("tab", { name: "Settings" }));
await user.click(await screen.findByRole("button", { name: /edit settings/i }));
await screen.findByLabelText("Team Name");
};
it("resends the stored agents and agent_access_groups when the selector is left untouched", async () => {
const user = userEvent.setup({ delay: null });
await openEditorWithAgents(user);
const payload = await save(user);
const objectPermission = wireBody(payload).object_permission as Record<string, unknown>;
expect(objectPermission.agents).toStrictEqual(["agent-1"]);
expect(objectPermission.agent_access_groups).toStrictEqual(["group-a"]);
});
it("sends empty agents and agent_access_groups arrays after the last agent chip is removed", async () => {
const user = userEvent.setup({ delay: null });
await openEditorWithAgents(user);
await user.click(within(screen.getByLabelText("agent-1")).getByRole("button"));
await user.click(within(screen.getByLabelText("group:group-a")).getByRole("button"));
expect(screen.queryByLabelText("agent-1")).not.toBeInTheDocument();
expect(screen.queryByLabelText("group:group-a")).not.toBeInTheDocument();
const payload = await save(user);
const objectPermission = wireBody(payload).object_permission as Record<string, unknown>;
expect(objectPermission.agents).toStrictEqual([]);
expect(objectPermission.agent_access_groups).toStrictEqual([]);
});
it("resends every stored value once both sections are opened", async () => {
const user = userEvent.setup({ delay: null });
await openEditor(user);

View file

@ -863,12 +863,8 @@ const TeamInfoView: React.FC<TeamInfoProps> = ({
agents: [],
accessGroups: [],
};
if (agents && agents.length > 0) {
updateData.object_permission.agents = agents;
}
if (agentAccessGroups && agentAccessGroups.length > 0) {
updateData.object_permission.agent_access_groups = agentAccessGroups;
}
updateData.object_permission.agents = agents;
updateData.object_permission.agent_access_groups = agentAccessGroups;
delete values.agents_and_groups;
// Handle vector stores permissions

View file

@ -425,6 +425,22 @@ describe("KeyEditView", () => {
expect(screen.getByText("Policies")).toBeInTheDocument();
});
it("lists a prompt existing in several environments once in the dropdown", async () => {
vi.mocked(getPromptsList).mockResolvedValueOnce({
prompts: [
{ prompt_id: "envgreet", litellm_params: {}, prompt_info: { prompt_type: "db" }, environment: "development" },
{ prompt_id: "envgreet", litellm_params: {}, prompt_info: { prompt_type: "db" }, environment: "production" },
],
});
renderAs("Admin");
const prompts = await screen.findByLabelText(/Prompts/);
await userEvent.type(prompts, "envgreet");
expect(await screen.findAllByRole("option", { name: "envgreet" })).toHaveLength(1);
});
it("should omit both fields and fire neither admin-only request for an internal user", async () => {
renderAs("Internal User");

View file

@ -165,7 +165,7 @@ export function KeyEditView({
if (!accessToken) return;
try {
const response = await getPromptsList(accessToken);
setPromptsList(response.prompts.map((prompt) => prompt.prompt_id));
setPromptsList(Array.from(new Set(response.prompts.map((prompt) => prompt.prompt_id))));
} catch (error) {
console.error("Failed to fetch prompts:", error);
}

View file

@ -16606,6 +16606,8 @@ export interface paths {
* Get list of users by sso_ids. Comma separated list of sso_ids.
* user_email: Optional[str]
* Filter users by partial email match
* search: Optional[str]
* Combined search: matches users whose user_id or user_email contains the value (case-insensitive)
* team: Optional[str]
* Filter users by team id. Will match if user has this team in their teams array.
* page: int
@ -59835,6 +59837,8 @@ export interface operations {
sso_user_ids?: string | null;
/** @description Filter users by partial email match */
user_email?: string | null;
/** @description Combined search: matches users whose 'user_id' or 'user_email' contains the value (case-insensitive). */
search?: string | null;
/** @description Filter users by team id */
team?: string | null;
/** @description Page number */