mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_lit_7022_azure_ai_passthrough_config
This commit is contained in:
commit
6ed72693dc
23 changed files with 1989 additions and 582 deletions
|
|
@ -46,6 +46,7 @@ from litellm.proxy._types import (
|
|||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.openai_files_endpoints.common_utils import (
|
||||
BATCH_CREATE_HIDDEN_PARAM,
|
||||
FILE_LIST_CONTINUATION_CHUNK_SIZE,
|
||||
MAX_FILE_LIST_LIMIT,
|
||||
_is_base64_encoded_unified_file_id,
|
||||
|
|
@ -1321,7 +1322,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
## Check if unified_file_id is in the response
|
||||
unified_file_id = response._hidden_params.get("unified_file_id") # managed file id
|
||||
unified_batch_id = response._hidden_params.get("unified_batch_id") # managed batch id
|
||||
is_batch_create: Final = unified_file_id is not None
|
||||
is_batch_create: Final = response._hidden_params.get(BATCH_CREATE_HIDDEN_PARAM) is True
|
||||
model_id = cast(Optional[str], response._hidden_params.get("model_id"))
|
||||
model_name = cast(Optional[str], response._hidden_params.get("model_name"))
|
||||
|
||||
|
|
@ -1410,10 +1411,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
|
||||
)
|
||||
|
||||
# Only record batch creation metric on actual create (not retrieve/cancel).
|
||||
# unified_file_id in _hidden_params is only set by the create_batch endpoint.
|
||||
original_unified_file_id = response._hidden_params.get("unified_file_id")
|
||||
if original_unified_file_id:
|
||||
if is_batch_create:
|
||||
prom_logger = self._get_prometheus_logger()
|
||||
if prom_logger:
|
||||
batch_provider = ""
|
||||
|
|
|
|||
|
|
@ -13,9 +13,19 @@ import json
|
|||
import os
|
||||
import re
|
||||
import time
|
||||
from collections.abc import AsyncIterator, Awaitable, Callable, Mapping, Sequence
|
||||
from collections.abc import (
|
||||
AsyncIterator,
|
||||
Awaitable,
|
||||
Callable,
|
||||
Container,
|
||||
Iterable,
|
||||
Mapping,
|
||||
MutableMapping,
|
||||
Sequence,
|
||||
)
|
||||
from contextlib import asynccontextmanager
|
||||
from dataclasses import dataclass, replace
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, TypeAlias, TypedDict, cast
|
||||
from urllib.parse import ParseResult, urlparse
|
||||
|
||||
|
|
@ -307,6 +317,7 @@ class MCPServerConfig(TypedDict, total=False):
|
|||
:meth:`MCPServerManager.load_servers_from_config`. Every key is optional: YAML supplies
|
||||
whatever the admin wrote, and each read applies its own default."""
|
||||
|
||||
server_id: ReadOnly[str]
|
||||
alias: str
|
||||
description: str
|
||||
mcp_info: MCPInfo
|
||||
|
|
@ -400,6 +411,164 @@ def _blank_to_none(value: str | None) -> str | None:
|
|||
return value.strip() or None
|
||||
|
||||
|
||||
def _pinned_config_server_id(raw_server_id: object, server_name: str) -> str | None:
|
||||
"""Return the ``server_id`` an admin pinned for this config.yaml server, or ``None`` when absent.
|
||||
|
||||
Without a pin the id is derived by hashing ``server_name|url|transport|auth_type|alias``, so
|
||||
editing any of those fields mints a new id and every ``object_permission.mcp_servers`` grant
|
||||
holding the old one silently stops matching. A pinned id is used verbatim and survives those
|
||||
edits. Blank and non-string values are rejected rather than silently falling back to the hash,
|
||||
because a config that pins an id and still churns is the failure this field exists to prevent.
|
||||
|
||||
Under ``LITELLM_USE_SHORT_MCP_TOOL_PREFIX`` the tool prefix is derived from the server_id, so
|
||||
pinning an id other than the one already in use renames every tool that server exposes.
|
||||
"""
|
||||
if raw_server_id is None:
|
||||
return None
|
||||
if not isinstance(raw_server_id, str) or not raw_server_id.strip():
|
||||
raise ValueError(
|
||||
f"Invalid config for MCP server '{server_name}': server_id must be a non-empty string "
|
||||
f"(got {raw_server_id!r})."
|
||||
)
|
||||
return raw_server_id.strip()
|
||||
|
||||
|
||||
def _first_mapped_alias(server_name: str, mcp_aliases: Mapping[str, str] | None) -> str | None:
|
||||
"""The ``mcp_aliases`` name ``load_servers_from_config`` will assign to this server, if any.
|
||||
|
||||
Mirrors that loop, which takes the first mapping pointing at the server and stops. A later
|
||||
mapping for the same server is never applied, so it stays free for another entry to pin.
|
||||
"""
|
||||
if mcp_aliases is None:
|
||||
return None
|
||||
return next(
|
||||
(alias_name for alias_name, target_server_name in mcp_aliases.items() if target_server_name == server_name),
|
||||
None,
|
||||
)
|
||||
|
||||
|
||||
def _assigned_alias(
|
||||
server_name: str, server_config: MCPServerConfig, mcp_aliases: Mapping[str, str] | None
|
||||
) -> str | None:
|
||||
"""The alias ``load_servers_from_config`` will give this entry: its own, else the first mapping.
|
||||
|
||||
``is None``, not falsiness: the loader only consults the mapping when the key is absent, so an
|
||||
entry that sets ``alias: ""`` gets no mapped alias and reserves nothing.
|
||||
"""
|
||||
alias: Final = server_config.get("alias")
|
||||
return _first_mapped_alias(server_name, mcp_aliases) if alias is None else alias
|
||||
|
||||
|
||||
def _validate_config_server_names(mcp_servers_config: Mapping[str, MCPServerConfig]) -> None:
|
||||
"""Reject bad server names before ``_config_identifier_owners`` reads any entry's body.
|
||||
|
||||
The identifier index walks every entry up front, so without this pass a malformed entry under
|
||||
a bad name would surface as an ``AttributeError`` from the index instead of the name error.
|
||||
"""
|
||||
for server_name in mcp_servers_config:
|
||||
validate_mcp_server_name(server_name)
|
||||
|
||||
|
||||
def _config_identifier_owners(
|
||||
mcp_servers_config: Mapping[str, MCPServerConfig],
|
||||
mcp_aliases: Mapping[str, str] | None,
|
||||
) -> Mapping[str, frozenset[str]]:
|
||||
"""Map every server_name and alias in the config to the entries that own it.
|
||||
|
||||
``expand_permission_list`` resolves a grant against the registry keys before it falls back to
|
||||
matching alias and server_name, so an id equal to another entry's name or alias captures that
|
||||
entry's grants. Derived ids are hashes and never collide with a name, so this only matters once
|
||||
an id is pinned.
|
||||
|
||||
An alias is either set on the entry or mapped to it from ``litellm_settings.mcp_aliases``. Only
|
||||
a name the loader below will really assign is reserved: the mapping is ignored for an entry that
|
||||
sets its own ``alias``, and only the first mapping wins for one that does not, so reserving every
|
||||
mapping would fail startup on a pin that was never going to collide.
|
||||
|
||||
One identifier can have several owners when an entry's alias equals another entry's name. All of
|
||||
them are kept: a grant naming that identifier resolves to every match while no id is pinned, and
|
||||
a pin equal to it would narrow the grant to the pinning entry alone, even when that entry is one
|
||||
of the owners.
|
||||
"""
|
||||
claims: Final = tuple(
|
||||
(identifier, server_name)
|
||||
for server_name, server_config in mcp_servers_config.items()
|
||||
for identifier in (server_name, _assigned_alias(server_name, server_config, mcp_aliases))
|
||||
if identifier
|
||||
)
|
||||
return MappingProxyType(
|
||||
{identifier: frozenset(owner for claimed, owner in claims if claimed == identifier) for identifier, _ in claims}
|
||||
)
|
||||
|
||||
|
||||
def _config_ids_capturing_db_identifiers(
|
||||
config_server_ids: Container[str],
|
||||
db_servers: Iterable[MCPServer],
|
||||
) -> frozenset[str]:
|
||||
"""Config server ids that are a database-backed server's name, server_name or alias.
|
||||
|
||||
``expand_permission_list`` matches a grant against the registry keys before it matches names, so
|
||||
such an id answers every grant written for the database server, and the database server itself
|
||||
stops being reachable by name. The config load cannot catch this because the database registry
|
||||
is not loaded yet, so it is reported from the reload that does have both halves.
|
||||
|
||||
An identifier equal to the database server's own id is skipped: ``get_registry`` is
|
||||
``config_mcp_servers | registry``, so there the database server wins the id outright and the
|
||||
shadow warning above is the accurate one. Reporting both would contradict. The skip is per
|
||||
identifier rather than per server, so a row that shadows one config id and captures another
|
||||
still reports the capture.
|
||||
"""
|
||||
return frozenset(
|
||||
identifier
|
||||
for server in db_servers
|
||||
for identifier in (server.name, server.server_name, server.alias)
|
||||
if identifier and identifier != server.server_id and identifier in config_server_ids
|
||||
)
|
||||
|
||||
|
||||
def _reject_config_server_id_collision(
|
||||
assigned_server_ids: Mapping[str, str],
|
||||
server_id: str,
|
||||
server_name: str,
|
||||
pinned: bool,
|
||||
db_backed_server_ids: Mapping[str, object],
|
||||
identifier_owners: Mapping[str, frozenset[str]],
|
||||
) -> None:
|
||||
"""Raise when ``server_id`` is already taken, either by an earlier config entry or by the database.
|
||||
|
||||
Two config entries sharing an id would silently overwrite each other in ``config_mcp_servers``,
|
||||
and an id already held by a database-backed server is hidden by it, because ``get_registry`` is
|
||||
``config_mcp_servers | registry`` and the right operand wins. A pinned id that is another
|
||||
entry's server_name or alias captures that entry's permission grants the same way. Derived ids
|
||||
cannot collide (the unique config key is part of the hash input), so all three only happen once
|
||||
an id is pinned.
|
||||
|
||||
Pinning an identifier this entry itself owns is allowed, because a grant naming it already
|
||||
resolved here, but only when no other entry owns it too. An entry whose alias is this entry's
|
||||
server_name shares the identifier, and pinning it would take that entry's grants.
|
||||
"""
|
||||
claimed_by = assigned_server_ids.get(server_id)
|
||||
if claimed_by is not None:
|
||||
raise ValueError(
|
||||
f"Invalid config for MCP server '{server_name}': server_id '{server_id}' is already "
|
||||
f"used by MCP server '{claimed_by}'. Each mcp_servers entry needs its own id."
|
||||
)
|
||||
if pinned and server_id in db_backed_server_ids:
|
||||
raise ValueError(
|
||||
f"Invalid config for MCP server '{server_name}': server_id '{server_id}' belongs to a "
|
||||
"database-backed MCP server. The database entry takes precedence over config.yaml, so "
|
||||
"this server would never be reachable."
|
||||
)
|
||||
other_owners: Final = identifier_owners.get(server_id, frozenset()) - frozenset((server_name,))
|
||||
if pinned and other_owners:
|
||||
owner_names: Final = "', '".join(sorted(other_owners))
|
||||
raise ValueError(
|
||||
f"Invalid config for MCP server '{server_name}': server_id '{server_id}' is the "
|
||||
f"server_name or alias of MCP server '{owner_names}'. Permission entries naming "
|
||||
f"'{server_id}' would resolve to '{server_name}' alone and no longer reach '{owner_names}'."
|
||||
)
|
||||
|
||||
|
||||
def _uses_issuer_anchor(manual_issuer: str | None, is_discovery_auth_type: bool) -> bool:
|
||||
"""Whether the endpoints are authoritatively anchored to an admin-pinned issuer (RFC 8414 §3.3).
|
||||
|
||||
|
|
@ -1565,6 +1734,11 @@ class MCPServerManager:
|
|||
# empty result, or failure). Used to throttle re-probes for servers that do
|
||||
# not return instructions, and to apply a short cooldown after failures.
|
||||
self._upstream_initialize_instructions_probed_at: dict[str, float] = {}
|
||||
# Last set of config server ids found shadowed by database rows. reload_servers_from_database
|
||||
# runs on the config-reload timer, so this keeps a standing misconfiguration from re-logging
|
||||
# the same warning every interval; a change in the set logs again.
|
||||
self._warned_shadowed_config_server_ids: frozenset[str] = frozenset()
|
||||
self._warned_capturing_config_server_ids: frozenset[str] = frozenset()
|
||||
self._oauth_discovery_on_startup = _mcp_oauth_discovery_on_startup_enabled()
|
||||
self._oauth_discovery_generation_counter = 0
|
||||
self._oauth_discovery_slots: tuple[_OAuthDiscoverySlot, ...] = ()
|
||||
|
|
@ -1958,10 +2132,14 @@ class MCPServerManager:
|
|||
|
||||
# Track which aliases have been used to ensure only first occurrence is used
|
||||
used_aliases: Final = set()
|
||||
# server_id -> the config server_name that claimed it, so a pinned id cannot silently
|
||||
# overwrite another server's entry in self.config_mcp_servers.
|
||||
assigned_server_ids: MutableMapping[str, str] = {} # mutable-ok: per-load collision index
|
||||
_validate_config_server_names(mcp_servers_config)
|
||||
identifier_owners: Final = _config_identifier_owners(mcp_servers_config, mcp_aliases)
|
||||
|
||||
for server_name, raw_server_config in mcp_servers_config.items():
|
||||
server_config: MCPServerConfig = raw_server_config
|
||||
validate_mcp_server_name(server_name)
|
||||
_mcp_info: MCPInfo = server_config.get("mcp_info", None) or {}
|
||||
# Preserve all custom fields from config while setting defaults for core fields
|
||||
mcp_info: MCPInfo = _mcp_info.copy()
|
||||
|
|
@ -1994,14 +2172,24 @@ class MCPServerManager:
|
|||
name_for_prefix = get_server_prefix(temp_server)
|
||||
|
||||
server_url = server_config.get("url", None) or ""
|
||||
# Generate stable server ID based on parameters
|
||||
server_id = self._generate_stable_server_id(
|
||||
# An explicitly pinned server_id wins; otherwise derive one from the parameters.
|
||||
pinned_server_id = _pinned_config_server_id(server_config.get("server_id"), server_name)
|
||||
server_id = pinned_server_id or self._generate_stable_server_id(
|
||||
server_name=server_name,
|
||||
url=server_url,
|
||||
transport=server_config.get("transport", MCPTransport.http),
|
||||
auth_type=server_config.get("auth_type", None),
|
||||
alias=alias,
|
||||
)
|
||||
_reject_config_server_id_collision(
|
||||
assigned_server_ids,
|
||||
server_id,
|
||||
server_name,
|
||||
pinned=pinned_server_id is not None,
|
||||
db_backed_server_ids=self.registry,
|
||||
identifier_owners=identifier_owners,
|
||||
)
|
||||
assigned_server_ids[server_id] = server_name
|
||||
|
||||
_warn_on_server_name_fields(
|
||||
server_id=server_id,
|
||||
|
|
@ -6123,6 +6311,33 @@ class MCPServerManager:
|
|||
|
||||
verbose_logger.debug("MCP registry refreshed (%s servers in registry)", len(registered_registry))
|
||||
|
||||
# get_registry() is ``config_mcp_servers | registry``, so a database row sharing an id with a
|
||||
# config.yaml server hides that server everywhere. Only reachable once an operator pins
|
||||
# ``server_id`` in config.yaml; say so rather than letting the server disappear silently.
|
||||
shadowed_config_server_ids: Final = frozenset(self.config_mcp_servers.keys() & registered_registry.keys())
|
||||
if shadowed_config_server_ids and shadowed_config_server_ids != self._warned_shadowed_config_server_ids:
|
||||
verbose_logger.warning(
|
||||
"config.yaml MCP server_id(s) %s are also database-backed MCP servers. The database "
|
||||
"entry takes precedence, so the config.yaml server is unreachable. Give the config "
|
||||
"entry a different server_id.",
|
||||
", ".join(sorted(shadowed_config_server_ids)),
|
||||
)
|
||||
self._warned_shadowed_config_server_ids = shadowed_config_server_ids
|
||||
|
||||
# The mirror image of the block above: a config server_id that is a database server's name
|
||||
# answers that server's grants instead, because ids are matched before names.
|
||||
capturing_config_server_ids: Final = _config_ids_capturing_db_identifiers(
|
||||
self.config_mcp_servers.keys(), registered_registry.values()
|
||||
)
|
||||
if capturing_config_server_ids and capturing_config_server_ids != self._warned_capturing_config_server_ids:
|
||||
verbose_logger.warning(
|
||||
"config.yaml MCP server_id(s) %s are the name or alias of a database-backed MCP "
|
||||
"server. Permission entries naming them resolve to the config.yaml server, not the "
|
||||
"database one. Give the config entry a different server_id.",
|
||||
", ".join(sorted(capturing_config_server_ids)),
|
||||
)
|
||||
self._warned_capturing_config_server_ids = capturing_config_server_ids
|
||||
|
||||
await self._hydrate_config_servers_dcr_clients()
|
||||
|
||||
def get_mcp_servers_from_ids(self, server_ids: list[str]) -> list[MCPServer]:
|
||||
|
|
|
|||
|
|
@ -2638,6 +2638,10 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
|
|||
default=None,
|
||||
description="Set-up pass-through endpoints for provider-specific endpoints. Docs - https://docs.litellm.ai/docs/proxy/pass_through",
|
||||
)
|
||||
enable_openai_websocket_passthrough: bool | None = Field(
|
||||
default=None,
|
||||
description="Serve the OpenAI pass-through WebSocket route, which relays frames to OpenAI under the proxy's own provider credential without reading them. Off by default.",
|
||||
)
|
||||
user_header_name: str | None = Field(
|
||||
None,
|
||||
description="[DEPRECATED] Use 'user_header_mappings' instead. When set, the header value is treated as the end user id unless overridden by user_header_mappings.",
|
||||
|
|
|
|||
|
|
@ -2352,6 +2352,13 @@ async def _backfill_null_user_email(
|
|||
return updated_row
|
||||
|
||||
|
||||
class UserNotFoundError(ValueError):
|
||||
"""The user row is provably absent, as opposed to merely unreadable, so a caller that reads a missing row as no user-level limits can key on it without also swallowing a database that would not answer."""
|
||||
|
||||
def __init__(self, user_id: str) -> None:
|
||||
super().__init__(f"User doesn't exist in db. 'user_id'={user_id}. Create user via `/user/new` call.")
|
||||
|
||||
|
||||
@log_db_metrics
|
||||
async def get_user_object(
|
||||
user_id: str | None,
|
||||
|
|
@ -2457,7 +2464,7 @@ async def get_user_object(
|
|||
value=None,
|
||||
last_db_access_time=last_db_access_time,
|
||||
)
|
||||
raise Exception
|
||||
raise UserNotFoundError(user_id=user_id)
|
||||
|
||||
if response.organization_memberships is not None and len(response.organization_memberships) > 0:
|
||||
# dump each organization membership to type LiteLLM_OrganizationMembershipTable
|
||||
|
|
@ -2493,7 +2500,9 @@ async def get_user_object(
|
|||
)
|
||||
|
||||
return _response
|
||||
except Exception as e: # if user not in db
|
||||
except UserNotFoundError:
|
||||
raise
|
||||
except Exception as e:
|
||||
_log_budget_lookup_failure("user", e)
|
||||
raise ValueError(
|
||||
f"User doesn't exist in db. 'user_id'={user_id}. Create user via `/user/new` call. Got error - {e}"
|
||||
|
|
@ -4155,6 +4164,79 @@ async def _granted_model_lists(
|
|||
)
|
||||
|
||||
|
||||
async def _user_object_or_none(
|
||||
valid_token: UserAPIKeyAuth,
|
||||
prisma_client: PrismaClient,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
) -> LiteLLM_UserTable | None:
|
||||
try:
|
||||
return await get_user_object(
|
||||
user_id=valid_token.user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
except UserNotFoundError:
|
||||
return None
|
||||
|
||||
|
||||
async def enforced_model_allowlists(
|
||||
valid_token: UserAPIKeyAuth,
|
||||
prisma_client: PrismaClient | None,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
) -> tuple[Sequence[str], ...]:
|
||||
"""One model allowlist per level that ``common_checks`` enforces on a request from this identity."""
|
||||
key_models: Final = _resolve_key_models_for_auth_check(valid_token=valid_token)
|
||||
if prisma_client is None:
|
||||
return (key_models, tuple(valid_token.team_models or ()))
|
||||
team_object: Final = (
|
||||
None
|
||||
if valid_token.team_id is None
|
||||
else await get_team_object(
|
||||
team_id=valid_token.team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
)
|
||||
user_object: Final = (
|
||||
None
|
||||
if team_object is not None
|
||||
else await _user_object_or_none(
|
||||
valid_token=valid_token,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
)
|
||||
project_object: Final = (
|
||||
None
|
||||
if valid_token.project_id is None
|
||||
else await get_project_object(
|
||||
project_id=valid_token.project_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
)
|
||||
return (
|
||||
key_models,
|
||||
team_object.models if team_object is not None else (),
|
||||
await _team_member_granted_models(
|
||||
valid_token=valid_token,
|
||||
team_object=team_object,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
),
|
||||
user_object.models if user_object is not None else (),
|
||||
project_object.models if project_object is not None else (),
|
||||
)
|
||||
|
||||
|
||||
async def collect_matched_model_access_groups(
|
||||
model: str | Sequence[str] | None,
|
||||
valid_token: UserAPIKeyAuth | None,
|
||||
|
|
|
|||
|
|
@ -25,6 +25,7 @@ from litellm.proxy.common_utils.openai_endpoint_utils import (
|
|||
get_custom_llm_provider_from_request_query,
|
||||
)
|
||||
from litellm.proxy.openai_files_endpoints.common_utils import (
|
||||
BATCH_CREATE_HIDDEN_PARAM,
|
||||
_is_base64_encoded_unified_file_id,
|
||||
add_internal_model_credentials,
|
||||
apply_team_provider_credentials,
|
||||
|
|
@ -347,6 +348,8 @@ async def create_batch(
|
|||
**_create_batch_data,
|
||||
)
|
||||
|
||||
response._hidden_params[BATCH_CREATE_HIDDEN_PARAM] = True
|
||||
|
||||
### CALL HOOKS ### - modify outgoing data
|
||||
response = await proxy_logging_obj.post_call_success_hook(
|
||||
data=data, user_api_key_dict=user_api_key_dict, response=response
|
||||
|
|
|
|||
|
|
@ -37,6 +37,8 @@ MAX_FILE_LIST_LIMIT: Final = 10000
|
|||
|
||||
FILE_LIST_CONTINUATION_CHUNK_SIZE: Final = 500
|
||||
|
||||
BATCH_CREATE_HIDDEN_PARAM: Final = "batch_create"
|
||||
|
||||
|
||||
def validate_file_list_limit(limit: int | None) -> None:
|
||||
"""Reject a ``limit`` outside the range OpenAI documents for GET /v1/files."""
|
||||
|
|
|
|||
|
|
@ -13,14 +13,16 @@ import inspect
|
|||
import json
|
||||
import os
|
||||
import re
|
||||
from collections.abc import AsyncGenerator, Callable, Mapping
|
||||
from collections.abc import AsyncGenerator, Callable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Annotated, Final, cast
|
||||
from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, cast
|
||||
|
||||
import httpx
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, Response, WebSocket
|
||||
from fastapi.responses import StreamingResponse
|
||||
from starlette.websockets import WebSocketState
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm import get_llm_provider
|
||||
|
|
@ -35,6 +37,7 @@ from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
|||
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
|
||||
from litellm.passthrough.main import AsyncPassthroughStreamingResponse
|
||||
from litellm.proxy._types import *
|
||||
from litellm.proxy.auth.auth_checks import enforced_model_allowlists
|
||||
from litellm.proxy.auth.handle_jwt import JWTHandler
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
from litellm.proxy.auth.user_api_key_auth import (
|
||||
|
|
@ -1780,7 +1783,7 @@ def _upstream_headers_for_vertex_route(endpoint: str, headers: Mapping[str, str]
|
|||
|
||||
|
||||
def get_vertex_pass_through_handler(
|
||||
call_type: Literal["discovery", "aiplatform"], # noqa: UP037 # ruff reports quoted Literal values here
|
||||
call_type: Literal["discovery", "aiplatform"],
|
||||
) -> BaseVertexAIPassThroughHandler:
|
||||
if call_type == "discovery":
|
||||
return VertexAIDiscoveryPassThroughHandler()
|
||||
|
|
@ -2347,9 +2350,102 @@ _OPENAI_WS_ALL_MODEL_ACCESS: Final = frozenset(
|
|||
)
|
||||
|
||||
|
||||
def _key_has_model_restrictions(user_api_key_dict: UserAPIKeyAuth) -> bool:
|
||||
scoped_models: Final = (*user_api_key_dict.models, *user_api_key_dict.team_models)
|
||||
return any(str(model) not in _OPENAI_WS_ALL_MODEL_ACCESS for model in scoped_models)
|
||||
def _has_model_restrictions(model_allowlists: tuple[Sequence[str], ...]) -> bool:
|
||||
return any(str(model) not in _OPENAI_WS_ALL_MODEL_ACCESS for allowlist in model_allowlists for model in allowlist)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _OpenAIWebsocketRefusal:
|
||||
close_reason: str
|
||||
message: str
|
||||
|
||||
|
||||
class _OpenAIWebsocketErrorDetail(TypedDict):
|
||||
type: ReadOnly[Literal["invalid_request_error"]]
|
||||
message: ReadOnly[str]
|
||||
|
||||
|
||||
class _OpenAIWebsocketErrorFrame(TypedDict):
|
||||
type: ReadOnly[Literal["error"]]
|
||||
error: ReadOnly[_OpenAIWebsocketErrorDetail]
|
||||
|
||||
|
||||
_OPENAI_WS_DISABLED_REFUSAL: Final = _OpenAIWebsocketRefusal(
|
||||
close_reason="OpenAI websocket passthrough is disabled",
|
||||
message=(
|
||||
"OpenAI websocket passthrough is disabled on this gateway. A proxy admin can turn it on by "
|
||||
"setting general_settings.enable_openai_websocket_passthrough to true."
|
||||
),
|
||||
)
|
||||
|
||||
_OPENAI_WS_MODEL_RESTRICTED_REFUSAL: Final = _OpenAIWebsocketRefusal(
|
||||
close_reason="Keys with model restrictions cannot use OpenAI websocket passthrough",
|
||||
message=(
|
||||
"Keys with model restrictions cannot use OpenAI websocket passthrough, because this route "
|
||||
"relays frames to the provider without reading which model they ask for."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _is_openai_websocket_passthrough_enabled(general_settings: Mapping[str, object]) -> bool:
|
||||
setting: Final = general_settings.get("enable_openai_websocket_passthrough")
|
||||
if isinstance(setting, str):
|
||||
return str_to_bool(setting) is True
|
||||
return setting is True
|
||||
|
||||
|
||||
class _OpenAIWebsocketModelAllowlists(Protocol):
|
||||
async def __call__(self, valid_token: UserAPIKeyAuth, /) -> tuple[Sequence[str], ...]: ...
|
||||
|
||||
|
||||
async def _openai_websocket_refusal(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
general_settings: Mapping[str, object],
|
||||
model_allowlists: _OpenAIWebsocketModelAllowlists,
|
||||
) -> _OpenAIWebsocketRefusal | None:
|
||||
if not _is_openai_websocket_passthrough_enabled(general_settings):
|
||||
return _OPENAI_WS_DISABLED_REFUSAL
|
||||
if _has_model_restrictions(await model_allowlists(user_api_key_dict)):
|
||||
return _OPENAI_WS_MODEL_RESTRICTED_REFUSAL
|
||||
return None
|
||||
|
||||
|
||||
class _OpenAIWebsocketRelay(Protocol):
|
||||
async def __call__(
|
||||
self,
|
||||
*,
|
||||
websocket: WebSocket,
|
||||
target: str,
|
||||
custom_headers: dict[str, str], # mutable-ok: the relay takes a plain dict of upstream headers
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
forward_headers: bool,
|
||||
endpoint: str,
|
||||
accept_websocket: bool,
|
||||
) -> None: ...
|
||||
|
||||
|
||||
def _proxy_general_settings() -> Mapping[str, object]:
|
||||
from litellm.proxy.proxy_server import general_settings
|
||||
|
||||
return general_settings
|
||||
|
||||
|
||||
def _openai_websocket_relay() -> _OpenAIWebsocketRelay:
|
||||
return websocket_passthrough_request
|
||||
|
||||
|
||||
def _proxy_model_allowlists() -> _OpenAIWebsocketModelAllowlists:
|
||||
from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache
|
||||
|
||||
async def resolve(valid_token: UserAPIKeyAuth, /) -> tuple[Sequence[str], ...]:
|
||||
return await enforced_model_allowlists(
|
||||
valid_token=valid_token,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
return resolve
|
||||
|
||||
|
||||
@router.websocket("/openai_passthrough/{endpoint:path}")
|
||||
|
|
@ -2358,13 +2454,27 @@ async def openai_websocket_proxy_route(
|
|||
websocket: WebSocket,
|
||||
endpoint: str,
|
||||
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth_websocket)],
|
||||
general_settings: Annotated[Mapping[str, object], Depends(_proxy_general_settings)],
|
||||
relay: Annotated[_OpenAIWebsocketRelay, Depends(_openai_websocket_relay)],
|
||||
model_allowlists: Annotated[_OpenAIWebsocketModelAllowlists, Depends(_proxy_model_allowlists)],
|
||||
) -> None:
|
||||
"""WebSocket passthrough for OpenAI prefixes (realtime / responses.connect)."""
|
||||
if _key_has_model_restrictions(user_api_key_dict):
|
||||
await websocket.close(
|
||||
code=1008,
|
||||
reason="Keys with model restrictions cannot use OpenAI websocket passthrough",
|
||||
)
|
||||
requested_subprotocols: Final = tuple(
|
||||
protocol.strip()
|
||||
for protocol in (websocket.headers.get("sec-websocket-protocol") or "").split(",")
|
||||
if protocol.strip()
|
||||
)
|
||||
negotiated_subprotocol: Final = requested_subprotocols[0] if requested_subprotocols else None
|
||||
|
||||
refusal: Final = await _openai_websocket_refusal(user_api_key_dict, general_settings, model_allowlists)
|
||||
if refusal is not None:
|
||||
await websocket.accept(subprotocol=negotiated_subprotocol)
|
||||
error_frame: Final[_OpenAIWebsocketErrorFrame] = {
|
||||
"type": "error",
|
||||
"error": {"type": "invalid_request_error", "message": refusal.message},
|
||||
}
|
||||
await websocket.send_text(json.dumps(error_frame))
|
||||
await websocket.close(code=1008, reason=refusal.close_reason)
|
||||
return
|
||||
|
||||
base_target_url: Final = os.getenv("OPENAI_API_BASE") or "https://api.openai.com/"
|
||||
|
|
@ -2400,14 +2510,9 @@ async def openai_websocket_proxy_route(
|
|||
"Authorization": f"Bearer {openai_api_key}"
|
||||
}
|
||||
|
||||
requested_subprotocols: Final = tuple(
|
||||
protocol.strip()
|
||||
for protocol in (websocket.headers.get("sec-websocket-protocol") or "").split(",")
|
||||
if protocol.strip()
|
||||
)
|
||||
await websocket.accept(subprotocol=requested_subprotocols[0] if requested_subprotocols else None)
|
||||
await websocket.accept(subprotocol=negotiated_subprotocol)
|
||||
|
||||
await websocket_passthrough_request(
|
||||
await relay(
|
||||
websocket=websocket,
|
||||
target=wss_target,
|
||||
custom_headers=custom_headers,
|
||||
|
|
|
|||
|
|
@ -6810,6 +6810,11 @@ class ProxyConfig:
|
|||
else:
|
||||
general_settings["apply_user_budget_to_team_keys"] = db_value if db_value is None else bool(db_value)
|
||||
|
||||
if "enable_openai_websocket_passthrough" not in self._yaml_general_settings_keys:
|
||||
general_settings["enable_openai_websocket_passthrough"] = _general_settings.get(
|
||||
"enable_openai_websocket_passthrough"
|
||||
)
|
||||
|
||||
## STORE MODEL IN DB ##
|
||||
if "store_model_in_db" in _general_settings:
|
||||
value = _general_settings["store_model_in_db"]
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@
|
|||
"limit": 733
|
||||
},
|
||||
"TQ002": {
|
||||
"limit": 741
|
||||
"limit": 737
|
||||
},
|
||||
"TQ003": {
|
||||
"limit": 62
|
||||
|
|
@ -21,6 +21,6 @@
|
|||
"limit": 117
|
||||
},
|
||||
"TQ008": {
|
||||
"limit": 11003
|
||||
"limit": 10993
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFi
|
|||
from litellm.caching import DualCache
|
||||
from litellm.proxy._types import CallTypes
|
||||
from litellm.proxy.openai_files_endpoints.common_utils import (
|
||||
BATCH_CREATE_HIDDEN_PARAM,
|
||||
_is_base64_encoded_unified_file_id,
|
||||
encode_file_id_with_model,
|
||||
)
|
||||
|
|
@ -3185,7 +3186,7 @@ def _batch_response(batch_id, output_file_id=None, is_create=False):
|
|||
output_file_id=output_file_id,
|
||||
)
|
||||
if is_create:
|
||||
batch._hidden_params["unified_file_id"] = "unified-input-file-id"
|
||||
batch._hidden_params[BATCH_CREATE_HIDDEN_PARAM] = True
|
||||
return batch
|
||||
|
||||
|
||||
|
|
@ -3411,11 +3412,8 @@ async def test_provider_format_file_without_ownership_row_stays_accessible():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_batch_create_stores_ownership_row():
|
||||
"""
|
||||
Batch creation (response hidden params carry the unified input file id)
|
||||
must write an ownership row attributed to the creating key.
|
||||
"""
|
||||
@pytest.mark.parametrize("batch_id", [MODEL_ENCODED_BATCH_ID, RAW_PROVIDER_BATCH_ID])
|
||||
async def test_post_call_batch_create_stores_ownership_row(batch_id):
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
prisma_client = AsyncMock()
|
||||
|
|
@ -3432,13 +3430,11 @@ async def test_post_call_batch_create_stores_ownership_row():
|
|||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_id="user_a", team_id="team_a", parent_otel_span=MagicMock()
|
||||
),
|
||||
response=_batch_response(MODEL_ENCODED_BATCH_ID, is_create=True),
|
||||
response=_batch_response(batch_id, is_create=True),
|
||||
)
|
||||
|
||||
upsert_call = prisma_client.db.litellm_managedobjecttable.upsert.await_args
|
||||
assert upsert_call.kwargs["where"] == {
|
||||
"unified_object_id": MODEL_ENCODED_BATCH_ID
|
||||
}
|
||||
assert upsert_call.kwargs["where"] == {"unified_object_id": batch_id}
|
||||
create_data = upsert_call.kwargs["data"]["create"]
|
||||
assert create_data["created_by"] == "user_a"
|
||||
assert create_data["team_id"] == "team_a"
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ from typing import Optional
|
|||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from litellm.proxy._types import LitellmUserRoles, ProxyException, UserAPIKeyAuth
|
||||
from litellm.proxy.openai_files_endpoints.common_utils import BATCH_CREATE_HIDDEN_PARAM
|
||||
from litellm.types.llms.openai import FileListPage, OpenAIFileObject
|
||||
from litellm.types.utils import LiteLLMBatch
|
||||
|
||||
|
|
@ -1540,6 +1541,11 @@ async def test_batch_create_hook_persists_creating_key_and_tags():
|
|||
managed_files = _make_managed_files_instance()
|
||||
creator = UserAPIKeyAuth(api_key="sk-the-creator", user_id="alice", parent_otel_span=None)
|
||||
create_response = _make_batch_response(status="validating", output_file_id=None)
|
||||
create_response._hidden_params = {
|
||||
BATCH_CREATE_HIDDEN_PARAM: True,
|
||||
"model_id": "model-deploy-xyz",
|
||||
"model_name": "azure/gpt-4",
|
||||
}
|
||||
|
||||
await managed_files.async_post_call_success_hook(
|
||||
data={"litellm_metadata": {"tags": ["env:prod", "team:ml"], "user_api_key": creator.api_key}},
|
||||
|
|
@ -1554,6 +1560,52 @@ async def test_batch_create_hook_persists_creating_key_and_tags():
|
|||
assert stored["user_api_key_dict"] is creator
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_batch_create_hook_records_created_metric_once():
|
||||
managed_files = _make_managed_files_instance()
|
||||
prometheus_logger = MagicMock()
|
||||
managed_files._get_prometheus_logger = MagicMock(return_value=prometheus_logger)
|
||||
create_response = _make_batch_response(status="validating", output_file_id=None)
|
||||
create_response._hidden_params = {
|
||||
BATCH_CREATE_HIDDEN_PARAM: True,
|
||||
"model_id": "model-deploy-xyz",
|
||||
"model_name": "azure/gpt-4",
|
||||
}
|
||||
|
||||
await managed_files.async_post_call_success_hook(
|
||||
data={},
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="sk-the-creator", user_id="alice", parent_otel_span=None),
|
||||
response=create_response,
|
||||
)
|
||||
|
||||
prometheus_logger.record_managed_batch_created.assert_called_once()
|
||||
recorded = prometheus_logger.record_managed_batch_created.call_args.kwargs
|
||||
assert recorded["model"] == "azure/gpt-4"
|
||||
assert recorded["api_provider"] == "azure"
|
||||
assert recorded["user"] == "alice"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_batch_retrieve_hook_does_not_record_created_metric():
|
||||
managed_files = _make_managed_files_instance()
|
||||
prometheus_logger = MagicMock()
|
||||
managed_files._get_prometheus_logger = MagicMock(return_value=prometheus_logger)
|
||||
retrieve_response = _make_batch_response(status="in_progress", output_file_id=None)
|
||||
retrieve_response._hidden_params = {
|
||||
"unified_batch_id": "some-unified-batch-id",
|
||||
"model_id": "model-deploy-xyz",
|
||||
"model_name": "azure/gpt-4",
|
||||
}
|
||||
|
||||
await managed_files.async_post_call_success_hook(
|
||||
data={},
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="sk-the-poller", user_id="bob", parent_otel_span=None),
|
||||
response=retrieve_response,
|
||||
)
|
||||
|
||||
prometheus_logger.record_managed_batch_created.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_batch_retrieve_hook_does_not_claim_attribution():
|
||||
"""A retrieve carries unified_batch_id but no unified_file_id, so it must not rewrite
|
||||
|
|
|
|||
|
|
@ -11428,6 +11428,569 @@ class TestOpenApiHandlerRelaysUpstreamAuth:
|
|||
assert "upstream returned HTTP 503" in result.content[0].text
|
||||
|
||||
|
||||
class TestConfigServerIdPinning:
|
||||
"""config.yaml servers may pin ``server_id`` so permission grants survive connection edits."""
|
||||
|
||||
@staticmethod
|
||||
def _config(**overrides: object) -> dict[str, dict[str, object]]:
|
||||
return {
|
||||
"docs_server": {
|
||||
"url": "https://example.com/mcp",
|
||||
"transport": MCPTransport.http,
|
||||
**overrides,
|
||||
}
|
||||
}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_derived_id_churns_when_connection_fields_change(self):
|
||||
"""The behavior the pin exists to escape: editing the url mints a brand-new id."""
|
||||
manager = MCPServerManager()
|
||||
|
||||
await manager.load_servers_from_config(self._config())
|
||||
before = next(iter(manager.config_mcp_servers))
|
||||
|
||||
manager.config_mcp_servers.clear()
|
||||
await manager.load_servers_from_config(self._config(url="https://prod.example.com/mcp"))
|
||||
after = next(iter(manager.config_mcp_servers))
|
||||
|
||||
assert before != after
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pinned_id_survives_url_transport_auth_and_alias_edits(self):
|
||||
manager = MCPServerManager()
|
||||
|
||||
await manager.load_servers_from_config(self._config(server_id="docs-prod-1"))
|
||||
assert list(manager.config_mcp_servers) == ["docs-prod-1"]
|
||||
assert manager.config_mcp_servers["docs-prod-1"].server_id == "docs-prod-1"
|
||||
|
||||
manager.config_mcp_servers.clear()
|
||||
await manager.load_servers_from_config(
|
||||
self._config(
|
||||
server_id="docs-prod-1",
|
||||
url="https://prod.example.com/mcp",
|
||||
transport=MCPTransport.sse,
|
||||
auth_type=MCPAuth.bearer_token,
|
||||
alias="docs",
|
||||
)
|
||||
)
|
||||
|
||||
assert list(manager.config_mcp_servers) == ["docs-prod-1"]
|
||||
assert manager.config_mcp_servers["docs-prod-1"].url == "https://prod.example.com/mcp"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_absent_server_id_keeps_the_derived_hash(self):
|
||||
manager = MCPServerManager()
|
||||
|
||||
await manager.load_servers_from_config(self._config())
|
||||
|
||||
derived = manager._generate_stable_server_id(
|
||||
server_name="docs_server",
|
||||
url="https://example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=None,
|
||||
alias=None,
|
||||
)
|
||||
assert list(manager.config_mcp_servers) == [derived]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("bad_value", ["", " ", 123, True, ["docs-prod-1"]])
|
||||
async def test_blank_or_non_string_server_id_is_rejected(self, bad_value: Any):
|
||||
manager = MCPServerManager()
|
||||
|
||||
with pytest.raises(ValueError, match="server_id must be a non-empty string"):
|
||||
await manager.load_servers_from_config(self._config(server_id=bad_value))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_two_servers_pinning_the_same_id_are_rejected(self):
|
||||
manager = MCPServerManager()
|
||||
config: Dict[str, Any] = {
|
||||
"docs_server": {"url": "https://a.example.com/mcp", "server_id": "shared-id"},
|
||||
"wiki_server": {"url": "https://b.example.com/mcp", "server_id": "shared-id"},
|
||||
}
|
||||
|
||||
with pytest.raises(ValueError, match="already used by MCP server 'docs_server'"):
|
||||
await manager.load_servers_from_config(config)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pinned_id_colliding_with_a_derived_id_is_rejected(self):
|
||||
"""A pin that lands on another entry's derived hash collides just as hard."""
|
||||
manager = MCPServerManager()
|
||||
derived = manager._generate_stable_server_id(
|
||||
server_name="docs_server",
|
||||
url="https://a.example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=None,
|
||||
alias=None,
|
||||
)
|
||||
config: Dict[str, Any] = {
|
||||
"docs_server": {"url": "https://a.example.com/mcp", "transport": MCPTransport.http},
|
||||
"wiki_server": {"url": "https://b.example.com/mcp", "server_id": derived},
|
||||
}
|
||||
|
||||
with pytest.raises(ValueError, match="already used by MCP server 'docs_server'"):
|
||||
await manager.load_servers_from_config(config)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pinned_id_colliding_with_a_db_backed_server_is_rejected(self):
|
||||
"""get_registry() is ``config | registry``, so the db row would hide the config server.
|
||||
|
||||
The registry is seeded by hand because on a real startup the config loads before the
|
||||
database does, so this check only fires on a later reload. The startup ordering is covered
|
||||
by ``test_db_row_arriving_on_a_pinned_config_id_warns``; the warning there is not redundant.
|
||||
"""
|
||||
manager = MCPServerManager()
|
||||
manager.registry["db-uuid-1"] = MCPServer(
|
||||
server_id="db-uuid-1",
|
||||
name="db_server",
|
||||
transport=MCPTransport.http,
|
||||
url="https://db.example.com/mcp",
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="belongs to a database-backed MCP server"):
|
||||
await manager.load_servers_from_config(self._config(server_id="db-uuid-1"))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_derived_id_matching_a_db_backed_server_is_not_rejected(self):
|
||||
"""Only a pinned id is an authoring error; a hash collision must not fail startup."""
|
||||
manager = MCPServerManager()
|
||||
derived = manager._generate_stable_server_id(
|
||||
server_name="docs_server",
|
||||
url="https://example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=None,
|
||||
alias=None,
|
||||
)
|
||||
manager.registry[derived] = MCPServer(
|
||||
server_id=derived,
|
||||
name="db_server",
|
||||
transport=MCPTransport.http,
|
||||
url="https://db.example.com/mcp",
|
||||
)
|
||||
|
||||
await manager.load_servers_from_config(self._config())
|
||||
|
||||
assert derived in manager.config_mcp_servers
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pinned_id_is_stripped_of_surrounding_whitespace(self):
|
||||
manager = MCPServerManager()
|
||||
|
||||
await manager.load_servers_from_config(self._config(server_id=" docs-prod-1 "))
|
||||
|
||||
assert list(manager.config_mcp_servers) == ["docs-prod-1"]
|
||||
|
||||
@staticmethod
|
||||
async def _reload_with_db_server(manager: MCPServerManager, server_id: str, db_name: str = "db_server") -> None:
|
||||
row = LiteLLM_MCPServerTable(
|
||||
server_id=server_id,
|
||||
server_name=db_name,
|
||||
alias=db_name,
|
||||
url="https://db.example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
)
|
||||
raw_row = MagicMock()
|
||||
raw_row.model_dump.return_value = row.model_dump()
|
||||
repository = MagicMock()
|
||||
repository.table.find_many = AsyncMock(return_value=[raw_row])
|
||||
built = MCPServer(
|
||||
server_id=server_id,
|
||||
name=db_name,
|
||||
server_name=db_name,
|
||||
url="https://db.example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
)
|
||||
with (
|
||||
patch( # test-quality-ok: the db reload path has no seam but its own repository
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository",
|
||||
return_value=repository,
|
||||
),
|
||||
patch( # test-quality-ok: same, the prisma client is fetched inside the reload
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
patch.object(manager, "build_mcp_server_from_table", new=AsyncMock(return_value=built)),
|
||||
):
|
||||
await manager.reload_servers_from_database()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_db_row_arriving_on_a_pinned_config_id_warns(self, caplog):
|
||||
"""The db row loads after config on startup, so the config server is hidden then, not at load."""
|
||||
manager = MCPServerManager()
|
||||
await manager.load_servers_from_config(self._config(server_id="docs-prod-1"))
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
|
||||
await self._reload_with_db_server(manager, "docs-prod-1")
|
||||
|
||||
assert any("docs-prod-1" in m and "database entry takes precedence" in m for m in caplog.messages)
|
||||
assert manager.get_registry()["docs-prod-1"].url == "https://db.example.com/mcp"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_db_row_with_a_distinct_id_does_not_warn(self, caplog):
|
||||
manager = MCPServerManager()
|
||||
await manager.load_servers_from_config(self._config(server_id="docs-prod-1"))
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
|
||||
await self._reload_with_db_server(manager, "db-uuid-1")
|
||||
|
||||
assert all("database entry takes precedence" not in m for m in caplog.messages)
|
||||
assert set(manager.get_registry()) == {"docs-prod-1", "db-uuid-1"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pinned_id_matching_another_entrys_server_name_is_rejected(self):
|
||||
"""expand_permission_list resolves against registry keys first, so this steals the grants."""
|
||||
manager = MCPServerManager()
|
||||
|
||||
with pytest.raises(ValueError, match="server_name or alias of MCP server 'wiki_server'"):
|
||||
await manager.load_servers_from_config(
|
||||
{
|
||||
"wiki_server": {"url": "https://wiki.example.com/mcp", "transport": MCPTransport.http},
|
||||
"docs_server": {
|
||||
"server_id": "wiki_server",
|
||||
"url": "https://example.com/mcp",
|
||||
"transport": MCPTransport.http,
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pinned_id_matching_another_entrys_alias_is_rejected(self):
|
||||
manager = MCPServerManager()
|
||||
|
||||
with pytest.raises(ValueError, match="server_name or alias of MCP server 'wiki_server'"):
|
||||
await manager.load_servers_from_config(
|
||||
{
|
||||
"wiki_server": {
|
||||
"alias": "wiki",
|
||||
"url": "https://wiki.example.com/mcp",
|
||||
"transport": MCPTransport.http,
|
||||
},
|
||||
"docs_server": {
|
||||
"server_id": "wiki",
|
||||
"url": "https://example.com/mcp",
|
||||
"transport": MCPTransport.http,
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pinning_a_servers_own_name_is_allowed(self):
|
||||
"""The most natural pin an operator writes; it resolves to the same server either way."""
|
||||
manager = MCPServerManager()
|
||||
|
||||
await manager.load_servers_from_config(self._config(server_id="docs_server"))
|
||||
|
||||
assert list(manager.config_mcp_servers) == ["docs_server"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pinning_a_servers_own_alias_is_allowed(self):
|
||||
manager = MCPServerManager()
|
||||
|
||||
await manager.load_servers_from_config(self._config(alias="docs", server_id="docs"))
|
||||
|
||||
assert list(manager.config_mcp_servers) == ["docs"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("aliasing_entry_first", [True, False])
|
||||
async def test_pinning_own_name_that_is_another_entrys_alias_is_rejected(self, aliasing_entry_first: bool):
|
||||
"""A grant naming 'docs_server' reaches both servers unpinned; the pin would narrow it to one."""
|
||||
manager = MCPServerManager()
|
||||
wiki = (
|
||||
"wiki_server",
|
||||
{"alias": "docs_server", "url": "https://wiki.example.com/mcp", "transport": MCPTransport.http},
|
||||
)
|
||||
docs = (
|
||||
"docs_server",
|
||||
{"server_id": "docs_server", "url": "https://example.com/mcp", "transport": MCPTransport.http},
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="server_name or alias of MCP server 'wiki_server'"):
|
||||
await manager.load_servers_from_config(dict((wiki, docs) if aliasing_entry_first else (docs, wiki)))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pinning_own_name_that_is_another_entrys_mapped_alias_is_rejected(self):
|
||||
manager = MCPServerManager()
|
||||
|
||||
with pytest.raises(ValueError, match="server_name or alias of MCP server 'wiki_server'"):
|
||||
await manager.load_servers_from_config(
|
||||
{
|
||||
"wiki_server": {"url": "https://wiki.example.com/mcp", "transport": MCPTransport.http},
|
||||
"docs_server": {
|
||||
"server_id": "docs_server",
|
||||
"url": "https://example.com/mcp",
|
||||
"transport": MCPTransport.http,
|
||||
},
|
||||
},
|
||||
mcp_aliases={"docs_server": "wiki_server"},
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pinning_own_alias_shared_with_a_later_entry_is_rejected(self):
|
||||
"""Nothing rejects duplicate aliases, so the first entry's pin would answer the second's grants."""
|
||||
manager = MCPServerManager()
|
||||
|
||||
with pytest.raises(ValueError, match="server_name or alias of MCP server 'docs_server'"):
|
||||
await manager.load_servers_from_config(
|
||||
{
|
||||
"wiki_server": {
|
||||
"alias": "shared",
|
||||
"server_id": "shared",
|
||||
"url": "https://wiki.example.com/mcp",
|
||||
"transport": MCPTransport.http,
|
||||
},
|
||||
"docs_server": {
|
||||
"alias": "shared",
|
||||
"url": "https://example.com/mcp",
|
||||
"transport": MCPTransport.http,
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_own_name_pin_resolves_grants_like_the_unpinned_name(self):
|
||||
"""The negative control: a sole-owner self-pin must keep loading and answer the same grants."""
|
||||
manager = MCPServerManager()
|
||||
|
||||
await manager.load_servers_from_config(
|
||||
{
|
||||
"wiki_server": {"alias": "wiki", "url": "https://wiki.example.com/mcp", "transport": MCPTransport.http},
|
||||
"docs_server": {
|
||||
"server_id": "docs_server",
|
||||
"url": "https://example.com/mcp",
|
||||
"transport": MCPTransport.http,
|
||||
},
|
||||
}
|
||||
)
|
||||
wiki_id = next(sid for sid, server in manager.config_mcp_servers.items() if server.alias == "wiki")
|
||||
|
||||
assert manager.expand_permission_list(["docs_server"]) == ["docs_server"]
|
||||
assert manager.expand_permission_list(["wiki"]) == [wiki_id]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_derived_id_is_not_checked_against_names(self):
|
||||
"""Unpinned configs must keep loading; only a pinned id can be an authoring error."""
|
||||
manager = MCPServerManager()
|
||||
|
||||
await manager.load_servers_from_config(
|
||||
{
|
||||
"wiki_server": {"url": "https://wiki.example.com/mcp", "transport": MCPTransport.http},
|
||||
"docs_server": {"url": "https://example.com/mcp", "transport": MCPTransport.http},
|
||||
}
|
||||
)
|
||||
|
||||
assert len(manager.config_mcp_servers) == 2
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_shadow_warning_is_not_repeated_on_every_reload(self, caplog):
|
||||
"""reload_servers_from_database runs on the config-reload timer; one warning, not one a tick."""
|
||||
manager = MCPServerManager()
|
||||
await manager.load_servers_from_config(self._config(server_id="docs-prod-1"))
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
|
||||
await self._reload_with_db_server(manager, "docs-prod-1")
|
||||
first_round = [m for m in caplog.messages if "database entry takes precedence" in m]
|
||||
await self._reload_with_db_server(manager, "docs-prod-1")
|
||||
second_round = [m for m in caplog.messages if "database entry takes precedence" in m]
|
||||
|
||||
assert len(first_round) == 1
|
||||
assert second_round == first_round
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_shadow_warning_fires_again_when_the_shadowed_set_changes(self, caplog):
|
||||
manager = MCPServerManager()
|
||||
await manager.load_servers_from_config(self._config(server_id="docs-prod-1"))
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
|
||||
await self._reload_with_db_server(manager, "docs-prod-1")
|
||||
await self._reload_with_db_server(manager, "db-uuid-1")
|
||||
await self._reload_with_db_server(manager, "docs-prod-1")
|
||||
|
||||
assert len([m for m in caplog.messages if "database entry takes precedence" in m]) == 2
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pinned_id_matching_a_mapped_alias_is_rejected(self):
|
||||
"""An alias can also arrive from litellm_settings.mcp_aliases; it is reserved just the same."""
|
||||
manager = MCPServerManager()
|
||||
|
||||
with pytest.raises(ValueError, match="server_name or alias of MCP server 'wiki_server'"):
|
||||
await manager.load_servers_from_config(
|
||||
{
|
||||
"wiki_server": {"url": "https://wiki.example.com/mcp", "transport": MCPTransport.http},
|
||||
"docs_server": {
|
||||
"server_id": "wiki",
|
||||
"url": "https://example.com/mcp",
|
||||
"transport": MCPTransport.http,
|
||||
},
|
||||
},
|
||||
{"wiki": "wiki_server"},
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pinning_a_servers_own_mapped_alias_is_allowed(self):
|
||||
manager = MCPServerManager()
|
||||
|
||||
await manager.load_servers_from_config(
|
||||
self._config(server_id="docs"),
|
||||
{"docs": "docs_server"},
|
||||
)
|
||||
|
||||
assert list(manager.config_mcp_servers) == ["docs"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mapped_alias_for_an_unknown_server_reserves_nothing(self):
|
||||
"""A dangling mcp_aliases entry is never applied, so it must not fail an unrelated pin."""
|
||||
manager = MCPServerManager()
|
||||
|
||||
await manager.load_servers_from_config(
|
||||
self._config(server_id="wiki"),
|
||||
{"wiki": "a_server_that_does_not_exist"},
|
||||
)
|
||||
|
||||
assert list(manager.config_mcp_servers) == ["wiki"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_config_id_that_is_a_db_server_name_warns(self, caplog):
|
||||
"""The mirror of the shadow case: here the config entry captures the db server's grants."""
|
||||
manager = MCPServerManager()
|
||||
await manager.load_servers_from_config(self._config(server_id="db_server"))
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
|
||||
await self._reload_with_db_server(manager, "db-uuid-1")
|
||||
|
||||
assert any("db_server" in m and "name or alias of a database-backed" in m for m in caplog.messages)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_capture_warning_is_not_repeated_on_every_reload(self, caplog):
|
||||
manager = MCPServerManager()
|
||||
await manager.load_servers_from_config(self._config(server_id="db_server"))
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
|
||||
await self._reload_with_db_server(manager, "db-uuid-1")
|
||||
await self._reload_with_db_server(manager, "db-uuid-1")
|
||||
|
||||
assert len([m for m in caplog.messages if "name or alias of a database-backed" in m]) == 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_config_id_unrelated_to_db_names_does_not_warn(self, caplog):
|
||||
manager = MCPServerManager()
|
||||
await manager.load_servers_from_config(self._config(server_id="docs-prod-1"))
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
|
||||
await self._reload_with_db_server(manager, "db-uuid-1")
|
||||
|
||||
assert all("name or alias of a database-backed" not in m for m in caplog.messages)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mapped_alias_for_a_server_with_its_own_alias_reserves_nothing(self):
|
||||
"""load_servers_from_config ignores the mapping when the entry sets alias, so it is free."""
|
||||
manager = MCPServerManager()
|
||||
|
||||
await manager.load_servers_from_config(
|
||||
{
|
||||
"wiki_server": {
|
||||
"alias": "wiki_prod",
|
||||
"url": "https://wiki.example.com/mcp",
|
||||
"transport": MCPTransport.http,
|
||||
},
|
||||
"docs_server": {
|
||||
"server_id": "wiki",
|
||||
"url": "https://example.com/mcp",
|
||||
"transport": MCPTransport.http,
|
||||
},
|
||||
},
|
||||
{"wiki": "wiki_server"},
|
||||
)
|
||||
|
||||
assert "wiki" in manager.config_mcp_servers
|
||||
assert len(manager.config_mcp_servers) == 2
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_only_the_first_mapped_alias_for_a_server_is_reserved(self):
|
||||
"""Only the first mapping is applied, so pinning the second one must still load."""
|
||||
manager = MCPServerManager()
|
||||
|
||||
await manager.load_servers_from_config(
|
||||
{
|
||||
"wiki_server": {"url": "https://wiki.example.com/mcp", "transport": MCPTransport.http},
|
||||
"docs_server": {
|
||||
"server_id": "wiki_two",
|
||||
"url": "https://example.com/mcp",
|
||||
"transport": MCPTransport.http,
|
||||
},
|
||||
},
|
||||
{"wiki_one": "wiki_server", "wiki_two": "wiki_server"},
|
||||
)
|
||||
|
||||
assert "wiki_two" in manager.config_mcp_servers
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invalid_name_is_reported_before_any_entry_body_is_read(self):
|
||||
"""The identifier index walks every entry up front, so a bad name must still fail on the name."""
|
||||
with pytest.raises(Exception, match="Server name cannot contain"):
|
||||
await MCPServerManager().load_servers_from_config({"my-server": None})
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_shadowing_db_server_reports_only_the_shadow_warning(self, caplog):
|
||||
"""The db row wins the id outright, so the capture message would contradict the shadow one."""
|
||||
manager = MCPServerManager()
|
||||
await manager.load_servers_from_config(self._config(server_id="db_server"))
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
|
||||
await self._reload_with_db_server(manager, "db_server")
|
||||
|
||||
assert any("database entry takes precedence" in m for m in caplog.messages)
|
||||
assert all("name or alias of a database-backed" not in m for m in caplog.messages)
|
||||
assert manager.get_registry()["db_server"].url == "https://db.example.com/mcp"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_an_explicitly_blank_alias_still_blocks_the_mapping(self):
|
||||
"""The loader only consults mcp_aliases when the key is absent, so a blank alias frees it."""
|
||||
manager = MCPServerManager()
|
||||
|
||||
await manager.load_servers_from_config(
|
||||
{
|
||||
"wiki_server": {
|
||||
"alias": "",
|
||||
"url": "https://wiki.example.com/mcp",
|
||||
"transport": MCPTransport.http,
|
||||
},
|
||||
"docs_server": {
|
||||
"server_id": "wiki",
|
||||
"url": "https://example.com/mcp",
|
||||
"transport": MCPTransport.http,
|
||||
},
|
||||
},
|
||||
{"wiki": "wiki_server"},
|
||||
)
|
||||
|
||||
assert "wiki" in manager.config_mcp_servers
|
||||
assert manager.config_mcp_servers["wiki"].url == "https://example.com/mcp"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_row_that_shadows_one_id_still_reports_capturing_another(self, caplog):
|
||||
"""Skipping is per identifier, not per row, so the second collision is not lost."""
|
||||
manager = MCPServerManager()
|
||||
await manager.load_servers_from_config(
|
||||
{
|
||||
"docs_server": {
|
||||
"server_id": "shadow_x",
|
||||
"url": "https://example.com/mcp",
|
||||
"transport": MCPTransport.http,
|
||||
},
|
||||
"wiki_server": {
|
||||
"server_id": "capture_y",
|
||||
"url": "https://wiki.example.com/mcp",
|
||||
"transport": MCPTransport.http,
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
|
||||
await self._reload_with_db_server(manager, "shadow_x", db_name="capture_y")
|
||||
|
||||
assert any("shadow_x" in m and "database entry takes precedence" in m for m in caplog.messages)
|
||||
assert any("capture_y" in m and "name or alias of a database-backed" in m for m in caplog.messages)
|
||||
|
||||
|
||||
class TestLitellmAdmissionKeyIsNeverTheSubjectToken:
|
||||
"""The bearer that admitted the request as a LiteLLM key must not be sent to the IdP as the
|
||||
RFC 8693 subject_token (or ID-JAG assertion). Only ``x-litellm-api-key`` disambiguates: with it
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -29,6 +29,7 @@ added to this layer raises instead of silently passing - the inventory of seams
|
|||
cannot drift without a test failure.
|
||||
"""
|
||||
|
||||
import base64
|
||||
import json
|
||||
from contextlib import ExitStack
|
||||
from dataclasses import dataclass
|
||||
|
|
@ -36,7 +37,7 @@ from typing import Any, Dict, Optional
|
|||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles
|
||||
|
||||
import litellm
|
||||
import litellm.proxy.batches_endpoints.endpoints as endpoints
|
||||
|
|
@ -989,6 +990,67 @@ async def test_create__uses_acreate_batch_route_type(harness, openai_env_creds):
|
|||
assert harness.pre_call.call_args.kwargs["route_type"] == "acreate_batch"
|
||||
|
||||
|
||||
def install_managed_files_hook(harness: Harness) -> AsyncMock:
|
||||
prisma_client = AsyncMock()
|
||||
managed_files = _PROXY_LiteLLMManagedFiles(MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client)
|
||||
harness.logging.post_call_success_hook = AsyncMock(side_effect=managed_files.async_post_call_success_hook)
|
||||
harness.router.model_list = []
|
||||
return prisma_client
|
||||
|
||||
|
||||
TEAM_A_KEY = UserAPIKeyAuth(api_key="sk-team-a", user_id="user_a", team_id="team_a")
|
||||
|
||||
|
||||
def assert_ownership_registered_for_team_a(prisma_client: AsyncMock, batch_id: str) -> None:
|
||||
upsert = prisma_client.db.litellm_managedobjecttable.upsert
|
||||
upsert.assert_awaited_once()
|
||||
assert upsert.await_args.kwargs["where"] == {"unified_object_id": batch_id}
|
||||
created = upsert.await_args.kwargs["data"]["create"]
|
||||
assert created["created_by"] == "user_a"
|
||||
assert created["team_id"] == "team_a"
|
||||
prisma_client.db.litellm_managedobjecttable.update_many.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"body",
|
||||
[
|
||||
{"input_file_id": AZURE_FILE_ID},
|
||||
{"input_file_id": "file-plain", "model": "vertex-model"},
|
||||
{"input_file_id": "file-plain"},
|
||||
],
|
||||
ids=["model_encoded_file_id", "model_param", "provider_fallback"],
|
||||
)
|
||||
async def test_create__registers_ownership_for_creator(harness, openai_env_creds, body):
|
||||
set_body(harness, {**body, "endpoint": "/v1/chat/completions", "completion_window": "24h"})
|
||||
prisma_client = install_managed_files_hook(harness)
|
||||
|
||||
resp = await call_create(harness, user=TEAM_A_KEY)
|
||||
|
||||
assert_ownership_registered_for_team_a(prisma_client, resp.id)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create__unified_file_id_registers_ownership_for_creator(harness):
|
||||
unified_input_file_id = base64.urlsafe_b64encode(
|
||||
b"litellm_proxy:application/octet-stream;unified_id,input-uuid;target_model_names,gpt-4o-mini"
|
||||
).decode()
|
||||
set_body(
|
||||
harness,
|
||||
{
|
||||
"input_file_id": unified_input_file_id,
|
||||
"endpoint": "/v1/chat/completions",
|
||||
"completion_window": "24h",
|
||||
},
|
||||
)
|
||||
prisma_client = install_managed_files_hook(harness)
|
||||
|
||||
resp = await call_create(harness, user=TEAM_A_KEY)
|
||||
|
||||
assert harness.router_acreate.call_count == 1
|
||||
assert_ownership_registered_for_team_a(prisma_client, resp.id)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create__metadata_sanitized_before_forwarding(harness, openai_env_creds):
|
||||
set_body(
|
||||
|
|
|
|||
|
|
@ -1,16 +1,39 @@
|
|||
"""OpenAI passthrough must register WebSocket catch-all routes (#36088)."""
|
||||
"""OpenAI passthrough WebSocket route: registration, opt-in gating, and refusals."""
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
import json
|
||||
from collections.abc import Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType, SimpleNamespace
|
||||
from typing import Final
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
from starlette.routing import WebSocketRoute
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
|
||||
_OPENAI_WS_DISABLED_REFUSAL,
|
||||
_OPENAI_WS_MODEL_RESTRICTED_REFUSAL,
|
||||
_has_model_restrictions,
|
||||
_openai_websocket_refusal,
|
||||
_proxy_model_allowlists,
|
||||
openai_websocket_proxy_route,
|
||||
router,
|
||||
)
|
||||
|
||||
Scopes = tuple[Sequence[str], ...]
|
||||
|
||||
ENABLED: Final = MappingProxyType({"enable_openai_websocket_passthrough": True})
|
||||
DISABLED_SETTINGS: Final = (
|
||||
MappingProxyType({}),
|
||||
MappingProxyType({"enable_openai_websocket_passthrough": False}),
|
||||
MappingProxyType({"enable_openai_websocket_passthrough": "false"}),
|
||||
MappingProxyType({"enable_openai_websocket_passthrough": None}),
|
||||
)
|
||||
GET_CREDENTIALS: Final = (
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials"
|
||||
)
|
||||
|
||||
|
||||
def test_openai_websocket_passthrough_routes_registered():
|
||||
ws_paths = {route.path for route in router.routes if isinstance(route, WebSocketRoute)}
|
||||
|
|
@ -18,164 +41,258 @@ def test_openai_websocket_passthrough_routes_registered():
|
|||
assert "/openai_passthrough/{endpoint:path}" in ws_paths
|
||||
|
||||
|
||||
def _mock_websocket(path: str, query: str, headers: dict[str, str] | None = None) -> MagicMock:
|
||||
websocket = MagicMock()
|
||||
websocket.url.path = path
|
||||
websocket.url.query = query
|
||||
websocket.headers = headers or {}
|
||||
websocket.accept = AsyncMock()
|
||||
websocket.close = AsyncMock()
|
||||
return websocket
|
||||
class _FakeWebSocket:
|
||||
def __init__(self, path: str, query: str, subprotocols: str | None = None) -> None:
|
||||
self.url = SimpleNamespace(path=path, query=query)
|
||||
self.headers = {"sec-websocket-protocol": subprotocols} if subprotocols else {}
|
||||
self.accepts: list[str | None] = []
|
||||
self.sent: list[str] = []
|
||||
self.closed: tuple[int, str] | None = None
|
||||
|
||||
async def accept(self, subprotocol: str | None = None) -> None:
|
||||
self.accepts.append(subprotocol)
|
||||
|
||||
async def send_text(self, data: str) -> None:
|
||||
self.sent.append(data)
|
||||
|
||||
async def close(self, code: int = 1000, reason: str = "") -> None:
|
||||
self.closed = (code, reason)
|
||||
|
||||
def error_message(self) -> str:
|
||||
assert len(self.sent) == 1
|
||||
frame = json.loads(self.sent[0])
|
||||
assert frame["type"] == "error"
|
||||
return frame["error"]["message"]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _RelayCall:
|
||||
target: str
|
||||
custom_headers: Mapping[str, str]
|
||||
forward_headers: bool
|
||||
endpoint: str
|
||||
accept_websocket: bool
|
||||
|
||||
|
||||
class _FakeRelay:
|
||||
def __init__(self) -> None:
|
||||
self.calls: list[_RelayCall] = []
|
||||
|
||||
async def __call__(
|
||||
self,
|
||||
*,
|
||||
websocket: _FakeWebSocket,
|
||||
target: str,
|
||||
custom_headers: dict[str, str],
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
forward_headers: bool,
|
||||
endpoint: str,
|
||||
accept_websocket: bool,
|
||||
) -> None:
|
||||
self.calls.append(
|
||||
_RelayCall(
|
||||
target=target,
|
||||
custom_headers=MappingProxyType(dict(custom_headers)),
|
||||
forward_headers=forward_headers,
|
||||
endpoint=endpoint,
|
||||
accept_websocket=accept_websocket,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
class _FakeModelAllowlists:
|
||||
def __init__(self, scopes: Scopes) -> None:
|
||||
self.scopes = scopes
|
||||
self.calls: list[UserAPIKeyAuth] = []
|
||||
|
||||
async def __call__(self, valid_token: UserAPIKeyAuth, /) -> Scopes:
|
||||
self.calls.append(valid_token)
|
||||
return self.scopes
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Served:
|
||||
relay: _FakeRelay
|
||||
allowlists: _FakeModelAllowlists
|
||||
|
||||
|
||||
async def _serve(
|
||||
websocket: _FakeWebSocket,
|
||||
endpoint: str,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
general_settings: Mapping[str, object],
|
||||
scopes: Scopes = (),
|
||||
) -> _Served:
|
||||
served = _Served(relay=_FakeRelay(), allowlists=_FakeModelAllowlists(scopes))
|
||||
await openai_websocket_proxy_route(
|
||||
websocket=websocket,
|
||||
endpoint=endpoint,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
general_settings=general_settings,
|
||||
relay=served.relay,
|
||||
model_allowlists=served.allowlists,
|
||||
)
|
||||
return served
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("prefix", ["openai", "openai_passthrough"])
|
||||
async def test_openai_websocket_forwards_query_and_keeps_provider_auth(prefix):
|
||||
websocket = _mock_websocket(f"/{prefix}/v1/realtime", "model=gpt-4o-realtime-preview")
|
||||
async def test_openai_websocket_forwards_query_and_keeps_provider_auth(prefix, monkeypatch):
|
||||
monkeypatch.delenv("OPENAI_API_BASE", raising=False)
|
||||
websocket = _FakeWebSocket(f"/{prefix}/v1/realtime", "model=gpt-4o-realtime-preview")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
|
||||
return_value="sk-provider",
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._join_url_paths",
|
||||
return_value="https://api.openai.com/v1/realtime",
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.websocket_passthrough_request",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_ws,
|
||||
):
|
||||
await openai_websocket_proxy_route(
|
||||
websocket=websocket,
|
||||
endpoint="v1/realtime",
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
with patch(GET_CREDENTIALS, return_value="sk-provider"):
|
||||
served = await _serve(websocket, "v1/realtime", UserAPIKeyAuth(), ENABLED)
|
||||
|
||||
assert served.relay.calls == [
|
||||
_RelayCall(
|
||||
target="wss://api.openai.com/v1/realtime?model=gpt-4o-realtime-preview",
|
||||
custom_headers=MappingProxyType({"Authorization": "Bearer sk-provider"}),
|
||||
forward_headers=False,
|
||||
endpoint=f"/{prefix}/v1/realtime",
|
||||
accept_websocket=False,
|
||||
)
|
||||
|
||||
kwargs = mock_ws.await_args.kwargs
|
||||
assert kwargs["target"] == "wss://api.openai.com/v1/realtime?model=gpt-4o-realtime-preview"
|
||||
assert kwargs["custom_headers"] == {"Authorization": "Bearer sk-provider"}
|
||||
assert kwargs["forward_headers"] is False
|
||||
assert kwargs["endpoint"] == f"/{prefix}/v1/realtime"
|
||||
assert kwargs["accept_websocket"] is False
|
||||
websocket.accept.assert_awaited_once_with(subprotocol=None)
|
||||
websocket.close.assert_not_awaited()
|
||||
]
|
||||
assert websocket.accepts == [None]
|
||||
assert websocket.sent == []
|
||||
assert websocket.closed is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_websocket_accepts_first_client_subprotocol():
|
||||
websocket = _mock_websocket(
|
||||
websocket = _FakeWebSocket(
|
||||
"/openai/v1/realtime",
|
||||
"model=gpt-4o-realtime-preview",
|
||||
headers={
|
||||
"sec-websocket-protocol": "realtime, openai-insecure-api-key.sk-abc, openai-beta.realtime-v1"
|
||||
},
|
||||
subprotocols="realtime, openai-insecure-api-key.sk-abc, openai-beta.realtime-v1",
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
|
||||
return_value="sk-provider",
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.websocket_passthrough_request",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_ws,
|
||||
):
|
||||
await openai_websocket_proxy_route(
|
||||
websocket=websocket,
|
||||
endpoint="v1/realtime",
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
)
|
||||
with patch(GET_CREDENTIALS, return_value="sk-provider"):
|
||||
served = await _serve(websocket, "v1/realtime", UserAPIKeyAuth(), ENABLED)
|
||||
|
||||
websocket.accept.assert_awaited_once_with(subprotocol="realtime")
|
||||
assert mock_ws.await_args.kwargs["accept_websocket"] is False
|
||||
websocket.close.assert_not_awaited()
|
||||
assert websocket.accepts == ["realtime"]
|
||||
assert [call.accept_websocket for call in served.relay.calls] == [False]
|
||||
assert websocket.closed is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_websocket_closes_cleanly_when_provider_credentials_missing():
|
||||
websocket = _mock_websocket("/openai/v1/realtime", "model=gpt-4o-realtime-preview")
|
||||
websocket = _FakeWebSocket("/openai/v1/realtime", "model=gpt-4o-realtime-preview")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
|
||||
return_value=None,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.websocket_passthrough_request",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_ws,
|
||||
):
|
||||
await openai_websocket_proxy_route(
|
||||
websocket=websocket,
|
||||
endpoint="v1/realtime",
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
)
|
||||
with patch(GET_CREDENTIALS, return_value=None):
|
||||
served = await _serve(websocket, "v1/realtime", UserAPIKeyAuth(), ENABLED)
|
||||
|
||||
websocket.close.assert_awaited_once()
|
||||
assert websocket.close.await_args.kwargs["code"] == 1011
|
||||
websocket.accept.assert_not_awaited()
|
||||
mock_ws.assert_not_awaited()
|
||||
assert websocket.closed is not None
|
||||
assert websocket.closed[0] == 1011
|
||||
assert "OPENAI_API_KEY" in websocket.closed[1]
|
||||
assert websocket.accepts == []
|
||||
assert served.relay.calls == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"user_api_key_dict",
|
||||
[
|
||||
UserAPIKeyAuth(models=["gpt-4o"]),
|
||||
UserAPIKeyAuth(team_models=["gpt-4o-realtime-preview"]),
|
||||
UserAPIKeyAuth(models=["all-team-models"], team_models=["gpt-4o"]),
|
||||
],
|
||||
)
|
||||
async def test_openai_websocket_rejects_model_restricted_keys(user_api_key_dict):
|
||||
websocket = _mock_websocket("/openai/v1/realtime", "model=gpt-4o-realtime-preview")
|
||||
@pytest.mark.parametrize("prefix", ["openai", "openai_passthrough"])
|
||||
@pytest.mark.parametrize("general_settings", DISABLED_SETTINGS)
|
||||
async def test_openai_websocket_refused_unless_explicitly_enabled(prefix, general_settings):
|
||||
websocket = _FakeWebSocket(f"/{prefix}/v1/realtime", "model=gpt-4o-realtime-preview")
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.websocket_passthrough_request",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_ws:
|
||||
await openai_websocket_proxy_route(
|
||||
websocket=websocket,
|
||||
endpoint="v1/realtime",
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
served = await _serve(websocket, "v1/realtime", UserAPIKeyAuth(), general_settings)
|
||||
|
||||
websocket.close.assert_awaited_once()
|
||||
assert websocket.close.await_args.kwargs["code"] == 1008
|
||||
websocket.accept.assert_not_awaited()
|
||||
mock_ws.assert_not_awaited()
|
||||
assert "enable_openai_websocket_passthrough" in websocket.error_message()
|
||||
assert websocket.accepts == [None]
|
||||
assert websocket.closed == (1008, _OPENAI_WS_DISABLED_REFUSAL.close_reason)
|
||||
assert served.relay.calls == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"user_api_key_dict",
|
||||
[
|
||||
UserAPIKeyAuth(),
|
||||
UserAPIKeyAuth(models=["all-proxy-models"]),
|
||||
UserAPIKeyAuth(models=["*"]),
|
||||
UserAPIKeyAuth(models=["all-team-models"], team_models=["all-proxy-models"]),
|
||||
],
|
||||
@pytest.mark.parametrize("general_settings", DISABLED_SETTINGS)
|
||||
async def test_openai_websocket_refusal_is_disabled_for_falsy_settings(general_settings):
|
||||
refusal = await _openai_websocket_refusal(UserAPIKeyAuth(), general_settings, _FakeModelAllowlists(()))
|
||||
assert refusal is _OPENAI_WS_DISABLED_REFUSAL
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("value", [True, "true", "True"])
|
||||
async def test_openai_websocket_refusal_is_none_for_truthy_settings(value):
|
||||
settings = MappingProxyType({"enable_openai_websocket_passthrough": value})
|
||||
assert await _openai_websocket_refusal(UserAPIKeyAuth(), settings, _FakeModelAllowlists(())) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_websocket_refusal_echoes_requested_subprotocol():
|
||||
websocket = _FakeWebSocket(
|
||||
"/openai_passthrough/v1/realtime",
|
||||
"model=gpt-4o-realtime-preview",
|
||||
subprotocols="realtime, openai-beta.realtime-v1",
|
||||
)
|
||||
|
||||
served = await _serve(websocket, "v1/realtime", UserAPIKeyAuth(), MappingProxyType({}))
|
||||
|
||||
assert websocket.accepts == ["realtime"]
|
||||
assert websocket.closed == (1008, _OPENAI_WS_DISABLED_REFUSAL.close_reason)
|
||||
assert served.relay.calls == []
|
||||
|
||||
|
||||
RESTRICTED_SCOPES: Final[tuple[Scopes, ...]] = (
|
||||
(("gpt-4o",),),
|
||||
((), ("gpt-4o-realtime-preview",)),
|
||||
(("all-team-models",), ("gpt-4o",)),
|
||||
((), ("all-proxy-models",), ("gpt-4o",)),
|
||||
((), (), (), ("gpt-4o",)),
|
||||
(("*",), (), (), (), ("gpt-4o",)),
|
||||
)
|
||||
UNRESTRICTED_SCOPES: Final[tuple[Scopes, ...]] = (
|
||||
(),
|
||||
((),),
|
||||
(("all-proxy-models",),),
|
||||
(("*",),),
|
||||
(("all-team-models",), ("all-proxy-models",)),
|
||||
((), (), (), (), ()),
|
||||
(("*",), ("all-proxy-models",), ("all-team-models",), (), ()),
|
||||
)
|
||||
async def test_openai_websocket_allows_unrestricted_keys(user_api_key_dict):
|
||||
websocket = _mock_websocket("/openai/v1/responses", "")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
|
||||
return_value="sk-provider",
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.websocket_passthrough_request",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_ws,
|
||||
):
|
||||
await openai_websocket_proxy_route(
|
||||
websocket=websocket,
|
||||
endpoint="v1/responses",
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
mock_ws.assert_awaited_once()
|
||||
websocket.close.assert_not_awaited()
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("scopes", RESTRICTED_SCOPES)
|
||||
async def test_openai_websocket_rejects_model_restricted_identities(scopes):
|
||||
websocket = _FakeWebSocket("/openai/v1/realtime", "model=gpt-4o-realtime-preview")
|
||||
user_api_key_dict = UserAPIKeyAuth(token="hashed-fake", user_id="user-fake", team_id="team-fake")
|
||||
|
||||
served = await _serve(websocket, "v1/realtime", user_api_key_dict, ENABLED, scopes)
|
||||
|
||||
assert "model restrictions" in websocket.error_message()
|
||||
assert websocket.closed == (1008, _OPENAI_WS_MODEL_RESTRICTED_REFUSAL.close_reason)
|
||||
assert served.relay.calls == []
|
||||
assert served.allowlists.calls == [user_api_key_dict]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("scopes", RESTRICTED_SCOPES)
|
||||
async def test_openai_websocket_disabled_refusal_skips_allowlist_lookups(scopes):
|
||||
allowlists = _FakeModelAllowlists(scopes)
|
||||
|
||||
refusal = await _openai_websocket_refusal(UserAPIKeyAuth(), MappingProxyType({}), allowlists)
|
||||
|
||||
assert refusal is _OPENAI_WS_DISABLED_REFUSAL
|
||||
assert allowlists.calls == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("scopes", UNRESTRICTED_SCOPES)
|
||||
async def test_openai_websocket_allows_unrestricted_identities(scopes):
|
||||
websocket = _FakeWebSocket("/openai/v1/responses", "")
|
||||
|
||||
with patch(GET_CREDENTIALS, return_value="sk-provider"):
|
||||
served = await _serve(websocket, "v1/responses", UserAPIKeyAuth(), ENABLED, scopes)
|
||||
|
||||
assert len(served.relay.calls) == 1
|
||||
assert websocket.sent == []
|
||||
assert websocket.closed is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_proxy_model_allowlists_reads_the_token_scopes_without_a_database():
|
||||
token: Final = UserAPIKeyAuth(models=[], team_id="team-fake", team_models=["gpt-4o"])
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", None):
|
||||
scopes = await _proxy_model_allowlists()(token)
|
||||
|
||||
assert tuple(tuple(scope) for scope in scopes) == ((), ("gpt-4o",))
|
||||
assert _has_model_restrictions(scopes)
|
||||
|
|
|
|||
|
|
@ -11568,14 +11568,10 @@ async def test_key_window_spend_row_is_enqueued_with_the_actual_cost():
|
|||
|
||||
reset_at = datetime.now(timezone.utc) + timedelta(days=10)
|
||||
key_obj = MagicMock()
|
||||
key_obj.budget_limits = [
|
||||
{"budget_duration": "30d", "max_budget": 100.0, "reset_at": reset_at.isoformat()}
|
||||
]
|
||||
key_obj.budget_limits = [{"budget_duration": "30d", "max_budget": 100.0, "reset_at": reset_at.isoformat()}]
|
||||
|
||||
with _window_spend_enqueue_env({"hashed-token": key_obj}) as queue:
|
||||
await increment_spend_counters(
|
||||
token="hashed-token", team_id=None, user_id=None, response_cost=0.25
|
||||
)
|
||||
await increment_spend_counters(token="hashed-token", team_id=None, user_id=None, response_cost=0.25)
|
||||
enqueued = await _drain(queue)
|
||||
|
||||
assert len(enqueued) == 1
|
||||
|
|
@ -11594,14 +11590,10 @@ async def test_team_window_spend_row_is_enqueued():
|
|||
|
||||
reset_at = datetime.now(timezone.utc) + timedelta(days=3)
|
||||
team_obj = MagicMock()
|
||||
team_obj.budget_limits = [
|
||||
{"budget_duration": "7d", "max_budget": 50.0, "reset_at": reset_at.isoformat()}
|
||||
]
|
||||
team_obj.budget_limits = [{"budget_duration": "7d", "max_budget": 50.0, "reset_at": reset_at.isoformat()}]
|
||||
|
||||
with _window_spend_enqueue_env({"team_id:team-1": team_obj}) as queue:
|
||||
await increment_spend_counters(
|
||||
token=None, team_id="team-1", user_id=None, response_cost=1.5
|
||||
)
|
||||
await increment_spend_counters(token=None, team_id="team-1", user_id=None, response_cost=1.5)
|
||||
enqueued = await _drain(queue)
|
||||
|
||||
assert len(enqueued) == 1
|
||||
|
|
@ -11620,9 +11612,7 @@ async def test_window_spend_row_is_enqueued_even_when_the_counter_was_reserved()
|
|||
|
||||
reset_at = datetime.now(timezone.utc) + timedelta(days=10)
|
||||
key_obj = MagicMock()
|
||||
key_obj.budget_limits = [
|
||||
{"budget_duration": "30d", "max_budget": 100.0, "reset_at": reset_at.isoformat()}
|
||||
]
|
||||
key_obj.budget_limits = [{"budget_duration": "30d", "max_budget": 100.0, "reset_at": reset_at.isoformat()}]
|
||||
reservation = {
|
||||
"entries": [
|
||||
{"counter_key": "spend:key:hashed-token", "reserved": 1.0},
|
||||
|
|
@ -11660,9 +11650,7 @@ async def test_sliding_window_without_reset_at_is_not_enqueued():
|
|||
key_obj.budget_limits = [{"budget_duration": "30d", "max_budget": 100.0}]
|
||||
|
||||
with _window_spend_enqueue_env({"hashed-token": key_obj}) as queue:
|
||||
await increment_spend_counters(
|
||||
token="hashed-token", team_id=None, user_id=None, response_cost=0.25
|
||||
)
|
||||
await increment_spend_counters(token="hashed-token", team_id=None, user_id=None, response_cost=0.25)
|
||||
enqueued = await _drain(queue)
|
||||
|
||||
assert enqueued == []
|
||||
|
|
@ -11680,9 +11668,7 @@ async def test_each_configured_window_gets_its_own_row_enqueue():
|
|||
]
|
||||
|
||||
with _window_spend_enqueue_env({"hashed-token": key_obj}) as queue:
|
||||
await increment_spend_counters(
|
||||
token="hashed-token", team_id=None, user_id=None, response_cost=0.25
|
||||
)
|
||||
await increment_spend_counters(token="hashed-token", team_id=None, user_id=None, response_cost=0.25)
|
||||
enqueued = await _drain(queue)
|
||||
|
||||
assert sorted(item["window_duration"] for item in enqueued) == ["1d", "30d"]
|
||||
|
|
@ -11697,9 +11683,7 @@ async def test_no_window_spend_row_enqueued_without_budget_limits():
|
|||
key_obj.budget_limits = None
|
||||
|
||||
with _window_spend_enqueue_env({"hashed-token": key_obj}) as queue:
|
||||
await increment_spend_counters(
|
||||
token="hashed-token", team_id=None, user_id=None, response_cost=0.25
|
||||
)
|
||||
await increment_spend_counters(token="hashed-token", team_id=None, user_id=None, response_cost=0.25)
|
||||
enqueued = await _drain(queue)
|
||||
|
||||
assert enqueued == []
|
||||
|
|
@ -11713,9 +11697,7 @@ async def test_window_spend_row_carries_the_request_start_time():
|
|||
|
||||
reset_at = datetime.now(timezone.utc) + timedelta(days=10)
|
||||
key_obj = MagicMock()
|
||||
key_obj.budget_limits = [
|
||||
{"budget_duration": "30d", "max_budget": 100.0, "reset_at": reset_at.isoformat()}
|
||||
]
|
||||
key_obj.budget_limits = [{"budget_duration": "30d", "max_budget": 100.0, "reset_at": reset_at.isoformat()}]
|
||||
|
||||
with _window_spend_enqueue_env({"hashed-token": key_obj}) as queue:
|
||||
await increment_spend_counters(
|
||||
|
|
@ -11736,9 +11718,7 @@ async def test_team_window_spend_row_carries_the_request_start_time():
|
|||
|
||||
reset_at = datetime.now(timezone.utc) + timedelta(days=3)
|
||||
team_obj = MagicMock()
|
||||
team_obj.budget_limits = [
|
||||
{"budget_duration": "7d", "max_budget": 50.0, "reset_at": reset_at.isoformat()}
|
||||
]
|
||||
team_obj.budget_limits = [{"budget_duration": "7d", "max_budget": 50.0, "reset_at": reset_at.isoformat()}]
|
||||
|
||||
with _window_spend_enqueue_env({"team_id:team-1": team_obj}) as queue:
|
||||
await increment_spend_counters(
|
||||
|
|
@ -12050,7 +12030,6 @@ async def test_init_guardrails_in_db_snapshots_and_reconciles_under_guardrail_re
|
|||
assert not GUARDRAIL_RECONCILE_LOCK.locked()
|
||||
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_init_prompts_in_db_reloads_rows_patched_on_another_worker(monkeypatch):
|
||||
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
|
||||
|
|
@ -12094,7 +12073,9 @@ 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 AHOY"
|
||||
|
||||
prisma_client.db.litellm_prompttable.find_many = AsyncMock(return_value=[db_row("Begin every reply with HOWDY")])
|
||||
prisma_client.db.litellm_prompttable.find_many = AsyncMock(
|
||||
return_value=[db_row("Begin every reply with HOWDY")]
|
||||
)
|
||||
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
|
||||
|
||||
assert served_content() == "Begin every reply with HOWDY"
|
||||
|
|
@ -12547,3 +12528,40 @@ def test_disabling_docs_does_not_disable_other_routes(monkeypatch):
|
|||
|
||||
assert client.get("/redoc").status_code == 404
|
||||
assert client.get("/health/liveliness").status_code == 200
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"db_general_settings, expected",
|
||||
[
|
||||
({"enable_openai_websocket_passthrough": True}, True),
|
||||
({"enable_openai_websocket_passthrough": False}, False),
|
||||
({}, None),
|
||||
],
|
||||
)
|
||||
async def test_update_general_settings_propagates_openai_websocket_passthrough(db_general_settings, expected):
|
||||
from litellm.proxy.proxy_server import ProxyConfig
|
||||
|
||||
proxy_config = ProxyConfig()
|
||||
|
||||
with patch("litellm.proxy.proxy_server.general_settings", {"enable_openai_websocket_passthrough": True}):
|
||||
await proxy_config._update_general_settings(db_general_settings=db_general_settings)
|
||||
|
||||
import litellm.proxy.proxy_server as ps
|
||||
|
||||
assert ps.general_settings["enable_openai_websocket_passthrough"] is expected
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_general_settings_keeps_yaml_openai_websocket_passthrough():
|
||||
from litellm.proxy.proxy_server import ProxyConfig
|
||||
|
||||
proxy_config = ProxyConfig()
|
||||
proxy_config._yaml_general_settings_keys = {"enable_openai_websocket_passthrough"}
|
||||
|
||||
with patch("litellm.proxy.proxy_server.general_settings", {"enable_openai_websocket_passthrough": False}):
|
||||
await proxy_config._update_general_settings(db_general_settings={"enable_openai_websocket_passthrough": True})
|
||||
|
||||
import litellm.proxy.proxy_server as ps
|
||||
|
||||
assert ps.general_settings["enable_openai_websocket_passthrough"] is False
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ import { RestrictedSection, restrictedBy } from "./TierRestrictions";
|
|||
import HeuristicScoringConfig from "./HeuristicScoringConfig";
|
||||
import ClassifierReasoningEffortSelect from "./ClassifierReasoningEffortSelect";
|
||||
import ClassifierCircuitBreakerConfig from "./ClassifierCircuitBreakerConfig";
|
||||
import ClassifierVisionConfig from "./ClassifierVisionConfig";
|
||||
import type { ReasoningEffort } from "./complexity_router_tiers";
|
||||
import { useComplexityScorerDefaults } from "@/app/(dashboard)/hooks/autoRouter/useComplexityScorerDefaults";
|
||||
import {
|
||||
|
|
@ -315,12 +316,13 @@ const ClassificationMethodConfig: React.FC<ClassificationMethodConfigProps> = ({
|
|||
timeout_ms: value.classifier_llm_config?.timeout_ms ?? DEFAULT_CLASSIFIER_TIMEOUT_MS,
|
||||
classification_rubric: selectedRubric,
|
||||
};
|
||||
onChange({
|
||||
const nextValue: ComplexityRouterConfigValue = {
|
||||
...value,
|
||||
...(selectedRubric && { classifier_llm_config: rubricConfig }),
|
||||
classification_prompt: classificationPrompt,
|
||||
classification_examples: classificationExamples,
|
||||
});
|
||||
};
|
||||
onChange(nextValue);
|
||||
};
|
||||
|
||||
const handleClassifierModelChange = (model: string) => {
|
||||
|
|
@ -577,6 +579,10 @@ const ClassificationMethodConfig: React.FC<ClassificationMethodConfigProps> = ({
|
|||
value={value.classifier_llm_config ?? { model: "", timeout_ms: DEFAULT_CLASSIFIER_TIMEOUT_MS }}
|
||||
onChange={(classifier_llm_config) => onChange({ ...value, classifier_llm_config })}
|
||||
/>
|
||||
<ClassifierVisionConfig
|
||||
value={value.classifier_llm_config ?? { model: "", timeout_ms: DEFAULT_CLASSIFIER_TIMEOUT_MS }}
|
||||
onChange={(classifier_llm_config) => onChange({ ...value, classifier_llm_config })}
|
||||
/>
|
||||
<div>
|
||||
<div className="flex items-center gap-2 mb-1">
|
||||
<strong className="font-semibold">Classifier Prompt</strong>
|
||||
|
|
|
|||
|
|
@ -0,0 +1,79 @@
|
|||
import { Input } from "@/components/ui/input";
|
||||
import { Label } from "@/components/ui/label";
|
||||
import { Switch } from "@/components/ui/switch";
|
||||
import React from "react";
|
||||
|
||||
import type { ClassifierLLMConfigWire } from "./build_complexity_router_config";
|
||||
|
||||
export const DEFAULT_CLASSIFIER_VISION_ENABLED = false;
|
||||
export const DEFAULT_CLASSIFIER_VISION_MAX_IMAGES = 1;
|
||||
|
||||
const MAX_IMAGES_ID = "classifier-vision-max-images";
|
||||
|
||||
interface ClassifierVisionConfigProps {
|
||||
value: ClassifierLLMConfigWire;
|
||||
onChange: (value: ClassifierLLMConfigWire) => void;
|
||||
}
|
||||
|
||||
const ClassifierVisionConfig: React.FC<ClassifierVisionConfigProps> = ({ value, onChange }) => {
|
||||
const [draftMaxImages, setDraftMaxImages] = React.useState<string | null>(null);
|
||||
const enabled = value.vision?.enabled ?? DEFAULT_CLASSIFIER_VISION_ENABLED;
|
||||
|
||||
const handleMaxImagesChange = (raw: string): void => {
|
||||
setDraftMaxImages(raw);
|
||||
const parsed = Number(raw);
|
||||
if (raw.trim() === "" || !Number.isFinite(parsed)) return;
|
||||
onChange({
|
||||
...value,
|
||||
vision: { ...value.vision, enabled, max_images: Math.max(1, Math.round(parsed)) },
|
||||
});
|
||||
};
|
||||
|
||||
return (
|
||||
<div className="space-y-2 rounded-md border border-border p-3">
|
||||
<div className="flex items-center gap-2">
|
||||
<Switch
|
||||
checked={enabled}
|
||||
onCheckedChange={(visionEnabled): void => {
|
||||
if (!visionEnabled) {
|
||||
const { vision: _vision, ...withoutVision } = value;
|
||||
onChange(withoutVision);
|
||||
return;
|
||||
}
|
||||
onChange({
|
||||
...value,
|
||||
vision: {
|
||||
...value.vision,
|
||||
enabled: true,
|
||||
max_images: value.vision?.max_images ?? DEFAULT_CLASSIFIER_VISION_MAX_IMAGES,
|
||||
},
|
||||
});
|
||||
}}
|
||||
aria-label="Use images for classification"
|
||||
/>
|
||||
<strong className="font-semibold">Use images for classification</strong>
|
||||
</div>
|
||||
<span className="block text-xs text-muted-foreground">
|
||||
Send inline image data to the classifier so it can choose a tier from what the image shows.
|
||||
</span>
|
||||
{enabled && (
|
||||
<div>
|
||||
<Label htmlFor={MAX_IMAGES_ID} className="block mb-1 font-semibold">
|
||||
Maximum images per request
|
||||
</Label>
|
||||
<Input
|
||||
id={MAX_IMAGES_ID}
|
||||
type="text"
|
||||
inputMode="numeric"
|
||||
value={draftMaxImages ?? String(value.vision?.max_images ?? DEFAULT_CLASSIFIER_VISION_MAX_IMAGES)}
|
||||
onChange={(event) => handleMaxImagesChange(event.target.value)}
|
||||
onBlur={() => setDraftMaxImages(null)}
|
||||
className="w-full"
|
||||
/>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
||||
export default ClassifierVisionConfig;
|
||||
|
|
@ -1,5 +1,6 @@
|
|||
import { fireEvent, renderWithProviders, screen, within } from "../../../tests/test-utils";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import React from "react";
|
||||
import { vi } from "vitest";
|
||||
import ComplexityRouterConfig, { ComplexityRouterConfigValue } from "./ComplexityRouterConfig";
|
||||
vi.mock(
|
||||
|
|
@ -1690,3 +1691,83 @@ describe("ComplexityRouterConfig tier editing", () => {
|
|||
expect(screen.queryByText("Display names rename the built-in tiers", { exact: false })).not.toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
describe("classifier vision settings", () => {
|
||||
const llmValue: ComplexityRouterConfigValue = {
|
||||
...defaultValue,
|
||||
classifier_type: "llm",
|
||||
classifier_llm_config: { model: "gpt-3.5-turbo", timeout_ms: 3000 },
|
||||
};
|
||||
|
||||
const VisionFixture = ({ onChange = vi.fn() }: { onChange?: ReturnType<typeof vi.fn> }) => {
|
||||
const [value, setValue] = React.useState(llmValue);
|
||||
return (
|
||||
<ComplexityRouterConfig
|
||||
modelInfo={mockModelInfo}
|
||||
value={value}
|
||||
onChange={(nextValue) => {
|
||||
setValue(nextValue);
|
||||
onChange(nextValue);
|
||||
}}
|
||||
/>
|
||||
);
|
||||
};
|
||||
|
||||
it("starts off and reveals the default cap when enabled", () => {
|
||||
renderWithProviders(<VisionFixture />);
|
||||
fireEvent.click(screen.getByText("Advanced: Classification Method"));
|
||||
|
||||
const vision = screen.getByRole("switch", { name: "Use images for classification" });
|
||||
expect(vision).not.toBeChecked();
|
||||
expect(screen.queryByLabelText("Maximum images per request")).not.toBeInTheDocument();
|
||||
|
||||
fireEvent.click(vision);
|
||||
|
||||
expect(screen.getByLabelText("Maximum images per request")).toHaveValue("1");
|
||||
});
|
||||
|
||||
it("writes the switch and a clamped image cap into the classifier config", () => {
|
||||
const onChange = vi.fn();
|
||||
renderWithProviders(<VisionFixture onChange={onChange} />);
|
||||
fireEvent.click(screen.getByText("Advanced: Classification Method"));
|
||||
|
||||
fireEvent.click(screen.getByRole("switch", { name: "Use images for classification" }));
|
||||
expect(onChange).toHaveBeenLastCalledWith({
|
||||
...llmValue,
|
||||
classifier_llm_config: { ...llmValue.classifier_llm_config, vision: { enabled: true, max_images: 1 } },
|
||||
});
|
||||
|
||||
fireEvent.change(screen.getByLabelText("Maximum images per request"), { target: { value: "1.7" } });
|
||||
expect(onChange).toHaveBeenLastCalledWith({
|
||||
...llmValue,
|
||||
classifier_llm_config: { ...llmValue.classifier_llm_config, vision: { enabled: true, max_images: 2 } },
|
||||
});
|
||||
});
|
||||
|
||||
it("keeps the image cap draft empty until a valid value is entered", () => {
|
||||
const onChange = vi.fn();
|
||||
renderWithProviders(<VisionFixture onChange={onChange} />);
|
||||
fireEvent.click(screen.getByText("Advanced: Classification Method"));
|
||||
fireEvent.click(screen.getByRole("switch", { name: "Use images for classification" }));
|
||||
onChange.mockClear();
|
||||
|
||||
const input = screen.getByLabelText("Maximum images per request");
|
||||
fireEvent.change(input, { target: { value: "" } });
|
||||
|
||||
expect(input).toHaveValue("");
|
||||
expect(onChange).not.toHaveBeenCalled();
|
||||
|
||||
fireEvent.change(input, { target: { value: "0" } });
|
||||
expect(onChange).toHaveBeenLastCalledWith({
|
||||
...llmValue,
|
||||
classifier_llm_config: { ...llmValue.classifier_llm_config, vision: { enabled: true, max_images: 1 } },
|
||||
});
|
||||
});
|
||||
|
||||
it("is absent when the classifier is heuristic", () => {
|
||||
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultValue} onChange={vi.fn()} />);
|
||||
fireEvent.click(screen.getByText("Advanced: Classification Method"));
|
||||
|
||||
expect(screen.queryByText("Use images for classification")).not.toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -1170,12 +1170,13 @@ describe("buildComplexityRouterConfig stall escalation", () => {
|
|||
});
|
||||
|
||||
it("emits the toggle and both knobs when it is on", () => {
|
||||
const config = buildComplexityRouterConfig({
|
||||
const params = {
|
||||
...baseParams,
|
||||
stallEscalationEnabled: true,
|
||||
stallEscalationWindow: 8,
|
||||
stallEscalationRepeatThreshold: 4,
|
||||
});
|
||||
};
|
||||
const config = buildComplexityRouterConfig(params);
|
||||
expect(config.stall_escalation_enabled).toBe(true);
|
||||
expect(config.stall_escalation_window).toBe(8);
|
||||
expect(config.stall_escalation_repeat_threshold).toBe(4);
|
||||
|
|
@ -1207,3 +1208,38 @@ describe("dryRunRejection", () => {
|
|||
expect(dryRunRejection({ valid: true, error: null })).toBeNull();
|
||||
});
|
||||
});
|
||||
|
||||
describe("classifier vision wire payload", () => {
|
||||
const vision = { enabled: true, max_images: 3 };
|
||||
const classifierLlmConfig = { model: "classifier", timeout_ms: 3000, vision };
|
||||
|
||||
it("keeps vision through the standard-tier payload", () => {
|
||||
const params = { ...baseParams, classifierType: "llm" as const, classifierLlmConfig };
|
||||
const payload = buildComplexityRouterConfig(params);
|
||||
|
||||
expect(payload.classifier_llm_config).toMatchObject({ vision });
|
||||
});
|
||||
|
||||
it("keeps vision through the custom-tier payload", () => {
|
||||
const customTierSet = {
|
||||
tiers: [
|
||||
{ id: "simple", name: "simple", definition: "small talk", models: ["gpt-4o-mini"] },
|
||||
{ id: "complex", name: "complex", definition: "hard work", models: ["gpt-4o"] },
|
||||
],
|
||||
fallback_tier_id: "simple",
|
||||
};
|
||||
const payload = buildComplexityRouterConfig({ ...baseParams, customTierSet, classifierLlmConfig });
|
||||
|
||||
expect(payload.classifier_llm_config).toMatchObject({ vision });
|
||||
});
|
||||
|
||||
it("keeps an untouched classifier config free of vision", () => {
|
||||
const payload = buildComplexityRouterConfig({
|
||||
...baseParams,
|
||||
classifierType: "llm",
|
||||
classifierLlmConfig: { model: "classifier", timeout_ms: 3000 },
|
||||
});
|
||||
|
||||
expect(payload.classifier_llm_config).not.toHaveProperty("vision");
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -1,8 +1,5 @@
|
|||
import { KeywordTierRule } from "./KeywordTierRules";
|
||||
|
||||
type ClassifierLLMConfigWire = ClassifierLLMConfig & { vision?: { enabled?: boolean; max_images?: number } };
|
||||
|
||||
import type { ModelGroup } from "../llm_calls/fetch_models";
|
||||
import { KeywordTierRule } from "./KeywordTierRules";
|
||||
import {
|
||||
type CustomTierSet,
|
||||
type TierRow,
|
||||
|
|
@ -42,6 +39,9 @@ import {
|
|||
usesLlmClassifier,
|
||||
} from "./ComplexityRouterConfig";
|
||||
|
||||
export type ClassifierVisionConfig = { enabled?: boolean; max_images?: number };
|
||||
export type ClassifierLLMConfigWire = ClassifierLLMConfig & { vision?: ClassifierVisionConfig };
|
||||
|
||||
/**
|
||||
* Drop an empty system_prompt so the payload carries an override only when there is one. The
|
||||
* backend rejects a blank string rather than reading it as "use the default", and sending `""`
|
||||
|
|
@ -124,7 +124,7 @@ export interface BuildComplexityRouterConfigParams {
|
|||
planModeMinTier: string | undefined;
|
||||
tierLabels: ComplexityTierLabels | undefined;
|
||||
classifierType: ClassifierType;
|
||||
classifierLlmConfig: ClassifierLLMConfig | undefined;
|
||||
classifierLlmConfig: ClassifierLLMConfigWire | undefined;
|
||||
classifierContextWindowSize: number | undefined;
|
||||
classifierContextBudgetChars: number | undefined;
|
||||
classifierContextIncludeAssistantTurns: boolean | undefined;
|
||||
|
|
|
|||
|
|
@ -1036,3 +1036,60 @@ describe("EditAutoRouterModal with a stored custom tier set", () => {
|
|||
expect(savedConfig().tier_model_configs).toEqual(CUSTOM_STORED.tier_model_configs);
|
||||
});
|
||||
});
|
||||
|
||||
describe("EditAutoRouterModal classifier vision", () => {
|
||||
beforeEach(() => {
|
||||
modelPatchUpdateCall.mockClear();
|
||||
});
|
||||
|
||||
const STORED_CONFIG = {
|
||||
tiers: { SIMPLE: ["gpt-4o-mini"], MEDIUM: ["gpt-4o-mini"], COMPLEX: ["gpt-4o-mini"], REASONING: ["gpt-4o-mini"] },
|
||||
classifier_type: "llm",
|
||||
classifier_llm_config: {
|
||||
model: "gpt-4o-mini",
|
||||
timeout_ms: 3000,
|
||||
vision: { enabled: true, max_images: 2 },
|
||||
},
|
||||
};
|
||||
|
||||
const renderModal = () =>
|
||||
renderWithProviders(
|
||||
<EditAutoRouterModal
|
||||
isVisible
|
||||
onCancel={vi.fn()}
|
||||
onSuccess={vi.fn()}
|
||||
modelData={{
|
||||
...MODEL_DATA,
|
||||
litellm_params: { ...MODEL_DATA.litellm_params, complexity_router_config: STORED_CONFIG },
|
||||
}}
|
||||
accessToken="token"
|
||||
userRole="Admin"
|
||||
/>,
|
||||
);
|
||||
|
||||
it("hydrates and keeps a stored vision setting through an untouched save", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderModal();
|
||||
|
||||
await user.click(await screen.findByText("Advanced: Classification Method"));
|
||||
expect(screen.getByRole("switch", { name: "Use images for classification" })).toBeChecked();
|
||||
expect(screen.getByLabelText("Maximum images per request")).toHaveValue("2");
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /save changes/i }));
|
||||
|
||||
await waitFor(() => expect(modelPatchUpdateCall).toHaveBeenCalled());
|
||||
expect(savedConfig().classifier_llm_config).toMatchObject({ vision: { enabled: true, max_images: 2 } });
|
||||
});
|
||||
|
||||
it("removes vision when the operator turns it off", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderModal();
|
||||
|
||||
await user.click(await screen.findByText("Advanced: Classification Method"));
|
||||
await user.click(screen.getByRole("switch", { name: "Use images for classification" }));
|
||||
await user.click(screen.getByRole("button", { name: /save changes/i }));
|
||||
|
||||
await waitFor(() => expect(modelPatchUpdateCall).toHaveBeenCalled());
|
||||
expect(savedConfig().classifier_llm_config).not.toHaveProperty("vision");
|
||||
});
|
||||
});
|
||||
|
|
|
|||
5
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
5
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -25776,6 +25776,11 @@ export interface components {
|
|||
* @description If True and SSO is configured (MICROSOFT_CLIENT_ID, GOOGLE_CLIENT_ID, GENERIC_CLIENT_ID, or SAML_IDP_METADATA_URL/XML), disables username/password login on /login, /v2/login, and /v3/login so SSO is the only way to reach the Admin UI. An admin locked out of the UI can still administer the proxy over the API with the master key; unset this setting and restart the proxy to restore UI username/password login. Default is False.
|
||||
*/
|
||||
disable_password_login_when_sso_enabled?: boolean | null;
|
||||
/**
|
||||
* Enable Openai Websocket Passthrough
|
||||
* @description Serve the OpenAI pass-through WebSocket route, which relays frames to OpenAI under the proxy's own provider credential without reading them. Off by default.
|
||||
*/
|
||||
enable_openai_websocket_passthrough?: boolean | null;
|
||||
/**
|
||||
* Enable Public Model Hub
|
||||
* @description Public model hub for users to see what models they have access to, supported openai params, etc.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue