mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_batch_ui_logs
This commit is contained in:
commit
ed2408f28a
20 changed files with 1429 additions and 50 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,
|
||||
|
|
@ -1347,7 +1348,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"))
|
||||
|
||||
|
|
@ -1436,10 +1437,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 = ""
|
||||
|
|
|
|||
|
|
@ -77,4 +77,16 @@ spec:
|
|||
volumes:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
{{- with .Values.migrationJob.nodeSelector }}
|
||||
nodeSelector:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
{{- with .Values.migrationJob.affinity }}
|
||||
affinity:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
{{- with .Values.migrationJob.tolerations }}
|
||||
tolerations:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
{{- end }}
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
suite: test migrations Job ServiceAccount resolution and pod hardening
|
||||
suite: test migrations Job ServiceAccount resolution, pod hardening, and scheduling
|
||||
templates:
|
||||
- migrations-job.yaml
|
||||
values:
|
||||
|
|
@ -188,3 +188,69 @@ tests:
|
|||
asserts:
|
||||
- notExists:
|
||||
path: spec.activeDeadlineSeconds
|
||||
|
||||
- it: renders no scheduling fields by default
|
||||
asserts:
|
||||
- isNull:
|
||||
path: spec.template.spec.nodeSelector
|
||||
- isNull:
|
||||
path: spec.template.spec.tolerations
|
||||
- isNull:
|
||||
path: spec.template.spec.affinity
|
||||
|
||||
- it: renders nodeSelector, tolerations, and affinity from the migrationJob values
|
||||
set:
|
||||
migrationJob.nodeSelector:
|
||||
intent: no-csi-nodes
|
||||
migrationJob.tolerations:
|
||||
- key: intent
|
||||
operator: Equal
|
||||
value: no-csi-nodes
|
||||
effect: NoSchedule
|
||||
migrationJob.affinity:
|
||||
nodeAffinity:
|
||||
requiredDuringSchedulingIgnoredDuringExecution:
|
||||
nodeSelectorTerms:
|
||||
- matchExpressions:
|
||||
- key: intent
|
||||
operator: In
|
||||
values:
|
||||
- no-csi-nodes
|
||||
asserts:
|
||||
- equal:
|
||||
path: spec.template.spec.nodeSelector
|
||||
value:
|
||||
intent: no-csi-nodes
|
||||
- equal:
|
||||
path: spec.template.spec.tolerations
|
||||
value:
|
||||
- key: intent
|
||||
operator: Equal
|
||||
value: no-csi-nodes
|
||||
effect: NoSchedule
|
||||
- equal:
|
||||
path: spec.template.spec.affinity
|
||||
value:
|
||||
nodeAffinity:
|
||||
requiredDuringSchedulingIgnoredDuringExecution:
|
||||
nodeSelectorTerms:
|
||||
- matchExpressions:
|
||||
- key: intent
|
||||
operator: In
|
||||
values:
|
||||
- no-csi-nodes
|
||||
|
||||
- it: does not inherit the gateway's scheduling values
|
||||
set:
|
||||
gateway.nodeSelector:
|
||||
intent: no-csi-nodes
|
||||
gateway.tolerations:
|
||||
- key: intent
|
||||
operator: Equal
|
||||
value: no-csi-nodes
|
||||
effect: NoSchedule
|
||||
asserts:
|
||||
- isNull:
|
||||
path: spec.template.spec.nodeSelector
|
||||
- isNull:
|
||||
path: spec.template.spec.tolerations
|
||||
|
|
|
|||
|
|
@ -152,6 +152,13 @@ migrationJob:
|
|||
# the writable scratch space a read-only root filesystem needs.
|
||||
volumes: []
|
||||
volumeMounts: []
|
||||
# Scheduling for the Job pod, same shape as gateway.nodeSelector /
|
||||
# gateway.tolerations / gateway.affinity. The Job does not inherit the other
|
||||
# components' scheduling values: a migration usually needs a larger node
|
||||
# than the gateway, so pin it here explicitly.
|
||||
nodeSelector: {}
|
||||
tolerations: []
|
||||
affinity: {}
|
||||
image:
|
||||
repository: ghcr.io/berriai/litellm-migrations
|
||||
tag: "" # defaults to .Chart.AppVersion
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ from litellm.litellm_core_utils.core_helpers import process_response_headers
|
|||
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
|
||||
_safe_convert_created_field,
|
||||
)
|
||||
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
|
||||
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
|
||||
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
|
||||
from litellm.llms.openai.chat.gpt_5_transformation import is_gpt_reasoning_series_name
|
||||
|
|
@ -205,29 +206,76 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
|
|||
`remove_cache_control_flag_from_messages_and_tools`; mirror that here.
|
||||
"""
|
||||
|
||||
input = self._validate_input_param(input)
|
||||
tools = response_api_optional_request_params.get("tools")
|
||||
input, tools = self.remove_cache_control_flag_from_input_and_tools(model=model, input=input, tools=tools)
|
||||
sanitized_tools: Final = self._flatten_tool_schema_combinators_for_openai(
|
||||
model=model, tools=tools, litellm_params=litellm_params
|
||||
replay_safe_input, sanitized_tools = self._prepared_input_and_tools(
|
||||
model=model,
|
||||
input=input,
|
||||
tools=response_api_optional_request_params.get("tools"),
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
if sanitized_tools is not None:
|
||||
response_api_optional_request_params["tools"] = sanitized_tools
|
||||
replay_safe_input: Final = self._drop_foreign_tool_call_item_ids(input)
|
||||
final_request_params: Final = dict(
|
||||
ResponsesAPIRequestParams(model=model, input=replay_safe_input, **response_api_optional_request_params)
|
||||
)
|
||||
|
||||
return final_request_params
|
||||
|
||||
def _prepared_input_and_tools(
|
||||
self,
|
||||
model: str,
|
||||
input: str | ResponseInputParam,
|
||||
tools: Sequence[ALL_RESPONSES_API_TOOL_PARAMS] | None,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
) -> tuple[str | ResponseInputParam, Sequence[ALL_RESPONSES_API_TOOL_PARAMS] | None]:
|
||||
validated_input: Final = self._validate_input_param(input)
|
||||
stripped_input, stripped_tools = self.remove_cache_control_flag_from_input_and_tools(
|
||||
model=model, input=validated_input, tools=tools
|
||||
)
|
||||
object_schema_tools: Final = self._tools_with_object_parameters(model=model, tools=stripped_tools)
|
||||
sanitized_tools: Final = self._flatten_tool_schema_combinators_for_openai(
|
||||
model=model, tools=object_schema_tools, litellm_params=litellm_params
|
||||
)
|
||||
return self._drop_foreign_tool_call_item_ids(stripped_input), sanitized_tools
|
||||
|
||||
def _tools_with_object_parameters(
|
||||
self, model: str, tools: Sequence[ALL_RESPONSES_API_TOOL_PARAMS] | None
|
||||
) -> Sequence[ALL_RESPONSES_API_TOOL_PARAMS] | None:
|
||||
"""Decode tool schemas handed over already JSON-encoded, which the Responses validator
|
||||
rejects with a 400 naming the routed model rather than the tool. A null or absent schema
|
||||
is left alone because the API accepts both."""
|
||||
if tools is None:
|
||||
return None
|
||||
decoded: Final = [ # mutable-ok: request tools are a JSON list
|
||||
self._tool_with_object_parameters(model=model, index=index, tool=tool) for index, tool in enumerate(tools)
|
||||
]
|
||||
return cast("Sequence[ALL_RESPONSES_API_TOOL_PARAMS]", decoded) # cast-ok: dict spread keeps each tool's shape
|
||||
|
||||
def _tool_with_object_parameters(self, model: str, index: int, tool: object) -> object:
|
||||
if not isinstance(tool, dict) or tool.get("parameters") is None:
|
||||
return tool
|
||||
parameters: Final = tool["parameters"]
|
||||
if isinstance(parameters, dict):
|
||||
return tool
|
||||
decoded: Final = safe_json_loads(parameters) if isinstance(parameters, str) else None
|
||||
if isinstance(decoded, dict):
|
||||
return {**tool, "parameters": decoded} # mutable-ok: request tools are JSON dicts
|
||||
raise litellm.BadRequestError(
|
||||
message=(
|
||||
f"Invalid type for 'tools[{index}].parameters': expected an object, "
|
||||
f"but got {type(parameters).__name__} instead."
|
||||
),
|
||||
model=model,
|
||||
llm_provider=self.custom_llm_provider,
|
||||
)
|
||||
|
||||
def remove_cache_control_flag_from_input_and_tools(
|
||||
self,
|
||||
model: str, # allows overrides to selectively run this
|
||||
input: str | ResponseInputParam,
|
||||
tools: list[ALL_RESPONSES_API_TOOL_PARAMS] | None = None,
|
||||
tools: Sequence[ALL_RESPONSES_API_TOOL_PARAMS] | None = None,
|
||||
) -> tuple[
|
||||
str | ResponseInputParam,
|
||||
list[ALL_RESPONSES_API_TOOL_PARAMS] | None,
|
||||
Sequence[ALL_RESPONSES_API_TOOL_PARAMS] | None,
|
||||
]:
|
||||
"""Sibling of `remove_cache_control_flag_from_messages_and_tools` on
|
||||
the chat path. Strips Anthropic-only `cache_control` markers from
|
||||
|
|
@ -272,9 +320,9 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
|
|||
def _flatten_tool_schema_combinators_for_openai(
|
||||
self,
|
||||
model: str,
|
||||
tools: list[ALL_RESPONSES_API_TOOL_PARAMS] | None, # mutable-ok: request tools are a JSON list
|
||||
tools: Sequence[ALL_RESPONSES_API_TOOL_PARAMS] | None,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
) -> list[ALL_RESPONSES_API_TOOL_PARAMS] | None: # mutable-ok: request tools are a JSON list
|
||||
) -> Sequence[ALL_RESPONSES_API_TOOL_PARAMS] | None:
|
||||
"""Flatten top-level schema combinators only where OpenAI's validator rejects them.
|
||||
|
||||
OpenAI-compatible backends reusing this config (and the ChatGPT backend
|
||||
|
|
@ -293,7 +341,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
|
|||
flattened: Final = [ # mutable-ok: request tools are a JSON list
|
||||
self._flattened_tool_or_passthrough(tool) for tool in tools
|
||||
]
|
||||
return cast("list[ALL_RESPONSES_API_TOOL_PARAMS]", flattened) # cast-ok: dict spread keeps each tool's shape
|
||||
return cast("Sequence[ALL_RESPONSES_API_TOOL_PARAMS]", flattened) # cast-ok: spread keeps each tool's shape
|
||||
|
||||
@staticmethod
|
||||
def _flattened_tool_or_passthrough(tool: object) -> object:
|
||||
|
|
@ -786,15 +834,14 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
|
|||
compact_path: Final = parsed_url.path.rstrip("/") + "/compact"
|
||||
url: Final = str(parsed_url.copy_with(path=compact_path))
|
||||
|
||||
input = self._validate_input_param(input)
|
||||
tools = response_api_optional_request_params.get("tools")
|
||||
input, tools = self.remove_cache_control_flag_from_input_and_tools(model=model, input=input, tools=tools)
|
||||
sanitized_tools: Final = self._flatten_tool_schema_combinators_for_openai(
|
||||
model=model, tools=tools, litellm_params=litellm_params
|
||||
replay_safe_input, sanitized_tools = self._prepared_input_and_tools(
|
||||
model=model,
|
||||
input=input,
|
||||
tools=response_api_optional_request_params.get("tools"),
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
if sanitized_tools is not None:
|
||||
response_api_optional_request_params["tools"] = sanitized_tools
|
||||
replay_safe_input: Final = self._drop_foreign_tool_call_item_ids(input)
|
||||
data: Final = dict(
|
||||
ResponsesAPIRequestParams(model=model, input=replay_safe_input, **response_api_optional_request_params)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -300,6 +300,89 @@ class TestOpenAIResponsesAPIConfig:
|
|||
|
||||
assert result["input"][0]["id"] == "toolu_01Foreign"
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"raw_parameters",
|
||||
[
|
||||
'{"type": "object", "properties": {"city": {"type": "string"}}, "required": ["city"]}',
|
||||
'{"type":"object","properties":{"city":{"type":"string"}},"required":["city"]}',
|
||||
],
|
||||
)
|
||||
def test_transform_decodes_json_string_tool_parameters(self, raw_parameters: str):
|
||||
"""A JSON-encoded schema must reach the provider as an object."""
|
||||
result = self.config.transform_responses_api_request(
|
||||
model=self.model,
|
||||
input="weather in Paris",
|
||||
response_api_optional_request_params={
|
||||
"tools": [{"type": "function", "name": "get_weather", "parameters": raw_parameters}]
|
||||
},
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert result["tools"][0]["parameters"] == {
|
||||
"type": "object",
|
||||
"properties": {"city": {"type": "string"}},
|
||||
"required": ["city"],
|
||||
}
|
||||
|
||||
def test_transform_decodes_json_string_tool_parameters_on_compact_request(self):
|
||||
"""The compact request path builds the same wire body, so it must decode too."""
|
||||
_url, data = self.config.transform_compact_response_api_request(
|
||||
model=self.model,
|
||||
input="weather in Paris",
|
||||
response_api_optional_request_params={
|
||||
"tools": [{"type": "function", "name": "get_weather", "parameters": '{"type": "object"}'}]
|
||||
},
|
||||
api_base="https://api.openai.com/v1/responses",
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert data["tools"][0]["parameters"] == {"type": "object"}
|
||||
|
||||
@pytest.mark.parametrize("raw_parameters", ['"just a string"', "not json at all", "[1, 2, 3]", 42])
|
||||
def test_transform_rejects_tool_parameters_that_are_not_an_object(self, raw_parameters: object):
|
||||
"""Neither an object nor a string encoding one is a client error naming the tool index."""
|
||||
with pytest.raises(litellm.BadRequestError) as exc_info:
|
||||
self.config.transform_responses_api_request(
|
||||
model=self.model,
|
||||
input="weather in Paris",
|
||||
response_api_optional_request_params={
|
||||
"tools": [
|
||||
{"type": "web_search_preview"},
|
||||
{"type": "function", "name": "get_weather", "parameters": raw_parameters},
|
||||
]
|
||||
},
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert "tools[1].parameters" in str(exc_info.value)
|
||||
|
||||
def test_transform_leaves_object_null_and_absent_tool_parameters_untouched(self):
|
||||
"""The API accepts an object schema, an explicit null, an omitted schema and a built-in
|
||||
tool, so decoding must forward all four unchanged rather than raising."""
|
||||
schema = {"type": "object", "properties": {"city": {"type": "string"}}}
|
||||
tools = [
|
||||
{"type": "function", "name": "get_weather", "parameters": schema},
|
||||
{"type": "function", "name": "null_args", "parameters": None},
|
||||
{"type": "function", "name": "no_args"},
|
||||
{"type": "web_search_preview"},
|
||||
]
|
||||
|
||||
result = self.config.transform_responses_api_request(
|
||||
model=self.model,
|
||||
input="weather in Paris",
|
||||
response_api_optional_request_params={"tools": tools},
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert result["tools"][0]["parameters"] == schema
|
||||
assert result["tools"][1]["parameters"] is None
|
||||
assert "parameters" not in result["tools"][2]
|
||||
assert result["tools"][3] == {"type": "web_search_preview"}
|
||||
|
||||
def test_transform_compact_drops_foreign_tool_call_item_ids(self):
|
||||
"""The compact request path replays input the same way, so it must
|
||||
apply the same id drop."""
|
||||
|
|
@ -864,6 +947,20 @@ class TestAzureResponsesAPIConfig:
|
|||
self.model = "gpt-4o"
|
||||
self.logging_obj = MagicMock()
|
||||
|
||||
def test_azure_decodes_json_string_tool_parameters(self):
|
||||
"""Azure reaches the same wire through `super()`, after un-nesting a chat-shaped tool."""
|
||||
result = self.config.transform_responses_api_request(
|
||||
model=self.model,
|
||||
input="weather in Paris",
|
||||
response_api_optional_request_params={
|
||||
"tools": [{"type": "function", "function": {"name": "get_weather", "parameters": '{"type":"object"}'}}]
|
||||
},
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert result["tools"][0]["parameters"] == {"type": "object"}
|
||||
|
||||
def test_azure_get_complete_url_with_version_types(self):
|
||||
"""Test Azure get_complete_url with different API version types"""
|
||||
base_url = "https://litellm8397336933.openai.azure.com"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,6 +1,6 @@
|
|||
{
|
||||
"LIT001": {
|
||||
"limit": 22183
|
||||
"limit": 22181
|
||||
},
|
||||
"LIT002": {
|
||||
"limit": 26745
|
||||
|
|
@ -27,10 +27,10 @@
|
|||
"limit": 0
|
||||
},
|
||||
"LIT010": {
|
||||
"limit": 16468
|
||||
"limit": 16464
|
||||
},
|
||||
"LIT011": {
|
||||
"limit": 5510
|
||||
"limit": 5506
|
||||
},
|
||||
"LIT012": {
|
||||
"limit": 4486
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
});
|
||||
});
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue