mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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
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:
commit
7d5b6456ba
65 changed files with 2461 additions and 805 deletions
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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])
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.",
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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, ...]:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -3612,6 +3612,7 @@ all_litellm_params = (
|
|||
"litellm_system_prompt",
|
||||
"provider_specific_header",
|
||||
"prompt_version",
|
||||
"prompt_environment",
|
||||
"api_base",
|
||||
"force_timeout",
|
||||
"logger_fn",
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"}}
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"]}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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=[
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -27,7 +27,7 @@
|
|||
"limit": 0
|
||||
},
|
||||
"LIT010": {
|
||||
"limit": 16477
|
||||
"limit": 16476
|
||||
},
|
||||
"LIT011": {
|
||||
"limit": 5519
|
||||
|
|
|
|||
|
|
@ -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 () => {
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
/>
|
||||
),
|
||||
},
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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');
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
<>
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -89,6 +89,7 @@ const PromptEditorHeader: React.FC<PromptEditorHeaderProps> = ({
|
|||
promptVariables={promptVariables}
|
||||
accessToken={accessToken}
|
||||
version={version?.replace("v", "") || "1"}
|
||||
environment={environment}
|
||||
proxySettings={proxySettings}
|
||||
/>
|
||||
{editMode && onShowHistory && (
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
|
|
|
|||
4
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
4
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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 */
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue