Merge branch 'litellm_internal_staging' into litellm_bump_mcp_1_28_1

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yassin 2026-07-16 22:47:28 +00:00
commit 4526d8a8ca
65 changed files with 1558 additions and 3807 deletions

View file

@ -5,6 +5,8 @@ on:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
permissions:
contents: read

View file

@ -4,7 +4,11 @@ on:
push:
branches: [main, litellm_internal_staging]
pull_request:
branches: [main, litellm_internal_staging]
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}

View file

@ -114,6 +114,7 @@ COPY --from=builder /app/litellm/proxy/prisma_migration.py /app/litellm/proxy/pr
# working directory on sys.path; litellm/proxy/hooks resolves
# enterprise.enterprise_hooks from it)
COPY --from=builder /app/enterprise /app/enterprise
COPY --from=builder /app/litellm-proxy-extras /app/litellm-proxy-extras
# Prisma binaries live in $HOME/.cache (default prisma-python location),
# which is /root/.cache here. Copy only the Prisma subdirs — copying the
# whole /root/.cache drags in the uv build cache (~660 MB, includes a

View file

@ -111,6 +111,7 @@ COPY --from=builder /app/litellm/proxy/prisma_migration.py /app/litellm/proxy/pr
# working directory on sys.path; litellm/proxy/hooks resolves
# enterprise.enterprise_hooks from it)
COPY --from=builder /app/enterprise /app/enterprise
COPY --from=builder /app/litellm-proxy-extras /app/litellm-proxy-extras
# Prisma binaries live in $HOME/.cache (default prisma-python location),
# which is /root/.cache here. Copy them from the builder so they survive
# deployments that volume-mount /app/.cache (e.g. readOnlyRootFilesystem

View file

@ -137,6 +137,7 @@ COPY --from=builder /app/litellm/proxy/prisma_migration.py /app/litellm/proxy/pr
# working directory on sys.path; litellm/proxy/hooks resolves
# enterprise.enterprise_hooks from it)
COPY --from=builder /app/enterprise /app/enterprise
COPY --from=builder /app/litellm-proxy-extras /app/litellm-proxy-extras
COPY --from=builder /app/.cache /app/.cache
COPY --from=builder /var/lib/litellm/ui /var/lib/litellm/ui
COPY --from=builder /var/lib/litellm/assets /var/lib/litellm/assets

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-enterprise"
version = "0.1.50"
version = "0.1.51"
description = "Package for LiteLLM Enterprise features"
readme = "README.md"
requires-python = ">=3.9"
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
module-root = ""
[tool.commitizen]
version = "0.1.50"
version = "0.1.51"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-enterprise==",

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-proxy-extras"
version = "0.4.77"
version = "0.4.78"
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
readme = "README.md"
requires-python = ">=3.9"
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
module-root = ""
[tool.commitizen]
version = "0.4.77"
version = "0.4.78"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-proxy-extras==",

View file

@ -103,6 +103,18 @@ class DualCache(BaseCache):
if default_redis_ttl is not None:
self.default_redis_ttl = default_redis_ttl
def _backfill_kwargs(self, kwargs: "dict[str, object]") -> "dict[str, object]":
"""
Kwargs for writing a Redis read result into the in-memory tier.
Applies ``default_in_memory_ttl`` exactly like the write paths do;
without it, backfilled entries fall to ``InMemoryCache``'s own default
TTL and can outlive the TTL this cache was configured with.
"""
if "ttl" not in kwargs and self.default_in_memory_ttl is not None:
return {**kwargs, "ttl": self.default_in_memory_ttl}
return kwargs
def set_cache(self, key, value, local_only: bool = False, **kwargs):
# Update both Redis and in-memory cache
try:
@ -160,7 +172,7 @@ class DualCache(BaseCache):
if redis_result is not None:
# Update in-memory cache with the value from Redis
self.in_memory_cache.set_cache(key, redis_result, **kwargs)
self.in_memory_cache.set_cache(key, redis_result, **self._backfill_kwargs(kwargs))
result = redis_result
@ -226,7 +238,7 @@ class DualCache(BaseCache):
if redis_result is not None:
# Update in-memory cache with the value from Redis
await self.in_memory_cache.async_set_cache(key, redis_result, **kwargs)
await self.in_memory_cache.async_set_cache(key, redis_result, **self._backfill_kwargs(kwargs))
result = redis_result
@ -318,7 +330,7 @@ class DualCache(BaseCache):
result[key_to_index[key]] = value
if value is not None and self.in_memory_cache is not None:
await self.in_memory_cache.async_set_cache(key, value, **kwargs)
await self.in_memory_cache.async_set_cache(key, value, **self._backfill_kwargs(kwargs))
return result
except Exception:

View file

@ -2369,6 +2369,10 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
None,
description="If True, stores request messages and responses in spend logs. Default is False.",
)
disable_auto_add_proxy_admin_to_teams: bool | None = Field(
None,
description="By default, the user calling /team/new is automatically added to the new team as a team admin. If True, proxy admins are no longer auto-added; members explicitly listed in members_with_roles are unaffected. Default is False.",
)
maximum_spend_logs_retention_period: Optional[str] = Field(
None,
description="Maximum retention period for spend logs (e.g., '7d' for 7 days). Logs older than this will be deleted.",

View file

@ -1996,16 +1996,6 @@ async def _user_api_key_auth_builder(
raise HTTPException(401, detail="Invalid API key, no token associated")
api_key = valid_token.token
# Add hashed token to cache
asyncio.create_task(
_cache_key_object(
hashed_token=api_key,
user_api_key_obj=valid_token,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
)
valid_token_dict = valid_token.model_dump(exclude_none=True)
valid_token_dict.pop("token", None)
# budget_throttle_pct is excluded from model_dump (it must not leak

View file

@ -61,6 +61,8 @@ from litellm.types.proxy.management_endpoints.common_daily_activity import (
)
from litellm.types.proxy.management_endpoints.scim_v2 import (
SCIM_ENTERPRISE_METADATA_KEY,
SCIM_ENTITLEMENTS_METADATA_KEY,
SCIM_ROLES_METADATA_KEY,
)
from litellm.types.proxy.management_endpoints.internal_user_endpoints import (
BulkUpdateUserRequest,
@ -690,15 +692,21 @@ async def _get_user_info_teams(
return team_list, teams_1
_SCIM_DIRECTORY_METADATA_KEYS = frozenset(
{SCIM_ENTERPRISE_METADATA_KEY, SCIM_ENTITLEMENTS_METADATA_KEY, SCIM_ROLES_METADATA_KEY}
)
def _redact_scim_enterprise_metadata(
metadata: Optional[Dict[str, Any]],
) -> Optional[Dict[str, Any]]:
"""SCIM enterprise attributes are persisted in user metadata so reporting can
group on them, but they are directory-only fields that generic user-info
endpoints must not surface; SCIM clients read them through the SCIM endpoints."""
if not isinstance(metadata, dict) or SCIM_ENTERPRISE_METADATA_KEY not in metadata:
"""SCIM enterprise attributes, entitlements, and roles are persisted in user
metadata so reporting can group on them, but they are directory-only fields
that generic user-info endpoints must not surface; SCIM clients read them
through the SCIM endpoints."""
if not isinstance(metadata, dict) or not _SCIM_DIRECTORY_METADATA_KEYS.intersection(metadata):
return metadata
return {k: v for k, v in metadata.items() if k != SCIM_ENTERPRISE_METADATA_KEY}
return {k: v for k, v in metadata.items() if k not in _SCIM_DIRECTORY_METADATA_KEYS}
def _build_user_info_response(

View file

@ -1,5 +1,8 @@
from typing import List, Union
from typing import Callable, List, TypeVar, Union
from pydantic import ValidationError
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import (
LiteLLM_TeamTable,
LiteLLM_UserTable,
@ -9,6 +12,8 @@ from litellm.proxy._types import (
from litellm.repositories.team_repository import TeamRepository
from litellm.types.proxy.management_endpoints.scim_v2 import *
T = TypeVar("T")
class ScimTransformations:
DEFAULT_SCIM_NAME = "Unknown User"
@ -47,11 +52,19 @@ class ScimTransformations:
active = True if scim_active is None else bool(scim_active)
schemas = ["urn:ietf:params:scim:schemas:core:2.0:User"]
enterprise_user = None
if metadata.get(SCIM_ENTERPRISE_METADATA_KEY):
enterprise_user = SCIMEnterpriseUser.model_validate(metadata[SCIM_ENTERPRISE_METADATA_KEY])
enterprise_user = ScimTransformations._parse_directory_metadata(
user, SCIM_ENTERPRISE_METADATA_KEY, SCIMEnterpriseUser.model_validate
)
if enterprise_user is not None:
schemas.append(SCIM_ENTERPRISE_USER_SCHEMA)
entitlements = ScimTransformations._parse_directory_metadata(
user, SCIM_ENTITLEMENTS_METADATA_KEY, SCIM_MULTI_VALUED_LIST_ADAPTER.validate_python
)
roles = ScimTransformations._parse_directory_metadata(
user, SCIM_ROLES_METADATA_KEY, SCIM_MULTI_VALUED_LIST_ADAPTER.validate_python
)
return SCIMUser(
schemas=schemas,
id=user.user_id,
@ -64,6 +77,8 @@ class ScimTransformations:
emails=emails,
groups=groups,
active=active,
entitlements=entitlements,
roles=roles,
enterprise_user=enterprise_user,
meta={
"resourceType": "User",
@ -72,6 +87,31 @@ class ScimTransformations:
},
)
@staticmethod
def _parse_directory_metadata(
user: Union[LiteLLM_UserTable, NewUserResponse],
key: str,
validate: Callable[[object], T],
) -> T | None:
"""A SCIM directory attribute parsed from user metadata, or None when absent or malformed.
Metadata is writable outside the SCIM surface, so a malformed value on one user must not
fail the whole directory response; the attribute is omitted and the corruption logged.
"""
metadata = user.metadata or {}
raw = metadata.get(key)
if not raw:
return None
try:
return validate(raw)
except ValidationError:
verbose_proxy_logger.warning(
"Skipping malformed %s metadata on user %s in SCIM response",
key,
user.user_id,
)
return None
@staticmethod
def _get_scim_user_name(user: Union[LiteLLM_UserTable, NewUserResponse]) -> str:
"""

View file

@ -17,7 +17,7 @@ from fastapi import (
Request,
Response,
)
from pydantic import BaseModel
from pydantic import BaseModel, ValidationError
from typing_extensions import TypedDict
import litellm
@ -125,6 +125,8 @@ class ScimUserData(TypedDict):
family_name: Optional[str]
active: Optional[bool]
enterprise: Optional[SCIMEnterpriseUser]
entitlements: list[SCIMMultiValuedAttribute] | None
roles: list[SCIMMultiValuedAttribute] | None
class GroupMemberExtractionResult(BaseModel):
@ -199,6 +201,8 @@ def _extract_scim_user_data(user: SCIMUser) -> ScimUserData:
"family_name": user.name.familyName if user.name else None,
"active": user.active,
"enterprise": user.enterprise_user,
"entitlements": user.entitlements,
"roles": user.roles,
}
@ -207,6 +211,8 @@ def _build_scim_metadata(
family_name: Optional[str],
active: Optional[bool] = None,
enterprise: Optional[SCIMEnterpriseUser] = None,
entitlements: list[SCIMMultiValuedAttribute] | None = None,
roles: list[SCIMMultiValuedAttribute] | None = None,
) -> Dict[str, Any]:
"""Build metadata dictionary with SCIM data."""
metadata: Dict[str, Any] = {
@ -222,6 +228,12 @@ def _build_scim_metadata(
if enterprise is not None:
metadata[SCIM_ENTERPRISE_METADATA_KEY] = enterprise.model_dump(by_alias=True, exclude_none=True)
if entitlements is not None:
metadata[SCIM_ENTITLEMENTS_METADATA_KEY] = [e.model_dump(exclude_none=True) for e in entitlements]
if roles is not None:
metadata[SCIM_ROLES_METADATA_KEY] = [r.model_dump(exclude_none=True) for r in roles]
return metadata
@ -739,6 +751,62 @@ def _get_schemas() -> list:
),
],
),
SCIMSchemaAttribute(
name="entitlements",
type="complex",
multiValued=True,
description="A list of entitlements for the user.",
subAttributes=[
SCIMSchemaAttribute(
name="value",
type="string",
description="The value of an entitlement.",
),
SCIMSchemaAttribute(
name="display",
type="string",
description="A human-readable name for the entitlement.",
),
SCIMSchemaAttribute(
name="type",
type="string",
description="A label indicating the entitlement's function.",
),
SCIMSchemaAttribute(
name="primary",
type="boolean",
description="Whether this is the primary entitlement.",
),
],
),
SCIMSchemaAttribute(
name="roles",
type="complex",
multiValued=True,
description="A list of roles for the user.",
subAttributes=[
SCIMSchemaAttribute(
name="value",
type="string",
description="The value of a role.",
),
SCIMSchemaAttribute(
name="display",
type="string",
description="A human-readable name for the role.",
),
SCIMSchemaAttribute(
name="type",
type="string",
description="A label indicating the role's function.",
),
SCIMSchemaAttribute(
name="primary",
type="boolean",
description="Whether this is the primary role.",
),
],
),
],
meta={
"location": "/scim/v2/Schemas/urn:ietf:params:scim:schemas:core:2.0:User",
@ -1074,6 +1142,8 @@ async def create_user(
user_data["given_name"],
user_data["family_name"],
enterprise=user_data["enterprise"],
entitlements=user_data["entitlements"],
roles=user_data["roles"],
)
default_role = _default_scim_user_role()
@ -1152,6 +1222,8 @@ async def update_user(
user_data["family_name"],
scim_active_for_metadata,
enterprise=user_data["enterprise"],
entitlements=user_data["entitlements"],
roles=user_data["roles"],
)
await _handle_team_membership_changes(
@ -1311,6 +1383,48 @@ def _handle_group_operations(op_type: str, value: Any, teams_set: Set[str]) -> O
return None
def _multi_valued_attribute_base(path: str) -> str:
"""The attribute name a SCIM path targets, stripped of any value filter or sub-attribute."""
return path.split("[", 1)[0].split(".", 1)[0]
def _handle_multi_valued_attribute_update(path: str, op_type: str, value: Any, metadata: dict[str, Any]) -> None:
"""Handle add/replace/remove for the entitlements and roles multi-valued attributes."""
base = _multi_valued_attribute_base(path)
metadata_key = SCIM_MULTI_VALUED_ATTRIBUTE_METADATA_KEYS[base]
if path != base:
raise HTTPException(
status_code=400,
detail={"error": f"Filtered or sub-attribute paths are not supported for {base}; PATCH the full attribute"},
)
if op_type == "remove":
metadata.pop(metadata_key, None)
return
if value is None:
raise HTTPException(
status_code=400,
detail={"error": f"The {op_type} operation on {base} requires a 'value' member (RFC 7644 Section 3.5.2)"},
)
normalized = value if isinstance(value, list) else [value]
try:
attrs = SCIM_MULTI_VALUED_LIST_ADAPTER.validate_python(normalized)
except ValidationError:
raise HTTPException(
status_code=400,
detail={"error": f"Invalid value for {base}: expected a list of objects with a 'value' sub-attribute"},
)
dumped = [attr.model_dump(exclude_none=True) for attr in attrs]
existing = metadata.get(metadata_key)
if op_type == "add" and isinstance(existing, list):
metadata[metadata_key] = existing + dumped
return
metadata[metadata_key] = dumped
def _handle_generic_metadata(path: str, op_type: str, value: Any, metadata: Dict[str, Any]) -> None:
"""Handle generic metadata operations for unknown paths."""
if op_type == "remove":
@ -1346,6 +1460,8 @@ def _apply_patch_ops(
_handle_displayname_update(op_type, val, update_data)
elif key_lower == "externalid":
_handle_externalid_update(op_type, val, update_data)
elif key_lower in SCIM_MULTI_VALUED_ATTRIBUTE_METADATA_KEYS:
_handle_multi_valued_attribute_update(key_lower, op_type, val, metadata)
elif key_lower == "name" and isinstance(val, dict):
for name_key, name_val in val.items():
name_key_lower = name_key.lower()
@ -1366,6 +1482,8 @@ def _apply_patch_ops(
_handle_active_update(op_type, value, metadata)
elif path in ("name.givenname", "name.familyname"):
_handle_name_update(path, op_type, value, scim_metadata)
elif _multi_valued_attribute_base(path) in SCIM_MULTI_VALUED_ATTRIBUTE_METADATA_KEYS:
_handle_multi_valued_attribute_update(path, op_type, value, metadata)
elif path.startswith("groups"):
new_replace_set = _handle_group_operations(op_type, value, teams_set)
if new_replace_set is not None:

View file

@ -14,7 +14,7 @@ import json
import math
import traceback
from datetime import datetime, timezone
from typing import Annotated, Any, Dict, List, Optional, Tuple, Union, cast
from typing import Annotated, Any, Dict, List, Mapping, Optional, Tuple, Union, cast
import fastapi
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
@ -894,6 +894,17 @@ def _check_team_budget_update_authority(
)
def _should_auto_add_team_creator(
user_api_key_dict: UserAPIKeyAuth,
general_settings: Mapping[str, object],
) -> bool:
if user_api_key_dict.user_id is None:
return False
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
return True
return general_settings.get("disable_auto_add_proxy_admin_to_teams") is not True
#### TEAM MANAGEMENT ####
@router.post(
"/team/new",
@ -997,6 +1008,7 @@ async def new_team(
from litellm.proxy.proxy_server import (
_license_check,
create_audit_log_for_update,
general_settings,
litellm_proxy_admin_name,
prisma_client,
user_api_key_cache,
@ -1123,13 +1135,11 @@ async def new_team(
user_api_key_cache=user_api_key_cache,
)
if user_api_key_dict.user_id is not None:
creating_user_in_list = False
for member in data.members_with_roles:
if member.user_id == user_api_key_dict.user_id:
creating_user_in_list = True
if creating_user_in_list is False:
if _should_auto_add_team_creator(user_api_key_dict, general_settings):
creating_user_in_list = any(
member.user_id == user_api_key_dict.user_id for member in data.members_with_roles
)
if not creating_user_in_list:
data.members_with_roles.append(Member(role="admin", user_id=user_api_key_dict.user_id))
_check_passthrough_routes_caller_permission(data, user_api_key_dict, entity="team")

View file

@ -2814,21 +2814,6 @@ async def update_cache(
)
# set cooldown on alert
if existing_spend_obj is not None and getattr(existing_spend_obj, "team_spend", None) is not None:
existing_team_spend = existing_spend_obj.team_spend or 0
# Calculate the new cost by adding the existing cost and response_cost
existing_spend_obj.team_spend = existing_team_spend + response_cost
if existing_spend_obj is not None and getattr(existing_spend_obj, "team_member_spend", None) is not None:
existing_team_member_spend = existing_spend_obj.team_member_spend or 0
# Calculate the new cost by adding the existing cost and response_cost
existing_spend_obj.team_member_spend = existing_team_member_spend + response_cost
# Existing spend_obj is mutated; UserApiKeyCache.async_set_cache_pipeline turns
# BaseModel values into dicts for Redis (same Codec path as async_set_cache).
existing_spend_obj.spend = new_spend
values_to_update_in_cache.append((hashed_token, existing_spend_obj))
### UPDATE USER SPEND ###
async def _update_user_cache():
## UPDATE CACHE FOR USER ID + GLOBAL PROXY
@ -3032,13 +3017,27 @@ async def update_cache(
if tags is not None:
await _update_tag_cache()
asyncio.create_task(
user_api_key_cache.async_set_cache_pipeline(
cache_list=values_to_update_in_cache,
ttl=get_management_object_ttl(user_api_key_cache),
litellm_parent_otel_span=parent_otel_span,
global_proxy_spend_key = "{}:spend".format(litellm_proxy_admin_name)
local_object_updates = tuple((k, v) for k, v in values_to_update_in_cache if k != global_proxy_spend_key)
shared_scalar_updates = tuple((k, v) for k, v in values_to_update_in_cache if k == global_proxy_spend_key)
if local_object_updates:
asyncio.create_task(
user_api_key_cache.async_set_cache_pipeline(
cache_list=list(local_object_updates),
ttl=get_management_object_ttl(user_api_key_cache),
litellm_parent_otel_span=parent_otel_span,
local_only=True,
)
)
if shared_scalar_updates:
asyncio.create_task(
user_api_key_cache.async_set_cache_pipeline(
cache_list=list(shared_scalar_updates),
ttl=get_management_object_ttl(user_api_key_cache),
litellm_parent_otel_span=parent_otel_span,
)
)
)
def run_ollama_serve():
@ -5794,6 +5793,13 @@ class ProxyConfig:
# For other types, convert to bool
general_settings["store_prompts_in_spend_logs"] = bool(value)
if "disable_auto_add_proxy_admin_to_teams" in _general_settings:
value = _general_settings["disable_auto_add_proxy_admin_to_teams"]
if isinstance(value, str):
general_settings["disable_auto_add_proxy_admin_to_teams"] = value.lower() == "true"
else:
general_settings["disable_auto_add_proxy_admin_to_teams"] = value if value is None else bool(value)
## STORE MODEL IN DB ##
if "store_model_in_db" in _general_settings:
value = _general_settings["store_model_in_db"]
@ -14907,6 +14913,7 @@ async def get_config_list(
"mcp_required_fields": {"type": "List"},
"cancel_on_disconnect": {"type": "Boolean"},
"skip_user_budget_on_team_key": {"type": "Boolean"},
"disable_auto_add_proxy_admin_to_teams": {"type": "Boolean"},
}
return_val = []

View file

@ -362,6 +362,8 @@ async def route_request(
for _key in _MOCK_TESTING_KWARG_NAMES:
data.pop(_key, None)
data.pop("enable_tag_filtering", None)
team_id = get_team_id_from_data(data)
router_model_names = llm_router.model_names if llm_router is not None else []
is_proxy_admin_without_team = team_id is None and _is_proxy_admin_request(data)
@ -410,6 +412,7 @@ async def route_request(
"timeout",
"model_group_retry_policy",
"routing_strategy",
"enable_tag_filtering",
]
# Merge override settings into data (only if not already set in request)

View file

@ -9778,6 +9778,7 @@ class Router:
"retry_policy",
"model_group_alias",
"enable_weighted_failover",
"enable_tag_filtering",
]
for var in vars_to_include:
@ -9814,6 +9815,7 @@ class Router:
"model_group_retry_policy",
"model_group_alias",
"enable_weighted_failover",
"enable_tag_filtering",
]
_int_settings = [

View file

@ -160,8 +160,14 @@ async def get_deployments_for_tag(
Returns a list of deployments that match the requested model and tags in the request.
Executes tag based filtering based on the tags in request metadata and the tags on the deployments
Runs when the router-level `enable_tag_filtering` is True or the request carries
`enable_tag_filtering=True` (set from key/team router_settings by the proxy).
A request-level False never disables a router-level True, so per-request settings
cannot escape an operator's global tag-routing policy.
"""
if llm_router_instance.enable_tag_filtering is not True:
request_enable_tag_filtering = request_kwargs.get("enable_tag_filtering") if request_kwargs else None
if request_enable_tag_filtering is not True and llm_router_instance.enable_tag_filtering is not True:
return healthy_deployments
if request_kwargs is None:

View file

@ -6,13 +6,17 @@ from pydantic import (
ConfigDict,
EmailStr,
Field,
TypeAdapter,
field_validator,
model_serializer,
model_validator,
)
from pydantic_core.core_schema import SerializerFunctionWrapHandler
SCIM_ENTERPRISE_USER_SCHEMA = "urn:ietf:params:scim:schemas:extension:enterprise:2.0:User"
SCIM_ENTERPRISE_METADATA_KEY = "scim_enterprise"
SCIM_ENTITLEMENTS_METADATA_KEY = "scim_entitlements"
SCIM_ROLES_METADATA_KEY = "scim_roles"
class LiteLLM_UserScimMetadata(BaseModel):
@ -53,6 +57,28 @@ class SCIMUserGroup(BaseModel):
type: Optional[str] = "direct" # direct or indirect
class SCIMMultiValuedAttribute(BaseModel):
value: str
display: Optional[str] = None
type: Optional[str] = None
primary: Optional[bool] = None
@model_validator(mode="before")
@classmethod
def coerce_bare_string(cls, data: object) -> object:
if isinstance(data, str):
return {"value": data}
return data
SCIM_MULTI_VALUED_LIST_ADAPTER = TypeAdapter(List[SCIMMultiValuedAttribute])
SCIM_MULTI_VALUED_ATTRIBUTE_METADATA_KEYS = {
"entitlements": SCIM_ENTITLEMENTS_METADATA_KEY,
"roles": SCIM_ROLES_METADATA_KEY,
}
class SCIMUserManager(BaseModel):
model_config = ConfigDict(populate_by_name=True)
@ -81,6 +107,8 @@ class SCIMUser(SCIMResource):
active: bool = True
emails: Optional[List[SCIMUserEmail]] = None
groups: Optional[List[SCIMUserGroup]] = None
entitlements: Optional[List[SCIMMultiValuedAttribute]] = None
roles: Optional[List[SCIMMultiValuedAttribute]] = None
enterprise_user: Optional[SCIMEnterpriseUser] = Field(
default=None,
alias=SCIM_ENTERPRISE_USER_SCHEMA,
@ -88,11 +116,15 @@ class SCIMUser(SCIMResource):
)
@model_serializer(mode="wrap")
def _omit_absent_enterprise(self, handler: SerializerFunctionWrapHandler) -> Dict[str, Any]:
def _omit_absent_optional_blocks(self, handler: SerializerFunctionWrapHandler) -> Dict[str, Any]:
dumped = handler(self)
if self.enterprise_user is None:
dumped.pop(SCIM_ENTERPRISE_USER_SCHEMA, None)
dumped.pop("enterprise_user", None)
if self.entitlements is None:
dumped.pop("entitlements", None)
if self.roles is None:
dumped.pop("roles", None)
return dumped

View file

@ -117,6 +117,7 @@ class UpdateRouterConfig(BaseModel):
fallbacks: Optional[List[dict]] = None
context_window_fallbacks: Optional[List[dict]] = None
model_group_alias: Optional[Dict[str, Union[str, Dict]]] = {}
enable_tag_filtering: Optional[bool] = None
model_config = ConfigDict(protected_namespaces=())

View file

@ -3222,6 +3222,7 @@ all_litellm_params = (
"shared_session",
"search_tool_name",
"order",
"enable_tag_filtering",
"enable_json_schema_validation",
"use_xai_oauth",
"_litellm_rate_limit_descriptors",

View file

@ -62,8 +62,8 @@ proxy = [
"azure-identity>=1.25.2,<2.0",
"azure-storage-blob>=12.28.0,<13.0",
"mcp>=1.28.1,<2.0",
"litellm-proxy-extras==0.4.77",
"litellm-enterprise==0.1.50",
"litellm-proxy-extras==0.4.78",
"litellm-enterprise==0.1.51",
"RestrictedPython>=8.1,<9.0",
"rich>=13.9.4,<14.0",
"InquirerPy>=0.3.4,<1.0",

View file

@ -17,7 +17,7 @@ Each subdirectory under `tests/e2e/` is one suite, scoped to an endpoint family
- `security/` - secret handling and log-leak protection
- `router/` - routing and reliability behavior (fallbacks, cooldowns)
- `gateway/` - proxy configuration only (`litellm-config.yml`); no tests
- `claude_code/` - the Claude Code compatibility matrix: drives the real `claude` CLI (and HTTP probes) against a proxy for each feature x provider cell, reporting tagged-union outcomes via the `compat_result` fixture; ships its own driver/builder/publisher plus `_*_unit_tests/` trees, and does not use the shared transport harness
- `claude_code/` - the Claude Code compatibility matrix: drives the real `claude` CLI (and HTTP probes) against a proxy for each feature x provider cell, reporting tagged-union outcomes via the `compat_result` fixture; ships its own driver/builder/publisher and does not use the shared transport harness
## Lay the pattern down in a class
@ -53,7 +53,7 @@ Each suite provides its own `client` fixture (see `llm_translation/passthrough_c
Request and response bodies are typed pydantic models in `models.py`; only the fields a test reads are modelled, and nothing passes raw dicts. Outcomes come back as a `Result[R]` tagged union (`Success`, `NetworkError`, `UnauthorizedError`, `RateLimitedError`, `ValidationError`, `UnknownApiError`). Handle them with `match`, or call `unwrap(...)` when a non-success should fail the test. The skip-vs-fail split is deliberate: a test marked `e2e` skips when no proxy answers its liveness probe, but once a request reaches the proxy any wrong behavior is a hard failure, never a skip
Mark live tests with `@pytest.mark.e2e` (on the class or the module). Pure coverage of the harness itself carries no marker and runs regardless. Use `scoped_key` for a fresh all-models key that auto-deletes, `resources` when you need to create and tear down more than a key, and `unique_marker()` from `e2e_config` to keep prompts, tags, and customer ids from colliding across concurrent runs and the shared response cache
Mark live tests with `@pytest.mark.e2e` (on the class or the module). `tests/e2e/` is for live proxy suites only; do not put unit tests here. Use `scoped_key` for a fresh all-models key that auto-deletes, `resources` when you need to create and tear down more than a key, and `unique_marker()` from `e2e_config` to keep prompts, tags, and customer ids from colliding across concurrent runs and the shared response cache
## Typing

View file

@ -1,38 +0,0 @@
{
"schema_version": "1",
"generated_at": "2026-04-25T00:00:00Z",
"litellm_version": "v1.83.0-stable",
"claude_code_version": "2.1.120",
"providers": [
"anthropic",
"bedrock_invoke"
],
"features": [
{
"id": "basic_messaging_non_streaming",
"name": "Basic messaging (non-streaming)",
"providers": {
"anthropic": {
"status": "pass"
},
"bedrock_invoke": {
"status": "not_tested"
}
}
},
{
"id": "tool_use",
"name": "Tool use",
"providers": {
"anthropic": {
"status": "fail",
"error": "[claude-sonnet-4-5] tool call dropped"
},
"bedrock_invoke": {
"status": "not_applicable",
"reason": "tool use not yet wired up for Bedrock Invoke"
}
}
}
]
}

View file

@ -1,9 +0,0 @@
schema_version: "1"
providers:
- anthropic
- bedrock_invoke
features:
- id: basic_messaging_non_streaming
name: Basic messaging (non-streaming)
- id: tool_use
name: Tool use

View file

@ -1,41 +0,0 @@
{
"schema_version": "1",
"results": [
{
"feature_id": "basic_messaging_non_streaming",
"provider": "anthropic",
"nodeid": "tests/e2e/claude_code/basic_messaging_non_streaming/test_anthropic.py::test_basic_messaging_non_streaming_anthropic[claude-haiku-4-5]",
"result": {"status": "pass"}
},
{
"feature_id": "basic_messaging_non_streaming",
"provider": "anthropic",
"nodeid": "tests/e2e/claude_code/basic_messaging_non_streaming/test_anthropic.py::test_basic_messaging_non_streaming_anthropic[claude-sonnet-4-5]",
"result": {"status": "pass"}
},
{
"feature_id": "basic_messaging_non_streaming",
"provider": "anthropic",
"nodeid": "tests/e2e/claude_code/basic_messaging_non_streaming/test_anthropic.py::test_basic_messaging_non_streaming_anthropic[claude-opus-4-7]",
"result": {"status": "pass"}
},
{
"feature_id": "tool_use",
"provider": "anthropic",
"nodeid": "tests/e2e/claude_code/tool_use/test_anthropic.py::test_x[claude-haiku-4-5]",
"result": {"status": "pass"}
},
{
"feature_id": "tool_use",
"provider": "anthropic",
"nodeid": "tests/e2e/claude_code/tool_use/test_anthropic.py::test_x[claude-sonnet-4-5]",
"result": {"status": "fail", "error": "[claude-sonnet-4-5] tool call dropped"}
},
{
"feature_id": "tool_use",
"provider": "bedrock_invoke",
"nodeid": "tests/e2e/claude_code/tool_use/test_bedrock_invoke.py::test_x[claude-haiku-4-5]",
"result": {"status": "not_applicable", "reason": "tool use not yet wired up for Bedrock Invoke"}
}
]
}

View file

@ -1,479 +0,0 @@
"""Golden-file tests for the Matrix JSON Builder.
These tests fix the published JSON schema. The builder is a pure function
from (manifest, results, metadata) matrix dict, so we feed it a fixture
input set and compare the produced dict to a checked-in expected output.
Any schema drift intentional or accidental surfaces as a diff in PR
review.
"""
from __future__ import annotations
import json
from pathlib import Path
import pytest
from claude_code.matrix_builder import (
ManifestError,
ResultsError,
build_from_paths,
build_matrix,
load_manifest,
load_results,
)
FIXTURES = Path(__file__).parent / "fixtures"
def test_build_matrix_matches_golden_file(tmp_path):
manifest = load_manifest(FIXTURES / "manifest.yaml")
results = load_results(FIXTURES / "results.json")
matrix = build_matrix(
manifest=manifest,
results=results,
litellm_version="v1.83.0-stable",
claude_code_version="2.1.120",
generated_at="2026-04-25T00:00:00Z",
)
expected = json.loads((FIXTURES / "expected_matrix.json").read_text())
assert matrix == expected
def test_build_matrix_pass_requires_all_models_pass():
"""Multiple results in one cell must all be pass for the cell to be pass."""
manifest = {
"schema_version": "1",
"providers": ["anthropic"],
"features": [{"id": "f", "name": "F"}],
}
results = [
{"feature_id": "f", "provider": "anthropic", "result": {"status": "pass"}},
{"feature_id": "f", "provider": "anthropic", "result": {"status": "pass"}},
{"feature_id": "f", "provider": "anthropic", "result": {"status": "pass"}},
]
matrix = build_matrix(
manifest=manifest,
results=results,
litellm_version="v",
claude_code_version="c",
generated_at="t",
)
assert matrix["features"][0]["providers"]["anthropic"] == {"status": "pass"}
def test_build_matrix_any_fail_makes_cell_fail():
manifest = {
"schema_version": "1",
"providers": ["anthropic"],
"features": [{"id": "f", "name": "F"}],
}
results = [
{"feature_id": "f", "provider": "anthropic", "result": {"status": "pass"}},
{
"feature_id": "f",
"provider": "anthropic",
"result": {"status": "fail", "error": "[claude-opus-4-7] timeout"},
},
{"feature_id": "f", "provider": "anthropic", "result": {"status": "pass"}},
]
matrix = build_matrix(
manifest=manifest,
results=results,
litellm_version="v",
claude_code_version="c",
generated_at="t",
)
cell = matrix["features"][0]["providers"]["anthropic"]
assert cell["status"] == "fail"
assert cell["error"] == "[claude-opus-4-7] timeout"
def test_build_matrix_joins_all_failure_errors_in_one_cell():
"""When multiple tiers fail for different reasons within the same cell,
every failure's error must appear in the published cell so triage
isn't reduced to a single tier's diagnostic.
"""
manifest = {
"schema_version": "1",
"providers": ["anthropic"],
"features": [{"id": "f", "name": "F"}],
}
results = [
{
"feature_id": "f",
"provider": "anthropic",
"result": {"status": "fail", "error": "[claude-haiku-4-5] 429"},
},
{"feature_id": "f", "provider": "anthropic", "result": {"status": "pass"}},
{
"feature_id": "f",
"provider": "anthropic",
"result": {"status": "fail", "error": "[claude-opus-4-7] timeout"},
},
]
matrix = build_matrix(
manifest=manifest,
results=results,
litellm_version="v",
claude_code_version="c",
generated_at="t",
)
cell = matrix["features"][0]["providers"]["anthropic"]
assert cell["status"] == "fail"
assert "[claude-haiku-4-5] 429" in cell["error"]
assert "[claude-opus-4-7] timeout" in cell["error"]
def test_build_matrix_mixed_pass_and_not_tested_surfaces_pass():
"""A `not_tested` row mixed with `pass` rows must not silently demote
the cell to `not_tested` `not_tested` is "absent data", not a
negative signal. Otherwise a partial crash mid-test, or a test that
explicitly recorded "tier didn't run", would discard real passing
results from the published cell.
"""
manifest = {
"schema_version": "1",
"providers": ["anthropic"],
"features": [{"id": "f", "name": "F"}],
}
results = [
{"feature_id": "f", "provider": "anthropic", "result": {"status": "pass"}},
{
"feature_id": "f",
"provider": "anthropic",
"result": {"status": "not_tested"},
},
{"feature_id": "f", "provider": "anthropic", "result": {"status": "pass"}},
]
matrix = build_matrix(
manifest=manifest,
results=results,
litellm_version="v",
claude_code_version="c",
generated_at="t",
)
assert matrix["features"][0]["providers"]["anthropic"] == {"status": "pass"}
def test_build_matrix_all_not_tested_stays_not_tested():
"""A cell whose every row is `not_tested` (or empty) must remain
`not_tested` the absent-data rule only drops `not_tested` rows
when there's other signal to surface.
"""
manifest = {
"schema_version": "1",
"providers": ["anthropic"],
"features": [{"id": "f", "name": "F"}],
}
results = [
{
"feature_id": "f",
"provider": "anthropic",
"result": {"status": "not_tested"},
},
{
"feature_id": "f",
"provider": "anthropic",
"result": {"status": "not_tested"},
},
]
matrix = build_matrix(
manifest=manifest,
results=results,
litellm_version="v",
claude_code_version="c",
generated_at="t",
)
assert matrix["features"][0]["providers"]["anthropic"] == {"status": "not_tested"}
def test_build_matrix_mixed_pass_and_not_applicable_surfaces_pass():
"""A `not_applicable` row mixed with `pass` rows must surface as
`pass`, not `not_applicable`. The published cell answers "does this
feature work on this provider?"; if any tier passes, the feature
works there. Discarding passing tiers because one tier is NA would
misrepresent the cell as unsupported.
"""
manifest = {
"schema_version": "1",
"providers": ["anthropic"],
"features": [{"id": "f", "name": "F"}],
}
results = [
{"feature_id": "f", "provider": "anthropic", "result": {"status": "pass"}},
{
"feature_id": "f",
"provider": "anthropic",
"result": {
"status": "not_applicable",
"reason": "haiku does not support extended thinking",
},
},
{"feature_id": "f", "provider": "anthropic", "result": {"status": "pass"}},
]
matrix = build_matrix(
manifest=manifest,
results=results,
litellm_version="v",
claude_code_version="c",
generated_at="t",
)
assert matrix["features"][0]["providers"]["anthropic"] == {"status": "pass"}
def test_build_matrix_all_not_applicable_stays_not_applicable():
"""When every observed row is `not_applicable`, the cell remains
`not_applicable` and the first row's reason carries through to the
published matrix.
"""
manifest = {
"schema_version": "1",
"providers": ["anthropic"],
"features": [{"id": "f", "name": "F"}],
}
results = [
{
"feature_id": "f",
"provider": "anthropic",
"result": {
"status": "not_applicable",
"reason": "feature unsupported on this provider",
},
},
{
"feature_id": "f",
"provider": "anthropic",
"result": {"status": "not_applicable", "reason": "ditto"},
},
]
matrix = build_matrix(
manifest=manifest,
results=results,
litellm_version="v",
claude_code_version="c",
generated_at="t",
)
assert matrix["features"][0]["providers"]["anthropic"] == {
"status": "not_applicable",
"reason": "feature unsupported on this provider",
}
def test_build_matrix_fills_not_tested_for_missing_cells():
manifest = {
"schema_version": "1",
"providers": ["anthropic", "azure"],
"features": [{"id": "f", "name": "F"}],
}
results = [
{"feature_id": "f", "provider": "anthropic", "result": {"status": "pass"}},
]
matrix = build_matrix(
manifest=manifest,
results=results,
litellm_version="v",
claude_code_version="c",
generated_at="t",
)
cells = matrix["features"][0]["providers"]
assert cells["anthropic"] == {"status": "pass"}
assert cells["azure"] == {"status": "not_tested"}
def test_build_matrix_preserves_provider_and_feature_order():
manifest = {
"schema_version": "1",
"providers": ["azure", "anthropic", "vertex_ai"],
"features": [
{"id": "z", "name": "Z"},
{"id": "a", "name": "A"},
],
}
matrix = build_matrix(
manifest=manifest,
results=[],
litellm_version="v",
claude_code_version="c",
generated_at="t",
)
assert matrix["providers"] == ["azure", "anthropic", "vertex_ai"]
assert [f["id"] for f in matrix["features"]] == ["z", "a"]
assert list(matrix["features"][0]["providers"].keys()) == [
"azure",
"anthropic",
"vertex_ai",
]
def test_build_matrix_emits_schema_version_one():
manifest = {
"schema_version": "1",
"providers": ["anthropic"],
"features": [{"id": "f", "name": "F"}],
}
matrix = build_matrix(
manifest=manifest,
results=[],
litellm_version="v",
claude_code_version="c",
generated_at="t",
)
assert matrix["schema_version"] == "1"
def test_load_manifest_rejects_wrong_schema_version(tmp_path):
bad = tmp_path / "manifest.yaml"
bad.write_text(
'schema_version: "2"\nproviders: [anthropic]\nfeatures:\n - id: f\n name: F\n'
)
with pytest.raises(ManifestError, match="schema_version"):
load_manifest(bad)
def test_load_manifest_rejects_empty_features(tmp_path):
bad = tmp_path / "manifest.yaml"
bad.write_text('schema_version: "1"\nproviders: [anthropic]\nfeatures: []\n')
with pytest.raises(ManifestError):
load_manifest(bad)
def test_load_results_rejects_missing_results_key(tmp_path):
bad = tmp_path / "results.json"
bad.write_text(json.dumps({"schema_version": "1"}))
with pytest.raises(ResultsError):
load_results(bad)
def test_build_matrix_6x5_grid_matches_published_sample():
"""Slice 5 acceptance: feeding the per-model results the full v0
row set produces reproduces the hand-authored 6x5 sample that the
docs page renders.
Inputs mirror the structure of `compat-results.json` after a real
run with the proxy configured for all five columns and all six
feature directories: every (feature, provider, model) cell yields a
`pass`. Anthropic announced Claude in Microsoft Foundry on
2025-11-18, so the Azure column is now exercised end-to-end like
the others rather than reporting `not_applicable`.
The aggregated matrix must equal the checked-in
`sample_compatibility-matrix.json` byte-for-byte (after JSON load),
so any future schema drift surfaces here in review.
"""
repo_root = Path(__file__).resolve().parents[1]
full_manifest = load_manifest(repo_root / "manifest.yaml")
# The v0 sample matrix is a frozen baseline: it covers exactly the
# six features the PRD shipped with, in their canonical order. The
# live manifest may carry additional rows (extensions added after
# v0 shipped), but the sample is derived only from the v0 slice so
# this test stays a meaningful regression gate for the v0 cell
# shape rather than chasing every new row added downstream.
v0_feature_ids = [
"basic_messaging_non_streaming",
"basic_messaging_streaming",
"tool_use",
"prompt_caching_5m",
"vision",
# Row 6 of the v0 PRD; originally shipped as `extended_thinking`.
# The id was renamed in-place to `thinking` to match Anthropic's
# current docs (which reserve "extended thinking" for the
# deprecated manual mode only). The row's *position* in v0 is
# the load-bearing invariant, not the id string.
"thinking",
]
v0_features = [
feature
for feature in full_manifest["features"]
if feature["id"] in v0_feature_ids
]
manifest = {**full_manifest, "features": v0_features}
feature_ids = [feature["id"] for feature in manifest["features"]]
providers = manifest["providers"]
models = ["claude-haiku-4-5", "claude-sonnet-4-5", "claude-opus-4-7"]
results = []
for feature_id in feature_ids:
for provider in providers:
for model in models:
results.append(
{
"feature_id": feature_id,
"provider": provider,
"nodeid": (
f"tests/e2e/claude_code/{feature_id}/test_{provider}.py"
f"::test[{model}]"
),
"result": {"status": "pass"},
}
)
matrix = build_matrix(
manifest=manifest,
results=results,
litellm_version="v1.83.0-stable",
claude_code_version="2.1.120",
generated_at="2026-04-25T00:00:00Z",
)
expected = json.loads((repo_root / "sample_compatibility-matrix.json").read_text())
assert matrix == expected
def test_build_matrix_1x5_grid_one_failing_model_breaks_cell():
"""If even one of three models fails on a provider, that cell is fail
and the error string carries the failing model id so the docs
tooltip can name the outlier."""
repo_root = Path(__file__).resolve().parents[1]
manifest = load_manifest(repo_root / "manifest.yaml")
results = [
{
"feature_id": "basic_messaging_non_streaming",
"provider": "bedrock_invoke",
"result": {"status": "pass"},
},
{
"feature_id": "basic_messaging_non_streaming",
"provider": "bedrock_invoke",
"result": {
"status": "fail",
"error": "[claude-opus-4-7-bedrock-invoke] claude CLI exited 1: throttled",
},
},
{
"feature_id": "basic_messaging_non_streaming",
"provider": "bedrock_invoke",
"result": {"status": "pass"},
},
]
matrix = build_matrix(
manifest=manifest,
results=results,
litellm_version="v",
claude_code_version="c",
generated_at="t",
)
cell = matrix["features"][0]["providers"]["bedrock_invoke"]
assert cell["status"] == "fail"
assert "claude-opus-4-7-bedrock-invoke" in cell["error"]
def test_build_from_paths_writes_output(tmp_path):
out = tmp_path / "compatibility-matrix.json"
matrix = build_from_paths(
manifest_path=FIXTURES / "manifest.yaml",
results_path=FIXTURES / "results.json",
litellm_version="v1.83.0-stable",
claude_code_version="2.1.120",
generated_at="2026-04-25T00:00:00Z",
output_path=out,
)
assert out.exists()
on_disk = json.loads(out.read_text())
assert on_disk == matrix
expected = json.loads((FIXTURES / "expected_matrix.json").read_text())
assert on_disk == expected

View file

@ -1,183 +0,0 @@
"""Structural tests for the full v0 6x5 matrix layout.
These tests don't run the `claude` CLI — they only verify that the
shape of the test suite on disk matches what the PRD declares: six
features in the prescribed order, and for each feature a directory
with one test file per provider column.
Catching layout drift here means the daily-cron VM and the PR gate
both see the same row set the docs page declares.
"""
from __future__ import annotations
from pathlib import Path
import pytest
import yaml
SUITE_ROOT = Path(__file__).resolve().parents[1]
MANIFEST_PATH = SUITE_ROOT / "manifest.yaml"
# The PRD's "Features in v0" section, in row order.
EXPECTED_FEATURE_IDS = [
"basic_messaging_non_streaming",
"basic_messaging_streaming",
"tool_use",
"prompt_caching_5m",
"vision",
# v0 originally shipped this row as `extended_thinking`. It was
# renamed in-place to `thinking` because Anthropic's docs reserve
# "extended thinking" for the deprecated manual API mode only; the
# single row exercises both manual and adaptive shapes since Claude
# Code picks per model. The PRD's "v0" identity is the *position*
# (row 6, 0-indexed 5), not the id string.
"thinking",
]
# The PRD's column order. Every feature directory must have one
# `test_<provider>.py` for each of these.
EXPECTED_PROVIDERS = [
"anthropic",
"bedrock_invoke",
"bedrock_converse",
"vertex_ai",
"azure",
]
def _all_manifest_feature_ids() -> list[str]:
"""Every feature_id currently declared in `manifest.yaml`.
Evaluated at import time so the result can drive parametrized
structural tests below. Used to catch layout drift on post-v0
feature rows added after the matrix shipped the v0 anchor
constants above only validate the original six rows by design.
"""
return [
feature["id"]
for feature in yaml.safe_load(MANIFEST_PATH.read_text())["features"]
]
ALL_FEATURE_IDS = _all_manifest_feature_ids()
@pytest.fixture(scope="module")
def manifest() -> dict:
return yaml.safe_load(MANIFEST_PATH.read_text())
def test_manifest_lists_all_six_v0_features_in_order(manifest):
"""The PRD's v0 row set must appear at the top of the manifest in
order. Features beyond v0 (extensions added after the matrix
shipped) are allowed but must not reorder or displace the v0
rows the docs page anchors row links by index, so v0 stays
pinned at positions [0:6] for the lifetime of the schema.
"""
ids = [feature["id"] for feature in manifest["features"]]
assert ids[: len(EXPECTED_FEATURE_IDS)] == EXPECTED_FEATURE_IDS
def test_manifest_lists_all_five_v0_providers_in_order(manifest):
assert manifest["providers"] == EXPECTED_PROVIDERS
def test_manifest_every_feature_has_human_readable_name(manifest):
for feature in manifest["features"]:
assert isinstance(feature["name"], str) and feature["name"].strip()
@pytest.mark.parametrize("feature_id", EXPECTED_FEATURE_IDS)
def test_feature_directory_exists(feature_id):
feature_dir = SUITE_ROOT / feature_id
assert feature_dir.is_dir(), f"missing feature directory: {feature_dir}"
@pytest.mark.parametrize("feature_id", EXPECTED_FEATURE_IDS)
@pytest.mark.parametrize("provider", EXPECTED_PROVIDERS)
def test_per_provider_test_file_exists(feature_id, provider):
test_file = SUITE_ROOT / feature_id / f"test_{provider}.py"
assert test_file.is_file(), f"missing per-provider test file: {test_file}"
@pytest.mark.parametrize("feature_id", EXPECTED_FEATURE_IDS)
def test_feature_directory_has_init_file(feature_id):
"""Each feature directory needs an __init__.py so pytest collects
the per-provider test files as a package matches the layout
established by `basic_messaging_non_streaming/`."""
init_file = SUITE_ROOT / feature_id / "__init__.py"
assert init_file.is_file(), f"missing __init__.py: {init_file}"
# Manifest-driven structural tests: every feature in `manifest.yaml`
# (v0 and post-v0 alike) must have the expected on-disk layout. The
# v0-only tests above pin the position of the original six rows; these
# extend the same structural guarantees to any row added afterward so
# a broken post-v0 directory still fails CI.
@pytest.mark.parametrize("feature_id", ALL_FEATURE_IDS)
def test_every_manifest_feature_has_directory(feature_id):
feature_dir = SUITE_ROOT / feature_id
assert feature_dir.is_dir(), (
f"manifest declares {feature_id!r} but {feature_dir} is missing — "
"feature_id MUST match its on-disk directory (see manifest.yaml header)."
)
@pytest.mark.parametrize("feature_id", ALL_FEATURE_IDS)
def test_every_manifest_feature_has_init_file(feature_id):
init_file = SUITE_ROOT / feature_id / "__init__.py"
assert init_file.is_file(), f"missing __init__.py: {init_file}"
@pytest.mark.parametrize("feature_id", ALL_FEATURE_IDS)
@pytest.mark.parametrize("provider", EXPECTED_PROVIDERS)
def test_every_manifest_feature_has_per_provider_test_file(feature_id, provider):
"""Every (feature, provider) cell in the rendered matrix must be
backed by a per-provider test file. Without this check, a missing
file silently becomes a `not_tested` cell in the published matrix
rather than a CI failure surfacing the layout drift."""
test_file = SUITE_ROOT / feature_id / f"test_{provider}.py"
assert test_file.is_file(), f"missing per-provider test file: {test_file}"
@pytest.mark.parametrize("feature_id", EXPECTED_FEATURE_IDS)
@pytest.mark.parametrize("provider", EXPECTED_PROVIDERS)
def test_per_provider_test_file_imports_and_parametrizes_three_models(
feature_id, provider
):
"""Every test file must reference the three Claude tiers required
by the PRD: Haiku 4.5, Sonnet 4.6, Opus 4.7. Implementations may
use plain aliases or per-provider-suffixed aliases (e.g.
`claude-opus-4-7-bedrock-invoke`), so we check for the tier
substrings rather than exact alias names."""
text = (SUITE_ROOT / feature_id / f"test_{provider}.py").read_text()
for tier in ("haiku-4-5", "sonnet-4-5", "opus-4-7"):
assert (
tier in text
), f"{feature_id}/test_{provider}.py does not reference {tier}"
@pytest.mark.parametrize("feature_id", EXPECTED_FEATURE_IDS)
def test_azure_test_file_drives_the_proxy(feature_id):
"""Azure (Microsoft Foundry) hosts Anthropic Claude as of 2025-11-18,
so every Azure cell in the v0 matrix exercises a real route through
the LiteLLM proxy same shape as the other provider columns. Pin
that here so a future regression doesn't silently revert these
cells to the old `not_applicable` boilerplate.
We accept either the direct `run_claude(...)` family of entrypoints
or a per-feature shared helper (e.g. `run_basic_messaging_cell`)
that wraps them both shapes drive the proxy, and we don't want
this layout pin to block legitimate de-duplication of test bodies.
"""
text = (SUITE_ROOT / feature_id / "test_azure.py").read_text()
assert "run_claude" in text or "run_basic_messaging_cell" in text, (
f"{feature_id}/test_azure.py must drive the claude CLI via run_claude() "
"or a shared helper that wraps it; the not_applicable stub was removed "
"when Foundry started hosting Claude."
)
assert '"status": "not_applicable"' not in text, (
f"{feature_id}/test_azure.py still reports not_applicable; Microsoft Foundry "
"now hosts Claude (Haiku 4.5, Sonnet 4.6, Opus 4.7), so this row must run."
)

View file

@ -1,32 +0,0 @@
"""Local conftest for the driver unit tests.
Installs a hermetic, no-op rate limiter for every test in this
subdirectory. Without this, importing `cli_driver` and calling
`run_claude(..., runner=fake)` would silently consume tokens from the
shared default limiter (which writes to `$TMPDIR/...`), polluting the
on-disk state another test run might rely on and adding flakiness if
the env vars say "rate=0.1/s".
A no-op limiter (rate=0 for every provider) returns immediately from
`acquire(...)`, so unit tests behave exactly as they did before the
limiter was added.
"""
from __future__ import annotations
import pytest
from claude_code.rate_limiter import (
ALL_PROVIDERS,
ProviderConfig,
RateLimiter,
use_limiter,
)
@pytest.fixture(autouse=True)
def _hermetic_rate_limiter(tmp_path):
config = {p: ProviderConfig(rate_per_sec=0.0, burst=0.0) for p in ALL_PROVIDERS}
limiter = RateLimiter(config=config, state_dir=tmp_path)
with use_limiter(limiter):
yield

View file

@ -1,229 +0,0 @@
"""Unit tests for the shared `run_basic_messaging_cell` helper.
These tests inject a fake ``ClaudeRunner`` and a fake env mapping so
they exercise the helper's branching (env-missing guard, per-model
pass/fail/empty-text, streaming wire check) without spawning the real
CLI or touching ``os.environ``. The streaming check is the regression
we care about: a proxy that buffers the upstream stream must turn the
cell red, not green.
"""
from __future__ import annotations
from typing import Any, Dict, List, Mapping, Optional, Sequence
import pytest
from claude_code._basic_messaging import (
MIN_STREAM_DELTA_EVENTS,
_count_stream_event_deltas,
run_basic_messaging_cell,
)
from claude_code.cli_driver import DriverResult
_PROXY_ENV: Mapping[str, str] = {
"LITELLM_PROXY_URL": "http://localhost:4000",
"LITELLM_MASTER_KEY": "sk-test",
}
class _FakeResult:
"""Stand-in for the test's `compat_result` fixture.
Records every `set` / `add` payload so assertions can inspect what
the cell reported, in order, without needing the real
`pytest_runtest_logreport` plumbing from `conftest.py`.
"""
def __init__(self) -> None:
self.rows: List[Dict[str, Any]] = []
self.single: Optional[Dict[str, Any]] = None
def set(self, payload: Mapping[str, Any]) -> None:
self.single = dict(payload)
def add(self, payload: Mapping[str, Any]) -> None:
self.rows.append(dict(payload))
def _streamed_events(n_deltas: int = 5) -> List[Dict[str, Any]]:
"""Build a stream-json event list that *looks* streamed.
Includes `n_deltas` `stream_event` records (matching what
`--include-partial-messages` produces) plus the usual
`system`/`assistant`/`result` boilerplate the CLI always emits.
"""
events: List[Dict[str, Any]] = [{"type": "system", "subtype": "init"}]
for i in range(n_deltas):
events.append(
{
"type": "stream_event",
"event": {
"type": "content_block_delta",
"delta": {"type": "text_delta", "text": str(i)},
},
}
)
events.append(
{
"type": "assistant",
"message": {"content": [{"type": "text", "text": "1\n2\n3"}]},
}
)
events.append({"type": "result"})
return events
def _buffered_events() -> List[Dict[str, Any]]:
"""Event list a buffering proxy would produce: zero `stream_event`s."""
return [
{"type": "system", "subtype": "init"},
{
"type": "assistant",
"message": {"content": [{"type": "text", "text": "1\n2\n3"}]},
},
{"type": "result"},
]
def _make_fake_runner(*, outcomes_by_model):
"""Build an injectable runner that returns canned outcomes and
records the kwargs the helper passed in.
Returns a ``(runner, captured)`` pair; ``captured`` is a dict the
test can assert against without any global mutation, which is why
we prefer DI over ``monkeypatch.setattr``: the helper takes a
``runner=`` kwarg, so tests bind their fake directly."""
captured: Dict[str, Any] = {}
def runner(*, models, prompt, base_url, api_key, extra_args=None, **_kwargs):
captured["models"] = list(models)
captured["prompt"] = prompt
captured["base_url"] = base_url
captured["api_key"] = api_key
captured["extra_args"] = list(extra_args) if extra_args else []
return {model: outcomes_by_model[model] for model in models}
return runner, captured
def test_count_stream_event_deltas_only_counts_records_with_event_payload():
events = [
{"type": "system"},
{"type": "stream_event", "event": {"type": "message_start"}},
{"type": "stream_event", "event": {"type": "content_block_delta"}},
{"type": "stream_event"},
{"type": "stream_event", "event": None},
{"type": "stream_event", "event": "not-a-dict"},
{"type": "assistant"},
{"type": "result"},
]
assert _count_stream_event_deltas(events) == 2
def test_verify_streaming_passes_when_proxy_streams():
fake_result = _FakeResult()
model = "claude-haiku-4-5"
outcome = DriverResult(text="1\n2\n3", events=_streamed_events(n_deltas=5))
runner, captured = _make_fake_runner(outcomes_by_model={model: outcome})
run_basic_messaging_cell(
compat_result=fake_result,
models=[model],
prompt="Count from 1 to 5, one number per line.",
verify_streaming=True,
env=_PROXY_ENV,
runner=runner,
)
assert captured["extra_args"] == ["--include-partial-messages"]
assert fake_result.rows == [{"status": "pass"}]
def test_verify_streaming_fails_when_proxy_buffers():
fake_result = _FakeResult()
model = "claude-haiku-4-5"
outcome = DriverResult(text="1\n2\n3", events=_buffered_events())
runner, _captured = _make_fake_runner(outcomes_by_model={model: outcome})
with pytest.raises(pytest.fail.Exception):
run_basic_messaging_cell(
compat_result=fake_result,
models=[model],
prompt="Count from 1 to 5, one number per line.",
verify_streaming=True,
env=_PROXY_ENV,
runner=runner,
)
assert len(fake_result.rows) == 1
row = fake_result.rows[0]
assert row["status"] == "fail"
assert "stream_event" in row["error"]
assert f"< {MIN_STREAM_DELTA_EVENTS}" in row["error"]
def test_non_streaming_variant_omits_partial_messages_flag():
"""Default `verify_streaming=False` keeps the non-streaming wire identical."""
fake_result = _FakeResult()
model = "claude-haiku-4-5"
outcome = DriverResult(text="pong", events=_buffered_events())
runner, captured = _make_fake_runner(outcomes_by_model={model: outcome})
run_basic_messaging_cell(
compat_result=fake_result,
models=[model],
prompt="Reply with the single word 'pong' and nothing else.",
env=_PROXY_ENV,
runner=runner,
)
assert captured["extra_args"] == []
assert fake_result.rows == [{"status": "pass"}]
def test_verify_streaming_requires_all_models_to_stream():
"""If any one tier buffers, the cell fails — same all-must-pass shape as
the non-streaming check."""
fake_result = _FakeResult()
outcomes = {
"claude-haiku-4-5": DriverResult(text="ok", events=_streamed_events(5)),
"claude-sonnet-4-5": DriverResult(text="ok", events=_buffered_events()),
"claude-opus-4-7": DriverResult(text="ok", events=_streamed_events(5)),
}
runner, _captured = _make_fake_runner(outcomes_by_model=outcomes)
with pytest.raises(pytest.fail.Exception):
run_basic_messaging_cell(
compat_result=fake_result,
models=list(outcomes.keys()),
prompt="Count from 1 to 5, one number per line.",
verify_streaming=True,
env=_PROXY_ENV,
runner=runner,
)
statuses = [row["status"] for row in fake_result.rows]
assert statuses == ["pass", "fail", "pass"]
def test_missing_proxy_env_hard_fails_regardless_of_runner():
"""The env guard fires before the runner is called, and takes the
env from the injected mapping (not os.environ). Passing an empty
env dict must hard-fail even if a happy runner is bound."""
fake_result = _FakeResult()
runner, captured = _make_fake_runner(outcomes_by_model={})
with pytest.raises(pytest.fail.Exception):
run_basic_messaging_cell(
compat_result=fake_result,
models=["claude-haiku-4-5"],
prompt="whatever",
env={},
runner=runner,
)
assert captured == {}, "runner must not be called when env resolution fails"

View file

@ -1,992 +0,0 @@
"""Unit tests for the Claude Code CLI Driver.
These tests mock the subprocess so they run anywhere no network, no
`claude` install, no API keys. They cover the behavior contract:
argument assembly, environment overlay, stream-JSON parsing, exit-code
plumbing, and the structured failure modes (CLI not found, timeout).
"""
from __future__ import annotations
import json
import subprocess
from dataclasses import dataclass
from typing import List, Optional
import pytest
from claude_code.cli_driver import (
ClaudeCLIError,
DriverResult,
failure_diagnostic,
is_rate_limit_shaped,
run_claude,
run_claude_models_parallel,
)
@dataclass
class _Completed:
returncode: int = 0
stdout: str = ""
stderr: str = ""
def _make_runner(*, stdout: str = "", returncode: int = 0, stderr: str = ""):
captured = {}
def runner(cmd, env, capture_output, text, timeout, check, input=None):
captured["cmd"] = cmd
captured["env"] = env
captured["timeout"] = timeout
captured["input"] = input
return _Completed(returncode=returncode, stdout=stdout, stderr=stderr)
return runner, captured
def test_run_claude_assembles_command_correctly():
runner, captured = _make_runner(
stdout='{"type":"assistant","message":{"content":[{"type":"text","text":"ok"}]}}\n'
)
run_claude(
prompt="hello",
model="claude-haiku-4-5",
base_url="http://localhost:4000",
api_key="sk-test",
runner=runner,
)
cmd = captured["cmd"]
assert cmd[0] == "claude"
assert "--print" in cmd
assert "--output-format" in cmd
assert "stream-json" in cmd
assert "--model" in cmd
assert "claude-haiku-4-5" in cmd
# prompt is the last positional after the `--` end-of-options marker.
assert cmd[-2:] == ["--", "hello"]
def test_run_claude_places_extra_args_before_prompt():
"""`claude --print` expects the prompt as the final positional arg.
Flags appearing after the prompt are ignored or eaten by the prompt
parser (especially variadic flags like `--allowed-tools <tools...>`),
which silently broke the tool_use, vision, and web_search cells
before the fix. Pin the ordering: every flag (including
caller-supplied `extra_args`) must precede the `--` end-of-options
marker, which itself precedes the prompt.
"""
runner, captured = _make_runner(stdout="")
run_claude(
prompt="say hi",
model="claude-haiku-4-5",
base_url="http://localhost:4000",
api_key="sk-test",
extra_args=["--allowed-tools", "Bash"],
runner=runner,
)
cmd = captured["cmd"]
# Prompt is last, `--` immediately precedes it, and the caller's
# extra_args sit somewhere earlier in the command.
assert cmd[-2:] == ["--", "say hi"]
assert "--allowed-tools" in cmd
assert cmd.index("--allowed-tools") < cmd.index("--")
def test_run_claude_overlays_proxy_env():
runner, captured = _make_runner(stdout="")
run_claude(
prompt="hi",
model="claude-opus-4-7",
base_url="http://proxy.example:4000",
api_key="sk-abc",
runner=runner,
)
env = captured["env"]
assert env["ANTHROPIC_BASE_URL"] == "http://proxy.example:4000"
assert env["ANTHROPIC_AUTH_TOKEN"] == "sk-abc"
def test_run_claude_extra_env_is_added_to_subprocess_env():
"""Caller-supplied extra_env entries land on the subprocess env."""
runner, captured = _make_runner(stdout="")
run_claude(
prompt="hi",
model="claude-opus-4-7",
base_url="http://localhost",
api_key="sk-abc",
extra_env={"MAX_THINKING_TOKENS": "4096"},
runner=runner,
)
assert captured["env"]["MAX_THINKING_TOKENS"] == "4096"
def test_run_claude_inherits_only_allowlisted_os_environ(monkeypatch):
"""Process-runtime vars (PATH) flow through; credentials don't.
The `claude` CLI is a Node binary installed dynamically from npm in
CI. If the package were ever compromised, inheriting the entire
parent environment would hand it every credential the surrounding
proxy job loads (AWS keys, Azure Foundry key, GitHub token, etc.).
Pin the contract: only the small allowlist of runtime vars is
inherited; everything else is dropped unless the caller passes it
explicitly via extra_env.
`HOME` is *not* on the allowlist anymore see the dedicated
isolated-HOME test below for the reason.
"""
monkeypatch.setenv("PATH", "/usr/bin:/usr/local/bin")
monkeypatch.setenv("HOME", "/home/runner")
monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "totally-secret")
monkeypatch.setenv("ANTHROPIC_API_KEY", "sk-proxy-only")
monkeypatch.setenv("AZURE_AI_API_KEY", "azure-secret")
monkeypatch.setenv("VERTEXAI_CREDENTIALS", '{"private_key": "leak"}')
monkeypatch.setenv("GITHUB_TOKEN", "ghs_xxx")
runner, captured = _make_runner(stdout="")
run_claude(
prompt="hi",
model="claude-opus-4-7",
base_url="http://localhost",
api_key="sk-abc",
runner=runner,
)
env = captured["env"]
assert env["PATH"] == "/usr/bin:/usr/local/bin"
assert "AWS_SECRET_ACCESS_KEY" not in env
assert "ANTHROPIC_API_KEY" not in env
assert "AZURE_AI_API_KEY" not in env
assert "VERTEXAI_CREDENTIALS" not in env
assert "GITHUB_TOKEN" not in env
def test_run_claude_uses_isolated_per_invocation_home(monkeypatch, tmp_path):
"""`claude` subprocess never sees the runtime user's real $HOME.
The CLI needs *a* HOME (it caches per-session state under
`$HOME/.claude/projects/<sha>/`), but it has no business reading
the runtime user's real one. On the cron VM the runtime user is a
real interactive account with a populated home directory
(~/.config/gh/hosts.yml carrying a GitHub token, ~/.ssh/, etc.);
handing /home/mateo to a compromised npm package or to a
model-directed `Read` tool call during the PDF/vision cells
would let it exfiltrate those files. We hand the CLI a fresh
empty per-invocation tmpdir instead.
"""
monkeypatch.setenv("HOME", "/home/runner")
runner, captured = _make_runner(stdout="")
run_claude(
prompt="hi",
model="claude-opus-4-7",
base_url="http://localhost",
api_key="sk-abc",
runner=runner,
)
env = captured["env"]
assert "HOME" in env, "claude CLI needs HOME to find ~/.claude session dir"
assert (
env["HOME"] != "/home/runner"
), "HOME must not leak the parent process's HOME to claude"
# The isolated HOME is a fresh tmpdir prefixed `claude-cli-home-`;
# see `_make_isolated_home` in cli_driver.py. It exists during the
# subprocess call and is removed afterwards (cleanup runs in a
# `finally`, so by the time this assertion runs the dir is gone —
# we only check the *prefix* of the path string we captured).
assert "claude-cli-home-" in env["HOME"]
def test_run_claude_isolated_home_is_distinct_per_invocation(monkeypatch):
"""Two consecutive calls get two different isolated HOMEs.
Reusing a single tmpdir across calls would defeat the isolation
in the parallel matrix run (a compromised CLI could plant a file
in HOME on one model's run and read it on the next). Pin: each
`run_claude` invocation gets its own freshly-created HOME.
"""
monkeypatch.setenv("HOME", "/home/runner")
runner, captured = _make_runner(stdout="")
run_claude(
prompt="hi",
model="claude-opus-4-7",
base_url="http://localhost",
api_key="sk-abc",
runner=runner,
)
home_a = captured["env"]["HOME"]
runner, captured = _make_runner(stdout="")
run_claude(
prompt="hi",
model="claude-opus-4-7",
base_url="http://localhost",
api_key="sk-abc",
runner=runner,
)
home_b = captured["env"]["HOME"]
assert home_a != home_b
def test_run_claude_isolated_home_cleaned_up_after_run(monkeypatch):
"""The per-invocation HOME tmpdir is rm-rf'd when run_claude returns.
Without cleanup, a long matrix run would accumulate one tmpdir
per cell × per model × per CLI call (~75 dirs per cron run,
growing without bound across days).
"""
import os as _os
monkeypatch.setenv("HOME", "/home/runner")
runner, captured = _make_runner(stdout="")
run_claude(
prompt="hi",
model="claude-opus-4-7",
base_url="http://localhost",
api_key="sk-abc",
runner=runner,
)
isolated_home = captured["env"]["HOME"]
assert not _os.path.exists(
isolated_home
), f"isolated HOME {isolated_home!r} should be removed after run_claude returns"
def test_run_claude_isolated_home_cleaned_up_on_subprocess_failure(monkeypatch):
"""Cleanup runs even when the CLI subprocess raises.
If the CLI is missing or times out, `run_claude` raises
`ClaudeCLIError` but the per-invocation HOME tmpdir must still
be removed (the `finally` clause), otherwise long failure-prone
runs leak tmpdirs.
"""
import os as _os
monkeypatch.setenv("HOME", "/home/runner")
captured: dict = {}
def runner(cmd, env, capture_output, text, timeout, check, input=None):
captured["env"] = env
raise subprocess.TimeoutExpired(cmd=cmd, timeout=timeout)
with pytest.raises(ClaudeCLIError):
run_claude(
prompt="hi",
model="claude-opus-4-7",
base_url="http://localhost",
api_key="sk-abc",
runner=runner,
)
isolated_home = captured["env"]["HOME"]
assert not _os.path.exists(
isolated_home
), f"isolated HOME {isolated_home!r} should be removed even on timeout"
def test_run_claude_extra_env_can_pass_through_otherwise_blocked_var(monkeypatch):
"""The allowlist applies to inherited os.environ; extra_env is the
sanctioned way for a test to opt-in to passing something extra."""
monkeypatch.setenv("ANTHROPIC_API_KEY", "from-os")
runner, captured = _make_runner(stdout="")
run_claude(
prompt="hi",
model="claude-opus-4-7",
base_url="http://localhost",
api_key="sk-abc",
extra_env={"ANTHROPIC_API_KEY": "from-arg"},
runner=runner,
)
assert captured["env"]["ANTHROPIC_API_KEY"] == "from-arg"
def test_run_claude_parses_stream_json_assistant_text():
events = [
{"type": "system", "session_id": "abc"},
{
"type": "assistant",
"message": {
"content": [
{"type": "text", "text": "Hello "},
{"type": "text", "text": "world"},
]
},
},
{"type": "result", "usage": {"input_tokens": 10, "output_tokens": 2}},
]
stdout = "\n".join(json.dumps(e) for e in events) + "\n"
runner, _ = _make_runner(stdout=stdout)
result = run_claude(
prompt="hi",
model="claude-haiku-4-5",
base_url="http://localhost",
api_key="sk-abc",
runner=runner,
)
assert isinstance(result, DriverResult)
assert result.text == "Hello world"
assert len(result.events) == 3
assert result.usage == {"input_tokens": 10, "output_tokens": 2}
assert result.exit_code == 0
def test_run_claude_handles_string_message_content():
"""Some CLI versions emit `message.content` as a plain string."""
stdout = (
json.dumps({"type": "assistant", "message": {"content": "bare text"}}) + "\n"
)
runner, _ = _make_runner(stdout=stdout)
result = run_claude(
prompt="hi",
model="m",
base_url="http://x",
api_key="k",
runner=runner,
)
assert result.text == "bare text"
def test_run_claude_skips_malformed_lines():
stdout = (
"not-json\n"
+ json.dumps(
{
"type": "assistant",
"message": {"content": [{"type": "text", "text": "x"}]},
}
)
+ "\n"
+ "{also-bad\n"
)
runner, _ = _make_runner(stdout=stdout)
result = run_claude(
prompt="hi",
model="m",
base_url="http://x",
api_key="k",
runner=runner,
)
assert result.text == "x"
assert len(result.events) == 1
def test_run_claude_propagates_nonzero_exit_code():
runner, _ = _make_runner(stdout="", returncode=2, stderr="auth failed")
result = run_claude(
prompt="hi",
model="m",
base_url="http://x",
api_key="k",
runner=runner,
)
assert result.exit_code == 2
assert result.stderr == "auth failed"
assert result.text == ""
def test_run_claude_raises_on_missing_cli():
def runner(*args, **kwargs):
raise FileNotFoundError(2, "no such file", "claude")
with pytest.raises(ClaudeCLIError, match="claude CLI not found"):
run_claude(
prompt="hi",
model="m",
base_url="http://x",
api_key="k",
runner=runner,
)
def test_run_claude_raises_on_timeout():
def runner(*args, **kwargs):
raise subprocess.TimeoutExpired(cmd="claude", timeout=1)
with pytest.raises(ClaudeCLIError, match="timed out"):
run_claude(
prompt="hi",
model="m",
base_url="http://x",
api_key="k",
timeout=1,
runner=runner,
)
def test_run_claude_validates_required_params():
runner, _ = _make_runner()
with pytest.raises(ValueError, match="prompt"):
run_claude(
prompt="",
model="m",
base_url="http://x",
api_key="k",
runner=runner,
)
with pytest.raises(ValueError, match="stdin_input"):
run_claude(
prompt=None,
stdin_input="",
model="m",
base_url="http://x",
api_key="k",
runner=runner,
)
with pytest.raises(ValueError, match="model"):
run_claude(
prompt="hi",
model="",
base_url="http://x",
api_key="k",
runner=runner,
)
with pytest.raises(ValueError, match="base_url"):
run_claude(
prompt="hi",
model="m",
base_url="",
api_key="k",
runner=runner,
)
with pytest.raises(ValueError, match="api_key"):
run_claude(
prompt="hi",
model="m",
base_url="http://x",
api_key="",
runner=runner,
)
# ---------------------------------------------------------------------------
# failure_diagnostic
#
# Regression coverage for the bring-up incident where the proxy was started
# with the wrong config and tests reported only `claude CLI exited 1` while
# the actual 400 from LiteLLM was sitting in stdout. The helper must surface
# api_status, the assistant text (where API errors land), stderr, and the
# exit code together — and gracefully degrade when individual pieces are
# missing.
# ---------------------------------------------------------------------------
def test_failure_diagnostic_surfaces_api_error_text_from_stdout():
"""The CLI hides 4xx/5xx from the proxy in `assistant.message.content` text."""
api_error_text = (
'API Error: 400 {"error":{"message":"litellm.BadRequestError: '
"You passed in model=claude-haiku-4-5. There are no healthy "
'deployments..."}}'
)
result = DriverResult(
text=api_error_text,
events=[
{
"type": "assistant",
"message": {"content": [{"type": "text", "text": api_error_text}]},
},
{
"type": "result",
"is_error": True,
"api_error_status": 400,
"result": api_error_text,
},
],
exit_code=1,
stderr="",
)
diag = failure_diagnostic(result)
assert "exit=1" in diag
assert "api_status=400" in diag
assert "There are no healthy deployments" in diag
def test_failure_diagnostic_falls_back_to_stderr_when_no_text():
result = DriverResult(text="", events=[], exit_code=2, stderr="boom\n")
diag = failure_diagnostic(result)
assert "exit=2" in diag
assert "stderr=boom" in diag
def test_failure_diagnostic_handles_completely_empty_result():
"""A run that produced literally nothing should still yield a useful string."""
result = DriverResult(text="", events=[], exit_code=137, stderr="")
diag = failure_diagnostic(result)
assert "exit=137" in diag
assert "no diagnostic output" in diag
def test_failure_diagnostic_truncates_long_text():
"""Don't let a 5MB HTML 502 page from a load balancer wreck the matrix JSON."""
huge = "x" * 5000
result = DriverResult(text=huge, events=[], exit_code=1, stderr="")
diag = failure_diagnostic(result, max_len=100)
assert "truncated" in diag
# Allow some slack for the prefix/suffix/separator characters.
assert len(diag) < 300
def test_failure_diagnostic_ignores_non_int_api_error_status():
"""The CLI sometimes emits api_error_status as a string; don't crash."""
result = DriverResult(
text="oops",
events=[{"type": "result", "api_error_status": "n/a"}],
exit_code=1,
stderr="",
)
diag = failure_diagnostic(result)
assert "api_status" not in diag
assert "text=oops" in diag
# ---------------------------------------------------------------------------
# run_claude_models_parallel
#
# The matrix runs three Claude tiers per cell, so the parallel helper has to
# (a) invoke `run_claude` once per model, (b) preserve each model's outcome
# separately, and (c) return errors as values rather than raising — callers
# need both the failed and the succeeded model results to report per-cell
# rows accurately.
# ---------------------------------------------------------------------------
def test_run_claude_models_parallel_returns_one_result_per_model():
"""Each model gets its own DriverResult keyed under the helper's dict."""
seen_models: List[str] = []
def runner(cmd, env, capture_output, text, timeout, check, input=None):
# The model id is two slots after `--model` in the assembled command.
idx = cmd.index("--model")
model = cmd[idx + 1]
seen_models.append(model)
return _Completed(
returncode=0,
stdout=json.dumps(
{
"type": "assistant",
"message": {
"content": [{"type": "text", "text": f"reply-{model}"}]
},
}
)
+ "\n",
)
outcomes = run_claude_models_parallel(
models=["a", "b", "c"],
prompt="hi",
base_url="http://x",
api_key="k",
runner=runner,
)
assert set(outcomes.keys()) == {"a", "b", "c"}
for model in ("a", "b", "c"):
result = outcomes[model]
assert isinstance(result, DriverResult)
assert result.text == f"reply-{model}"
assert result.exit_code == 0
assert sorted(seen_models) == ["a", "b", "c"]
def test_run_claude_models_parallel_returns_errors_as_values():
"""A model whose CLI is missing surfaces as a ClaudeCLIError, not a raise."""
def runner(cmd, env, capture_output, text, timeout, check, input=None):
idx = cmd.index("--model")
model = cmd[idx + 1]
if model == "boom":
raise FileNotFoundError(2, "no such file", "claude")
return _Completed(
returncode=0,
stdout=json.dumps(
{
"type": "assistant",
"message": {"content": [{"type": "text", "text": "ok"}]},
}
)
+ "\n",
)
outcomes = run_claude_models_parallel(
models=["ok-model", "boom"],
prompt="hi",
base_url="http://x",
api_key="k",
runner=runner,
)
assert isinstance(outcomes["ok-model"], DriverResult)
assert outcomes["ok-model"].text == "ok"
assert isinstance(outcomes["boom"], ClaudeCLIError)
assert "claude CLI not found" in str(outcomes["boom"])
def test_run_claude_models_parallel_preserves_nonzero_exit_codes():
"""Mixed success/failure on exit code should not collapse into one verdict."""
def runner(cmd, env, capture_output, text, timeout, check, input=None):
idx = cmd.index("--model")
model = cmd[idx + 1]
if model == "fail":
return _Completed(returncode=2, stdout="", stderr="auth failed")
return _Completed(
returncode=0,
stdout=json.dumps(
{
"type": "assistant",
"message": {"content": [{"type": "text", "text": "ok"}]},
}
)
+ "\n",
)
outcomes = run_claude_models_parallel(
models=["ok-model", "fail"],
prompt="hi",
base_url="http://x",
api_key="k",
runner=runner,
)
assert outcomes["ok-model"].exit_code == 0
assert outcomes["fail"].exit_code == 2
assert outcomes["fail"].stderr == "auth failed"
def test_run_claude_models_parallel_rejects_empty_models():
with pytest.raises(ValueError, match="non-empty"):
run_claude_models_parallel(
models=[],
prompt="hi",
base_url="http://x",
api_key="k",
)
def test_run_claude_models_parallel_stamps_duration_on_each_result():
"""Each DriverResult carries the per-model wall time so callers can
attribute slow cells without re-timing the work themselves.
The fake runner sleeps for very different durations per model so
we can prove each result is timing its own work (not the batch
wall time). We use generous absolute bounds because thread-pool
scheduling on a loaded CI box adds noise on the order of tens of
milliseconds.
"""
import time
def runner(cmd, env, capture_output, text, timeout, check, input=None):
idx = cmd.index("--model")
model = cmd[idx + 1]
time.sleep(0.05 if model == "fast" else 0.40)
return _Completed(returncode=0, stdout="")
outcomes = run_claude_models_parallel(
models=["fast", "slow"],
prompt="hi",
base_url="http://x",
api_key="k",
runner=runner,
)
fast_ms = outcomes["fast"].duration_ms
slow_ms = outcomes["slow"].duration_ms
assert fast_ms is not None and slow_ms is not None
# 50ms sleep ⇒ ~50250ms after scheduling overhead; 400ms sleep ⇒
# 400700ms. We just need the two distributions to be non-overlapping
# so we know each row's duration is its own work, not the batch's.
assert fast_ms < 300, fast_ms
assert slow_ms >= 350, slow_ms
assert slow_ms > fast_ms
def test_run_claude_models_parallel_breakdown_logs_to_stderr(capsys):
"""The breakdown helper must emit a per-model timing block so users
can answer "why didn't parallel help?" without re-instrumenting."""
def runner(cmd, env, capture_output, text, timeout, check, input=None):
return _Completed(returncode=0, stdout="")
run_claude_models_parallel(
models=["model-x", "model-y"],
prompt="hi",
base_url="http://x",
api_key="k",
runner=runner,
)
captured = capsys.readouterr()
assert "[parallel] per-model wall time:" in captured.err
assert "model-x" in captured.err
assert "model-y" in captured.err
assert "speedup=" in captured.err
assert "slowest=" in captured.err
def test_run_claude_models_parallel_breakdown_marks_cli_errors(capsys):
"""When a model raises ClaudeCLIError, the breakdown should still
show its row tagged as `cli-error` rather than crashing or omitting it."""
def runner(cmd, env, capture_output, text, timeout, check, input=None):
idx = cmd.index("--model")
if cmd[idx + 1] == "boom":
raise FileNotFoundError(2, "no such file", "claude")
return _Completed(returncode=0, stdout="")
run_claude_models_parallel(
models=["ok-model", "boom"],
prompt="hi",
base_url="http://x",
api_key="k",
runner=runner,
)
captured = capsys.readouterr()
assert "ok-model" in captured.err
assert "boom" in captured.err
assert "cli-error" in captured.err
def test_run_claude_models_parallel_forwards_extra_args_and_env():
"""Shared kwargs must reach every per-model invocation unchanged."""
captured_envs: List[dict] = []
captured_cmds: List[List[str]] = []
def runner(cmd, env, capture_output, text, timeout, check, input=None):
captured_envs.append(env)
captured_cmds.append(cmd)
return _Completed(returncode=0, stdout="")
run_claude_models_parallel(
models=["a", "b"],
prompt="hi",
base_url="http://x",
api_key="k",
extra_env={"MAX_THINKING_TOKENS": "4096"},
extra_args=["--allowed-tools", "Bash"],
runner=runner,
)
assert all(env["MAX_THINKING_TOKENS"] == "4096" for env in captured_envs)
for cmd in captured_cmds:
assert "--allowed-tools" in cmd
assert "Bash" in cmd
def test_failure_diagnostic_uses_last_result_event_status():
"""If multiple `result` events appear, the most recent status wins."""
result = DriverResult(
text="",
events=[
{"type": "result", "api_error_status": 500},
{"type": "assistant", "message": {"content": []}},
{"type": "result", "api_error_status": 429},
],
exit_code=1,
stderr="",
)
diag = failure_diagnostic(result)
assert "api_status=429" in diag
assert "500" not in diag
_RATE_LIMITED_STDOUT = (
json.dumps(
{
"type": "assistant",
"message": {
"content": [
{"type": "text", "text": "API Error: 429 Too Many Requests"}
]
},
}
)
+ "\n"
+ json.dumps({"type": "result", "api_error_status": 429})
+ "\n"
)
_OK_STDOUT = (
json.dumps(
{
"type": "assistant",
"message": {"content": [{"type": "text", "text": "pong"}]},
}
)
+ "\n"
)
class _FlakyRunner:
"""Fake runner that rate-limits each model N times before succeeding.
Keeps a per-model call count so tests can assert exactly how many
attempts the retry loop made the load-bearing detail a canned
single-response runner can't express.
"""
def __init__(self, failures_before_success: dict):
self.failures_before_success = dict(failures_before_success)
self.calls: dict = {}
def __call__(self, cmd, env, capture_output, text, timeout, check, input=None):
model = cmd[cmd.index("--model") + 1]
self.calls[model] = self.calls.get(model, 0) + 1
if self.calls[model] <= self.failures_before_success.get(model, 0):
return _Completed(returncode=1, stdout=_RATE_LIMITED_STDOUT)
return _Completed(returncode=0, stdout=_OK_STDOUT)
@pytest.mark.parametrize(
"outcome,expected",
[
(ClaudeCLIError("claude CLI timed out after 120.0s"), True),
(ClaudeCLIError("claude CLI not found at 'claude'"), False),
(
DriverResult(
text="",
events=[{"type": "result", "api_error_status": 429}],
exit_code=1,
),
True,
),
(DriverResult(text="Too Many Requests", exit_code=1), True),
(DriverResult(text="", stderr="throttled by upstream", exit_code=1), True),
(DriverResult(text="rate limit exceeded", exit_code=0), False),
(DriverResult(text="", stderr="auth failed", exit_code=2), False),
],
)
def test_is_rate_limit_shaped_classification(outcome, expected):
"""The retry trigger must match 429/throttle/timeout markers on
failures only a passing result mentioning '429' in its reply text
must never be classified as retryable."""
assert is_rate_limit_shaped(outcome) is expected
def test_run_claude_models_parallel_retries_rate_limited_model_until_success():
"""A model that 429s once must be retried after the backoff sleep and
end up green, while an untroubled sibling model runs exactly once."""
runner = _FlakyRunner({"flaky": 1})
sleeps: List[float] = []
outcomes = run_claude_models_parallel(
models=["flaky", "steady"],
prompt="hi",
base_url="http://x",
api_key="k",
runner=runner,
rate_limit_retries=2,
rate_limit_backoff_seconds=0.5,
sleep=sleeps.append,
)
assert isinstance(outcomes["flaky"], DriverResult)
assert outcomes["flaky"].exit_code == 0
assert outcomes["flaky"].text == "pong"
assert runner.calls == {"flaky": 2, "steady": 1}
assert sleeps == [0.5]
def test_run_claude_models_parallel_does_not_retry_non_rate_limit_failures():
"""A deterministic failure (bad auth) must fail fast: no sleeps, one
attempt retrying it would just triple the matrix wall time."""
def runner(cmd, env, capture_output, text, timeout, check, input=None):
return _Completed(returncode=2, stdout="", stderr="auth failed")
sleeps: List[float] = []
outcomes = run_claude_models_parallel(
models=["a"],
prompt="hi",
base_url="http://x",
api_key="k",
runner=runner,
rate_limit_retries=2,
rate_limit_backoff_seconds=0.5,
sleep=sleeps.append,
)
assert outcomes["a"].exit_code == 2
assert sleeps == []
def test_run_claude_models_parallel_returns_last_failure_when_retries_exhausted():
"""A persistently rate-limited model exhausts its budget (initial
attempt + N retries, each preceded by one backoff sleep) and still
surfaces the 429 diagnostic instead of masking it."""
runner = _FlakyRunner({"stuck": 99})
sleeps: List[float] = []
outcomes = run_claude_models_parallel(
models=["stuck"],
prompt="hi",
base_url="http://x",
api_key="k",
runner=runner,
rate_limit_retries=2,
rate_limit_backoff_seconds=0.25,
sleep=sleeps.append,
)
assert runner.calls == {"stuck": 3}
assert sleeps == [0.25, 0.25]
assert outcomes["stuck"].exit_code == 1
assert "429" in failure_diagnostic(outcomes["stuck"])
def test_run_claude_models_parallel_retries_timeout_shaped_cli_errors():
"""CLI timeouts are how saturated upstreams usually present (the CLI
retries 429s internally until the harness kills it), so a timeout
must be retried like an explicit 429."""
calls: List[int] = []
def runner(cmd, env, capture_output, text, timeout, check, input=None):
calls.append(1)
if len(calls) == 1:
raise subprocess.TimeoutExpired(cmd="claude", timeout=1)
return _Completed(returncode=0, stdout=_OK_STDOUT)
sleeps: List[float] = []
outcomes = run_claude_models_parallel(
models=["a"],
prompt="hi",
base_url="http://x",
api_key="k",
runner=runner,
rate_limit_retries=1,
rate_limit_backoff_seconds=0.5,
sleep=sleeps.append,
)
assert isinstance(outcomes["a"], DriverResult)
assert outcomes["a"].text == "pong"
assert len(calls) == 2
assert sleeps == [0.5]
def test_run_claude_models_parallel_zero_retries_disables_backoff():
"""`rate_limit_retries=0` must restore the old single-attempt
behavior exactly: one call, no sleeps, failure returned as-is."""
runner = _FlakyRunner({"stuck": 99})
sleeps: List[float] = []
outcomes = run_claude_models_parallel(
models=["stuck"],
prompt="hi",
base_url="http://x",
api_key="k",
runner=runner,
rate_limit_retries=0,
rate_limit_backoff_seconds=0.5,
sleep=sleeps.append,
)
assert runner.calls == {"stuck": 1}
assert sleeps == []
assert outcomes["stuck"].exit_code == 1

View file

@ -1,138 +0,0 @@
"""Tests for the `compat_result` fixture's tagged-union validation.
The conftest's `pytest_runtest_makereport` hook is exercised end-to-end by
the matrix-builder golden-file tests (which consume a results.json that
the harness would produce). Here we just test the input-validation
contract on `CompatResult.set()`.
"""
from __future__ import annotations
import pytest
from claude_code.conftest import CompatResult
def test_set_pass_is_accepted():
r = CompatResult()
r.set({"status": "pass"})
assert r.value == {"status": "pass"}
def test_set_fail_requires_error():
r = CompatResult()
with pytest.raises(ValueError, match="requires 'error'"):
r.set({"status": "fail"})
def test_set_fail_with_error_is_accepted():
r = CompatResult()
r.set({"status": "fail", "error": "boom"})
assert r.value == {"status": "fail", "error": "boom"}
def test_set_not_applicable_requires_reason():
r = CompatResult()
with pytest.raises(ValueError, match="requires 'reason'"):
r.set({"status": "not_applicable"})
def test_set_not_applicable_with_reason_is_accepted():
r = CompatResult()
r.set({"status": "not_applicable", "reason": "Bedrock has no /thinking"})
assert r.value == {"status": "not_applicable", "reason": "Bedrock has no /thinking"}
def test_set_not_tested_is_accepted():
r = CompatResult()
r.set({"status": "not_tested"})
assert r.value == {"status": "not_tested"}
def test_set_rejects_unknown_status():
r = CompatResult()
with pytest.raises(ValueError, match="status must be one of"):
r.set({"status": "maybe"})
def test_set_rejects_non_dict():
r = CompatResult()
with pytest.raises(TypeError):
r.set("pass") # type: ignore[arg-type]
def test_set_copies_input():
"""Mutating the dict after set() must not change the stored value."""
r = CompatResult()
payload = {"status": "fail", "error": "x"}
r.set(payload)
payload["error"] = "mutated"
assert r.value["error"] == "x"
# ---------------------------------------------------------------------------
# add() / collected()
#
# When a single test exercises three Claude tiers in parallel, each tier
# needs its own row in the results artifact so the matrix builder can
# apply its "all three must pass" aggregation. `add()` is the per-tier
# recorder; `collected()` is what the conftest hook reads.
# ---------------------------------------------------------------------------
def test_add_appends_each_call_to_values():
r = CompatResult()
r.add({"status": "pass"})
r.add({"status": "fail", "error": "bad"})
assert r.values == [
{"status": "pass"},
{"status": "fail", "error": "bad"},
]
def test_add_validates_like_set():
"""The add() and set() validators are the same; both must reject bad payloads."""
r = CompatResult()
with pytest.raises(ValueError, match="requires 'error'"):
r.add({"status": "fail"})
with pytest.raises(ValueError, match="requires 'reason'"):
r.add({"status": "not_applicable"})
with pytest.raises(ValueError, match="status must be one of"):
r.add({"status": "maybe"})
with pytest.raises(TypeError):
r.add("pass") # type: ignore[arg-type]
def test_add_copies_input():
"""Same defensive copy contract as set()."""
r = CompatResult()
payload = {"status": "fail", "error": "x"}
r.add(payload)
payload["error"] = "mutated"
assert r.values[0]["error"] == "x"
def test_collected_returns_values_when_added():
r = CompatResult()
r.add({"status": "pass"})
r.add({"status": "pass"})
assert r.collected() == [{"status": "pass"}, {"status": "pass"}]
def test_collected_returns_single_value_when_only_set_called():
"""Legacy single-result tests should still surface their one outcome."""
r = CompatResult()
r.set({"status": "pass"})
assert r.collected() == [{"status": "pass"}]
def test_collected_prefers_added_values_over_set_value():
"""If both are populated, the per-tier list wins — that's the multi-model shape."""
r = CompatResult()
r.set({"status": "pass"})
r.add({"status": "fail", "error": "tier-2 broke"})
assert r.collected() == [{"status": "fail", "error": "tier-2 broke"}]
def test_collected_returns_empty_when_nothing_reported():
assert CompatResult().collected() == []

View file

@ -1,221 +0,0 @@
"""Unit tests for the shared `run_passthrough_cell` helper.
These tests inject a fake `run_models` callable and an explicit `env`
mapping (both are first-class parameters, no monkeypatching), so they
exercise the helper's branching -- env-missing guard, base-URL
assembly, extra-env forwarding, per-model pass/fail -- without
spawning the real CLI.
The env-builder tests pin the provider-mode contract itself: the
CLAUDE_CODE_USE_* / CLAUDE_CODE_SKIP_*_AUTH flags and the passthrough
route each mode must target. Those values are the feature -- e.g.
dropping the `/v1` from the vertex base URL produces a request Google
404s on -- so a mutation to any of them must fail here before it burns
a live matrix run.
"""
from __future__ import annotations
from typing import Any, Dict, List, Mapping, Optional
import pytest
from claude_code._env import (
PRIMARY_API_KEY_ENV,
PRIMARY_BASE_URL_ENV,
)
from claude_code._passthrough import (
ANTHROPIC_PASSTHROUGH_BASE_PATH,
CLIENT_SIDE_AWS_REGION,
VERTEX_PLACEHOLDER_PROJECT,
VERTEX_PLACEHOLDER_REGION,
bedrock_extra_env,
foundry_extra_env,
run_passthrough_cell,
vertex_extra_env,
)
from claude_code.cli_driver import ClaudeCLIError, DriverResult
PROXY_ENV = {
PRIMARY_BASE_URL_ENV: "http://localhost:4000",
PRIMARY_API_KEY_ENV: "sk-test",
}
class _FakeResult:
def __init__(self) -> None:
self.rows: List[Dict[str, Any]] = []
self.single: Optional[Dict[str, Any]] = None
def set(self, payload: Mapping[str, Any]) -> None:
self.single = dict(payload)
def add(self, payload: Mapping[str, Any]) -> None:
self.rows.append(dict(payload))
def _fake_run_models(outcomes_by_model, captured: Dict[str, Any]):
def fake(*, models, prompt, base_url, api_key, extra_env=None, **_kwargs):
captured["models"] = list(models)
captured["prompt"] = prompt
captured["base_url"] = base_url
captured["api_key"] = api_key
captured["extra_env"] = dict(extra_env) if extra_env is not None else None
return {model: outcomes_by_model[model] for model in models}
return fake
def test_env_missing_guard_reports_fail_and_aborts():
fake_result = _FakeResult()
with pytest.raises(pytest.fail.Exception):
run_passthrough_cell(
compat_result=fake_result,
models=["claude-haiku-4-5"],
prompt="ping",
env={},
)
assert fake_result.single is not None
assert fake_result.single["status"] == "fail"
assert PRIMARY_BASE_URL_ENV in fake_result.single["error"]
assert PRIMARY_API_KEY_ENV in fake_result.single["error"]
def test_suite_wide_env_reaches_the_proxy():
"""Passthrough cells resolve via the same suite-wide env names as
every other e2e cell, so EKS wiring that only exports
LITELLM_PROXY_URL + LITELLM_MASTER_KEY reaches the ALB."""
fake_result = _FakeResult()
captured: Dict[str, Any] = {}
outcome = DriverResult(text="pong")
run_passthrough_cell(
compat_result=fake_result,
models=["claude-haiku-4-5"],
prompt="ping",
run_models=_fake_run_models({"claude-haiku-4-5": outcome}, captured),
env=PROXY_ENV,
)
assert captured["base_url"] == "http://localhost:4000"
assert captured["api_key"] == "sk-test"
assert fake_result.rows == [{"status": "pass"}]
def test_anthropic_base_path_appended_to_normalized_proxy_url():
fake_result = _FakeResult()
captured: Dict[str, Any] = {}
outcome = DriverResult(text="pong")
run_passthrough_cell(
compat_result=fake_result,
models=["claude-haiku-4-5"],
prompt="ping",
passthrough_base_path=ANTHROPIC_PASSTHROUGH_BASE_PATH,
run_models=_fake_run_models({"claude-haiku-4-5": outcome}, captured),
env={**PROXY_ENV, PRIMARY_BASE_URL_ENV: "http://localhost:4000/"},
)
assert captured["base_url"] == "http://localhost:4000/anthropic"
assert captured["extra_env"] is None
assert fake_result.rows == [{"status": "pass"}]
def test_extra_env_builder_receives_normalized_base_and_is_forwarded():
fake_result = _FakeResult()
captured: Dict[str, Any] = {}
outcome = DriverResult(text="pong")
seen_bases: List[str] = []
def build(proxy_base: str) -> Dict[str, str]:
seen_bases.append(proxy_base)
return {"SOME_FLAG": "1"}
run_passthrough_cell(
compat_result=fake_result,
models=["claude-haiku-4-5"],
prompt="ping",
build_extra_env=build,
run_models=_fake_run_models({"claude-haiku-4-5": outcome}, captured),
env={**PROXY_ENV, PRIMARY_BASE_URL_ENV: "http://localhost:4000/"},
)
assert seen_bases == ["http://localhost:4000"]
assert captured["extra_env"] == {"SOME_FLAG": "1"}
assert captured["base_url"] == "http://localhost:4000"
def test_per_model_failures_reported_individually():
fake_result = _FakeResult()
captured: Dict[str, Any] = {}
outcomes = {
"claude-haiku-4-5": DriverResult(text="pong"),
"claude-sonnet-4-5": ClaudeCLIError("claude CLI timed out after 120s"),
"claude-opus-4-7": DriverResult(text="", exit_code=1),
}
with pytest.raises(pytest.fail.Exception):
run_passthrough_cell(
compat_result=fake_result,
models=list(outcomes.keys()),
prompt="ping",
run_models=_fake_run_models(outcomes, captured),
env=PROXY_ENV,
)
statuses = [row["status"] for row in fake_result.rows]
assert statuses == ["pass", "fail", "fail"]
assert "timed out" in fake_result.rows[1]["error"]
assert "claude CLI failed" in fake_result.rows[2]["error"]
def test_empty_assistant_text_is_a_fail():
fake_result = _FakeResult()
captured: Dict[str, Any] = {}
outcomes = {"claude-haiku-4-5": DriverResult(text=" ")}
with pytest.raises(pytest.fail.Exception):
run_passthrough_cell(
compat_result=fake_result,
models=["claude-haiku-4-5"],
prompt="ping",
run_models=_fake_run_models(outcomes, captured),
env=PROXY_ENV,
)
assert fake_result.rows == [
{
"status": "fail",
"error": "[claude-haiku-4-5] claude returned empty assistant text",
}
]
def test_bedrock_extra_env_targets_proxy_bedrock_route():
env = bedrock_extra_env("http://localhost:4000")
assert env == {
"CLAUDE_CODE_USE_BEDROCK": "1",
"CLAUDE_CODE_SKIP_BEDROCK_AUTH": "1",
"ANTHROPIC_BEDROCK_BASE_URL": "http://localhost:4000/bedrock",
"AWS_REGION": CLIENT_SIDE_AWS_REGION,
}
def test_vertex_extra_env_keeps_the_api_version_in_the_base_url():
env = vertex_extra_env("http://localhost:4000")
assert env == {
"CLAUDE_CODE_USE_VERTEX": "1",
"CLAUDE_CODE_SKIP_VERTEX_AUTH": "1",
"ANTHROPIC_VERTEX_BASE_URL": "http://localhost:4000/vertex_ai/v1",
"ANTHROPIC_VERTEX_PROJECT_ID": VERTEX_PLACEHOLDER_PROJECT,
"CLOUD_ML_REGION": VERTEX_PLACEHOLDER_REGION,
}
def test_foundry_extra_env_targets_proxy_azure_route():
env = foundry_extra_env("http://localhost:4000")
assert env == {
"CLAUDE_CODE_USE_FOUNDRY": "1",
"CLAUDE_CODE_SKIP_FOUNDRY_AUTH": "1",
"ANTHROPIC_FOUNDRY_BASE_URL": "http://localhost:4000/azure",
}

View file

@ -1,329 +0,0 @@
"""Unit tests for the cross-process token-bucket rate limiter.
The tests cover three layers:
1. Provider inference from model alias the matrix-column mapping the
live tests rely on (`-bedrock-converse` vs `-bedrock-invoke` vs
`-azure` vs `-vertex` vs bare = anthropic).
2. Config parsing env-var precedence, fallback to default, malformed
input handling, burst override semantics. These run against
`os.environ`-shaped dicts so we don't have to monkeypatch globals.
3. Token-bucket behavior enforcing rate, accumulating burst, never
over-spending across a fake clock. Filesystem state is exercised
with a real `tmp_path` because the persistence is the whole point;
the only injected seam is `clock` (and `sleep`, so tests don't
actually wait on wall time).
The cross-process flock semantics are exercised indirectly: every
test creates a fresh `RateLimiter` rooted at `tmp_path`, so the same
file lock that protects production is exercised here too. We don't
fork to test multi-process behavior in this file because pytest
fixtures + xdist already do that for the integration suite.
"""
from __future__ import annotations
import json
import time
from pathlib import Path
from typing import List
import pytest
from claude_code.rate_limiter import (
ALL_PROVIDERS,
BURST_ENV,
DEFAULT_RATE,
PROVIDER_ANTHROPIC,
PROVIDER_AZURE,
PROVIDER_BEDROCK_CONVERSE,
PROVIDER_BEDROCK_INVOKE,
PROVIDER_VERTEX_AI,
ProviderConfig,
RateLimiter,
get_default_limiter,
infer_provider,
load_config,
reset_default_limiter,
use_limiter,
)
# ---------------------------------------------------------------------------
# Provider inference
# ---------------------------------------------------------------------------
@pytest.mark.parametrize(
"model, expected",
[
("claude-haiku-4-5", PROVIDER_ANTHROPIC),
("claude-sonnet-4-5", PROVIDER_ANTHROPIC),
("claude-opus-4-7", PROVIDER_ANTHROPIC),
("claude-haiku-4-5-azure", PROVIDER_AZURE),
("claude-sonnet-4-5-azure", PROVIDER_AZURE),
("claude-opus-4-7-vertex", PROVIDER_VERTEX_AI),
("claude-haiku-4-5-bedrock-converse", PROVIDER_BEDROCK_CONVERSE),
("claude-haiku-4-5-bedrock-invoke", PROVIDER_BEDROCK_INVOKE),
],
)
def test_infer_provider_maps_alias_suffix_to_column(model, expected):
assert infer_provider(model) == expected
def test_infer_provider_bedrock_converse_beats_bedrock_invoke_lookup_order():
"""Both bedrock suffixes contain `bedrock`; the more-specific suffix wins."""
assert infer_provider("claude-foo-bedrock-converse") == PROVIDER_BEDROCK_CONVERSE
assert infer_provider("claude-foo-bedrock-invoke") == PROVIDER_BEDROCK_INVOKE
def test_infer_provider_rejects_empty_string():
with pytest.raises(ValueError, match="non-empty"):
infer_provider("")
def test_infer_provider_is_case_insensitive():
"""Aliases in the proxy config sometimes drift between cases; we
should still route them to the right column."""
assert infer_provider("CLAUDE-OPUS-4-7-AZURE") == PROVIDER_AZURE
# ---------------------------------------------------------------------------
# Config parsing
# ---------------------------------------------------------------------------
def test_load_config_uses_default_rate_when_env_absent():
cfg = load_config(env={})
for provider in ALL_PROVIDERS:
assert cfg[provider].rate_per_sec == DEFAULT_RATE
assert cfg[provider].burst == DEFAULT_RATE
def test_load_config_reads_per_provider_rate():
cfg = load_config(
env={
"LITELLM_COMPAT_RATE_ANTHROPIC": "10",
"LITELLM_COMPAT_RATE_AZURE": "0.5",
}
)
assert cfg[PROVIDER_ANTHROPIC].rate_per_sec == 10.0
assert cfg[PROVIDER_AZURE].rate_per_sec == 0.5
assert cfg[PROVIDER_VERTEX_AI].rate_per_sec == DEFAULT_RATE
def test_load_config_zero_rate_disables_provider():
cfg = load_config(env={"LITELLM_COMPAT_RATE_BEDROCK_INVOKE": "0"})
assert cfg[PROVIDER_BEDROCK_INVOKE].enabled is False
def test_load_config_burst_override_applies_to_every_provider():
cfg = load_config(
env={
"LITELLM_COMPAT_RATE_ANTHROPIC": "5",
BURST_ENV: "20",
}
)
for provider in ALL_PROVIDERS:
assert cfg[provider].burst == 20.0
def test_load_config_falls_back_on_malformed_value():
cfg = load_config(env={"LITELLM_COMPAT_RATE_ANTHROPIC": "not-a-number"})
assert cfg[PROVIDER_ANTHROPIC].rate_per_sec == DEFAULT_RATE
def test_load_config_burst_floors_at_one_when_rate_is_low():
"""A 0.5/s rate with no burst override must still allow at least
one immediate request otherwise the very first call would block."""
cfg = load_config(env={"LITELLM_COMPAT_RATE_ANTHROPIC": "0.5"})
assert cfg[PROVIDER_ANTHROPIC].burst == 1.0
# ---------------------------------------------------------------------------
# Token bucket
# ---------------------------------------------------------------------------
@pytest.fixture
def fake_clock():
"""A controllable monotonic clock + sleep for the limiter under test.
Tests advance `clock.now` to simulate elapsed wall time. `sleep`
adds the requested duration to `clock.now` instead of actually
sleeping, so a "wait 200ms" code path runs in microseconds and
is deterministic.
"""
class Clock:
def __init__(self):
self.now = 1_000.0
self.sleeps: List[float] = []
def __call__(self):
return self.now
def sleep(self, seconds: float) -> None:
self.sleeps.append(seconds)
self.now += seconds
return Clock()
def _make_limiter(tmp_path: Path, fake_clock, *, rate=10.0, burst=None):
cfg = {
p: ProviderConfig(rate_per_sec=rate, burst=burst if burst is not None else rate)
for p in ALL_PROVIDERS
}
return RateLimiter(
config=cfg,
state_dir=tmp_path,
clock=fake_clock,
sleep=fake_clock.sleep,
)
def test_acquire_first_call_does_not_wait(tmp_path, fake_clock):
"""A freshly-initialized bucket starts full; the first acquire is free."""
limiter = _make_limiter(tmp_path, fake_clock, rate=10.0, burst=10.0)
waited = limiter.acquire(PROVIDER_ANTHROPIC)
assert waited == 0.0
assert fake_clock.sleeps == []
def test_acquire_disabled_provider_returns_immediately(tmp_path, fake_clock):
"""rate=0 ⇒ no throttling, even if every other provider is throttled."""
cfg = {p: ProviderConfig(rate_per_sec=0.0, burst=0.0) for p in ALL_PROVIDERS}
limiter = RateLimiter(
config=cfg, state_dir=tmp_path, clock=fake_clock, sleep=fake_clock.sleep
)
for _ in range(100):
assert limiter.acquire(PROVIDER_ANTHROPIC) == 0.0
assert fake_clock.sleeps == []
def test_acquire_burns_through_burst_then_throttles(tmp_path, fake_clock):
"""`burst` immediate requests succeed; the next one waits 1/rate seconds."""
limiter = _make_limiter(tmp_path, fake_clock, rate=2.0, burst=3.0)
for _ in range(3):
assert limiter.acquire(PROVIDER_ANTHROPIC) == 0.0
# Bucket is empty; next call must sleep ~0.5s to earn one token at 2/s.
waited = limiter.acquire(PROVIDER_ANTHROPIC)
assert waited == pytest.approx(0.5, abs=0.01)
def test_acquire_refills_with_elapsed_time(tmp_path, fake_clock):
"""Advancing the clock between calls credits tokens at the configured rate."""
limiter = _make_limiter(tmp_path, fake_clock, rate=4.0, burst=1.0)
assert limiter.acquire(PROVIDER_ANTHROPIC) == 0.0 # consumes the 1-token burst
fake_clock.now += 0.25 # 0.25s × 4/s = 1 token earned
assert limiter.acquire(PROVIDER_ANTHROPIC) == 0.0
def test_acquire_caps_refill_at_burst(tmp_path, fake_clock):
"""A long quiet period must not let the bucket grow past `burst`."""
limiter = _make_limiter(tmp_path, fake_clock, rate=10.0, burst=2.0)
fake_clock.now += 1_000 # would earn 10_000 tokens uncapped
# Only `burst` (=2) immediate calls should succeed before throttling.
assert limiter.acquire(PROVIDER_ANTHROPIC) == 0.0
assert limiter.acquire(PROVIDER_ANTHROPIC) == 0.0
waited = limiter.acquire(PROVIDER_ANTHROPIC)
assert waited > 0
def test_acquire_independent_buckets_per_provider(tmp_path, fake_clock):
"""Anthropic exhaustion must not throttle Azure (each column has its own bucket)."""
limiter = _make_limiter(tmp_path, fake_clock, rate=2.0, burst=1.0)
assert limiter.acquire(PROVIDER_ANTHROPIC) == 0.0
# Anthropic bucket is now empty; Azure is untouched.
assert limiter.acquire(PROVIDER_AZURE) == 0.0
def test_acquire_persists_state_across_limiter_instances(tmp_path):
"""A fresh RateLimiter must read the on-disk state, not start fresh.
This is the property that makes the limiter cross-process: an
xdist worker created mid-run sees the credit consumed by other
workers, instead of getting its own private bucket.
"""
cfg = {p: ProviderConfig(rate_per_sec=10.0, burst=2.0) for p in ALL_PROVIDERS}
state = {"now": 1_000.0, "sleeps": []}
def clock():
return state["now"]
def sleep(seconds):
state["sleeps"].append(seconds)
state["now"] += seconds
first = RateLimiter(config=cfg, state_dir=tmp_path, clock=clock, sleep=sleep)
first.acquire(PROVIDER_ANTHROPIC)
first.acquire(PROVIDER_ANTHROPIC)
# bucket is now empty
second = RateLimiter(config=cfg, state_dir=tmp_path, clock=clock, sleep=sleep)
waited = second.acquire(PROVIDER_ANTHROPIC)
assert waited > 0 # had to wait, didn't see a fresh full bucket
def test_acquire_recovers_from_corrupt_state_file(tmp_path, fake_clock):
"""A truncated/garbage state file must not crash the test session."""
state_file = tmp_path / f"{PROVIDER_ANTHROPIC}.json"
state_file.write_text("not-json {{")
limiter = _make_limiter(tmp_path, fake_clock, rate=5.0, burst=5.0)
assert limiter.acquire(PROVIDER_ANTHROPIC) == 0.0
def test_acquire_handles_clock_going_backward(tmp_path, fake_clock):
"""Across a host suspend/resume the monotonic clock can briefly
go backward; we must not interpret that as removing tokens."""
limiter = _make_limiter(tmp_path, fake_clock, rate=1.0, burst=2.0)
limiter.acquire(PROVIDER_ANTHROPIC)
fake_clock.now -= 10 # clock moved backward
# Bucket should still have ~1 token left from the burst, not -9.
assert limiter.acquire(PROVIDER_ANTHROPIC) == 0.0
# ---------------------------------------------------------------------------
# Process-default singleton
# ---------------------------------------------------------------------------
def test_use_limiter_swaps_default_for_block(tmp_path):
sentinel_cfg = {
p: ProviderConfig(rate_per_sec=0.0, burst=0.0) for p in ALL_PROVIDERS
}
sentinel = RateLimiter(config=sentinel_cfg, state_dir=tmp_path)
reset_default_limiter()
try:
with use_limiter(sentinel):
assert get_default_limiter() is sentinel
# After the context exits, the default goes back to whatever it
# was — in this test that's "rebuilt on next access" because we
# called reset_default_limiter() above.
assert get_default_limiter() is not sentinel
finally:
reset_default_limiter()
# ---------------------------------------------------------------------------
# Persistence shape
# ---------------------------------------------------------------------------
def test_state_file_is_json_after_acquire(tmp_path, fake_clock):
limiter = _make_limiter(tmp_path, fake_clock, rate=5.0, burst=5.0)
limiter.acquire(PROVIDER_ANTHROPIC)
state_file = tmp_path / f"{PROVIDER_ANTHROPIC}.json"
payload = json.loads(state_file.read_text())
assert "tokens" in payload
assert "last_refill" in payload
assert payload["tokens"] == pytest.approx(4.0)

View file

@ -1,162 +0,0 @@
"""Pin tests for the `Bash`-using compat cells.
Every cell that passes `--allowed-tools Bash` to the `claude` CLI is
giving a model-controlled response the ability to run host commands.
On the PR-gate CircleCI machine executor, those commands have access
to the Docker socket and can read `docker inspect compat-proxy` to
recover the provider credentials living inside the proxy container.
To narrow that surface, every Bash-using cell must:
1. Restrict the allow rule to the *exact* command `Bash(echo pong)` so
a compromised provider response cannot turn `Bash` into arbitrary
host execution by emitting a `tool_use` with a different command.
2. Pair it with `--permission-mode dontAsk` so anything not matching
an allow rule is auto-denied instead of prompting (which would
abort the CLI in headless mode, but auto-denial is the explicit
contract).
These restrictions are enforced by the `claude` CLI, not by the
model see https://code.claude.com/docs/en/permissions for the
permission-rule precedence (`deny` `ask` `allow`).
This test scans every cell under the three Bash-using feature
directories (`tool_use`, `tool_use_streaming`, `thinking_with_tool_use`)
and pins both requirements so a future test refactor cannot silently
revert any cell to the broad `Bash` allow that was originally
flagged by Veria.
"""
from __future__ import annotations
from pathlib import Path
from typing import Iterable
import pytest
CLAUDE_CODE_DIR = Path(__file__).resolve().parents[1]
# Feature directories whose cells drive the `Bash` built-in tool. Add
# new entries here when a new Bash-using feature is added; the test
# fails loudly for any unhandled directory so we never miss one by
# silent omission.
BASH_FEATURE_DIRS = (
"tool_use",
"tool_use_streaming",
"thinking_with_tool_use",
)
def _bash_cells() -> Iterable[Path]:
for feature in BASH_FEATURE_DIRS:
feature_dir = CLAUDE_CODE_DIR / feature
assert feature_dir.is_dir(), (
f"{feature_dir} is missing — BASH_FEATURE_DIRS is out of sync "
f"with the layout under tests/e2e/claude_code/."
)
for path in sorted(feature_dir.glob("test_*.py")):
yield path
def _has_bare_bash_token(text: str) -> bool:
"""Return True if `text` contains a `"Bash"` token outside the
`"Bash(echo pong)"` allow rule.
Extracted as a pure helper so the negative path can be unit-tested
directly. Without it, the previous structure of this assertion was
`'"Bash"' not in text or '"Bash(echo pong)"' in text`, which
short-circuits to True any time the allow rule is present and lets
a stray bare `"Bash"` slip through the security pin undetected.
"""
return '"Bash"' in text.replace('"Bash(echo pong)"', "")
@pytest.mark.parametrize(
"cell", list(_bash_cells()), ids=lambda p: str(p.relative_to(CLAUDE_CODE_DIR))
)
def test_bash_allow_rule_is_pinned_to_exact_echo_pong(cell: Path) -> None:
"""The cell must pass `Bash(echo pong)` as the allow rule, not the
unrestricted `Bash` value that was originally flagged."""
text = cell.read_text()
assert '"Bash(echo pong)"' in text, (
f"{cell.relative_to(CLAUDE_CODE_DIR)} must restrict `--allowed-tools` to "
f'`Bash(echo pong)` (exact-match pattern). Unrestricted `"Bash"` '
f"grants arbitrary host command execution to model-controlled "
f"tool_use blocks, which can read `docker inspect compat-proxy` "
f"to exfiltrate provider credentials from the proxy container."
)
# The only place `"Bash"` (the bare token, surrounded by quotes
# exactly as it would appear in `--allowed-tools` lists) is allowed
# to appear is *inside* the exact-match `"Bash(echo pong)"` rule.
# `_has_bare_bash_token` keeps that scan independent of the first
# assertion — otherwise `'"Bash"' not in text or '"Bash(echo pong)"'
# in text` short-circuits to True and lets a stray bare `"Bash"`
# slip through silently.
assert not _has_bare_bash_token(text), (
f"{cell.relative_to(CLAUDE_CODE_DIR)} still references the unrestricted "
f'`"Bash"` value outside the `"Bash(echo pong)"` allow rule — '
f"sweep it out before merging."
)
def test_has_bare_bash_token_flags_unrestricted_value():
"""A file that allows the bare `"Bash"` token alongside the
exact-match rule must be flagged. Without this guard the security
pin reverts to the dead-code `or` it had originally, which let
arbitrary host commands through under the noise of a passing test.
"""
text = '--allowed-tools "Bash" "Bash(echo pong)"'
assert _has_bare_bash_token(text)
def test_has_bare_bash_token_accepts_only_exact_match():
"""The standard pattern — only the exact-match allow rule, no bare
`"Bash"` must be accepted. This is the shape every Bash-using
cell in the suite is required to take.
"""
text = '--allowed-tools "Bash(echo pong)" --permission-mode "dontAsk"'
assert not _has_bare_bash_token(text)
def test_has_bare_bash_token_ignores_unrelated_substrings():
"""`Bash(echo pong)` is the only allowed shape; substrings like
`BashTool` or `Bashing` are unrelated identifiers and must not be
confused with the bare `"Bash"` token (i.e. the exact quoted
string `"Bash"`)."""
text = "BashTool helper used by the bashing harness"
assert not _has_bare_bash_token(text)
@pytest.mark.parametrize(
"cell", list(_bash_cells()), ids=lambda p: str(p.relative_to(CLAUDE_CODE_DIR))
)
def test_bash_cell_uses_dontask_permission_mode(cell: Path) -> None:
"""The cell must pair the allow rule with `--permission-mode dontAsk`
so tool calls that don't match the allow rule are auto-denied (as
opposed to defaulting to "ask", which in headless mode would
succeed without ever surfacing the security issue)."""
text = cell.read_text()
assert '"--permission-mode"' in text and '"dontAsk"' in text, (
f"{cell.relative_to(CLAUDE_CODE_DIR)} must pass `--permission-mode dontAsk` "
f"alongside the `Bash(echo pong)` allow rule. Without dontAsk, "
f"commands outside the allow rule fall back to the default ask-"
f"mode behavior, which in `--print` (headless) mode is non-"
f"interactive — defeating the explicit-allow contract."
)
def test_claude_code_dir_anchor_is_layout_independent() -> None:
"""CLAUDE_CODE_DIR must resolve to the `claude_code/` directory that
contains this test file, regardless of how deep the repository is
mounted. The previous anchor `Path(__file__).resolve().parents[4]`
baked in the host layout (repo root sits four levels up) and broke
when the suite runs inside the stage container, where tests/e2e/ is
mounted at /app/e2e/ so `parents[4]` resolves to filesystem root and
the BASH_FEATURE_DIRS assertion looks for `/tests/e2e/claude_code/
tool_use`. Anchoring at `parents[1]` (the sibling of this file's
parent) is the same directory in both layouts.
"""
assert CLAUDE_CODE_DIR.name == "claude_code"
assert CLAUDE_CODE_DIR.is_dir()
assert (CLAUDE_CODE_DIR / "_pr_gate_unit_tests" / Path(__file__).name).is_file()

View file

@ -1,166 +0,0 @@
"""Regression tests for the compat-model registration loader.
The compat cells hardcode virtual model names like ``claude-sonnet-4-5``
and expect them to be registered on the proxy before the cell runs. The
session fixture in ``conftest.py`` reads ``test_config.yaml`` and POSTs
those deployments via ``/model/new``. These tests pin the invariants
that make that safe:
- The yaml declares an entry for every virtual name a cell references -
otherwise a cell probes a name the fixture never registered, and the
cell hits an ``Invalid model name`` 400 that is much harder to trace.
- The ``vertex_ai_*`` yaml keys get normalized to the ``vertex_*``
pydantic-body names before ``LiteLLMParamsBody(**)`` sees them, so
the vertex project/location aren't silently dropped by pydantic's
``extra="ignore"`` default.
"""
from __future__ import annotations
import re
from pathlib import Path
import pytest
from claude_code._compat_models import (
all_expected_model_names,
load_all_deployments,
)
CLAUDE_CODE_DIR = Path(__file__).resolve().parents[1]
def _cell_declared_model_names() -> frozenset[str]:
"""Every ``"claude-*"`` model name a compat cell hardcodes in a
``*_MODELS`` list. Uses a simple regex rather than importing every
cell because the cells depend on the harness which depends on env
the unit-test run does not have."""
pattern = re.compile(r'"(claude-[a-zA-Z0-9._-]+)"')
found: set[str] = set()
for path in CLAUDE_CODE_DIR.glob("*/test_*.py"):
if path.parent.name.startswith("_"):
continue
for match in pattern.finditer(path.read_text()):
name = match.group(1)
# Skip upstream model references (they carry a version
# suffix or the ``anthropic/`` provider prefix - we only
# want proxy-side virtual names here).
if "/" in name or "@" in name:
continue
found.add(name)
return frozenset(found)
def test_yaml_covers_every_cell_declared_model_name() -> None:
"""Every ``"claude-..."`` string a cell probes must have a
corresponding ``model_list`` entry in ``test_config.yaml``. A new
cell that adds a probe for a name the yaml doesn't know fails this
test - the alternative is a 400 at runtime that is much harder to
diagnose."""
yaml_names = all_expected_model_names()
cell_names = _cell_declared_model_names()
missing = cell_names - yaml_names
assert not missing, (
f"compat cells reference model names not declared in "
f"test_config.yaml: {sorted(missing)}. Add a matching "
f"model_list entry so the session fixture can register them."
)
def test_yaml_has_no_unused_declarations() -> None:
"""Every declaration in ``test_config.yaml`` is referenced by at
least one cell. A yaml entry no test exercises is dead
configuration and drift-prone; delete it or add the cell."""
yaml_names = all_expected_model_names()
cell_names = _cell_declared_model_names()
unused = yaml_names - cell_names
assert not unused, (
f"test_config.yaml declares model names no cell references: "
f"{sorted(unused)}. Delete them or add the cell."
)
def test_load_returns_fifteen_deployments() -> None:
"""The compat matrix is 3 tiers x 5 provider surfaces = 15. Pin the
count so a future edit to the yaml can't silently drop a tier."""
assert len(load_all_deployments()) == 15
def test_deployments_are_hashable_and_frozen() -> None:
"""``CompatDeployment`` is frozen so tests cannot accidentally
mutate the shared list mid-session."""
d = load_all_deployments()[0]
with pytest.raises((AttributeError, TypeError)):
d.model_name = "mutated" # type: ignore[misc]
def test_vertex_yaml_keys_populate_pydantic_body() -> None:
"""The yaml spells vertex fields ``vertex_ai_project`` /
``vertex_ai_location`` but ``LiteLLMParamsBody`` names them
``vertex_project`` / ``vertex_location``. Without the alias
normalization the pydantic body silently drops the yaml keys, and
the deployment gets registered with no vertex project - a real
incident the drift regressed twice historically."""
all_deployments = load_all_deployments()
vertex = [
d for d in all_deployments if d.model_name.endswith("-vertex")
]
assert vertex, "no vertex deployments found in yaml"
for d in vertex:
assert d.litellm_params.vertex_project, (
f"{d.model_name} lost its vertex_project after normalization"
)
assert d.litellm_params.vertex_location, (
f"{d.model_name} lost its vertex_location after normalization"
)
def test_vertex_deployments_keep_use_in_pass_through() -> None:
"""Vertex passthrough cells need the deployment registered with
``use_in_pass_through: true`` so the proxy wires project/location
credentials into the /vertex_ai passthrough router. ``LiteLLMParamsBody``
defaults to ``extra="ignore"``, so a missing field on the body silently
strips the yaml flag and every vertex passthrough cell fails at runtime
with "No credentials found on proxy for project_name=..."."""
vertex = [
d
for d in load_all_deployments()
if d.model_name.endswith("-vertex")
]
assert vertex, "no vertex deployments found in yaml"
for d in vertex:
assert d.litellm_params.use_in_pass_through is True, (
f"{d.model_name} lost use_in_pass_through after load; "
f"serialized body would be "
f"{d.litellm_params.model_dump(exclude_none=True)}"
)
def test_yaml_litellm_params_are_all_known_body_fields() -> None:
"""Every key under ``litellm_params`` in ``test_config.yaml`` must map
to a ``LiteLLMParamsBody`` field (after the vertex alias rewrite).
Without this pin, a new yaml flag can land in the fixture config and
be silently dropped by pydantic before ``/model/new`` ever sees it."""
import yaml
from models import LiteLLMParamsBody
from claude_code._compat_models import (
CONFIG_PATH,
_YAML_TO_PYDANTIC_ALIASES,
)
known = frozenset(LiteLLMParamsBody.model_fields)
doc = yaml.safe_load(CONFIG_PATH.read_text())
model_list = doc.get("model_list") or []
unknown = tuple(
(entry["model_name"], key)
for entry in model_list
for key in entry["litellm_params"]
if _YAML_TO_PYDANTIC_ALIASES.get(key, key) not in known
)
assert not unknown, (
f"test_config.yaml litellm_params keys not on LiteLLMParamsBody "
f"(will be silently dropped at register time): {unknown}"
)

View file

@ -1,162 +0,0 @@
"""Regression tests for ``claude_code/_env.py``.
Pin the resolution rules so a future edit cannot silently reintroduce
a private spelling that makes every ``claude_code`` cell fail with
"not configured" even when the surrounding e2e suite has a live proxy
configured under the suite-wide ``LITELLM_PROXY_URL`` /
``LITELLM_MASTER_KEY`` names.
"""
from __future__ import annotations
import pytest
from claude_code._env import (
PRIMARY_API_KEY_ENV,
PRIMARY_BASE_URL_ENV,
ProxyConfig,
require_proxy,
resolve_proxy_from,
)
class _CompatResultStub:
"""Minimal stand-in for the compat_result fixture used by cells."""
def __init__(self) -> None:
self.calls: list[dict[str, str]] = []
def set(self, payload: dict[str, str]) -> None:
self.calls.append(payload)
def test_primary_env_names_match_suite_wide_config() -> None:
"""The names claude_code reads must exactly match the ones
``e2e_config.py`` reads for the rest of the suite. Anything else
silently reintroduces the drift this refactor cleaned up."""
assert PRIMARY_BASE_URL_ENV == "LITELLM_PROXY_URL"
assert PRIMARY_API_KEY_ENV == "LITELLM_MASTER_KEY"
def test_returns_none_when_no_env_is_set() -> None:
assert resolve_proxy_from({}) is None
def test_returns_none_when_only_url_is_set() -> None:
assert (
resolve_proxy_from({PRIMARY_BASE_URL_ENV: "http://localhost:4000"}) is None
)
def test_returns_none_when_only_key_is_set() -> None:
assert resolve_proxy_from({PRIMARY_API_KEY_ENV: "sk-1234"}) is None
def test_primary_pair_resolves() -> None:
cfg = resolve_proxy_from(
{
PRIMARY_BASE_URL_ENV: "http://localhost:4000",
PRIMARY_API_KEY_ENV: "sk-1234",
}
)
assert cfg == ProxyConfig("http://localhost:4000", "sk-1234")
def test_legacy_pair_is_ignored() -> None:
"""``LITELLM_PROXY_BASE_URL`` / ``LITELLM_PROXY_API_KEY`` are not
accepted. A runner that only exports those must fail closed rather
than silently use a private spelling that the rest of the suite
does not know about."""
assert (
resolve_proxy_from(
{
"LITELLM_PROXY_BASE_URL": "http://legacy:4000",
"LITELLM_PROXY_API_KEY": "sk-legacy",
}
)
is None
)
def test_empty_string_env_is_treated_as_unset() -> None:
"""``os.environ.get`` on an exported-but-empty var returns "" which
is falsy. The resolver must treat that as unset so a shell that
accidentally exports ``LITELLM_PROXY_URL=`` doesn't turn into a
"" base_url that hits the wrong endpoint."""
assert (
resolve_proxy_from(
{PRIMARY_BASE_URL_ENV: "", PRIMARY_API_KEY_ENV: "sk-1234"}
)
is None
)
def test_require_proxy_fails_with_helpful_message_when_env_empty() -> None:
"""The error the user sees must name the suite-wide env vars so
they know exactly what to export."""
compat = _CompatResultStub()
with pytest.raises(pytest.fail.Exception) as excinfo:
require_proxy(compat, env={})
assert PRIMARY_BASE_URL_ENV in str(excinfo.value)
assert PRIMARY_API_KEY_ENV in str(excinfo.value)
assert compat.calls and compat.calls[0]["status"] == "fail"
assert PRIMARY_BASE_URL_ENV in compat.calls[0]["error"]
assert PRIMARY_API_KEY_ENV in compat.calls[0]["error"]
assert "LITELLM_PROXY_BASE_URL" not in compat.calls[0]["error"]
def test_require_proxy_returns_config_when_primary_env_supplied() -> None:
cfg = require_proxy(
_CompatResultStub(),
env={
PRIMARY_BASE_URL_ENV: "http://localhost:4000",
PRIMARY_API_KEY_ENV: "sk-1234",
},
)
assert cfg == ProxyConfig("http://localhost:4000", "sk-1234")
class TestControlGatewayFollowsResolvedProxy:
"""The session fixture that registers the compat deployments must talk
to the *same* proxy the cells do.
Building its Gateway off ``e2e_config``'s own env read instead of the
resolved ``ProxyConfig`` would send ``/model/new`` to whatever
``e2e_config`` defaults to when the process env is empty, while the
cells drive a different host so registration silently lands
somewhere else and every cell 400s with "Invalid model name".
"""
RESOLVED = ProxyConfig("http://eks-alb.internal:4000", "sk-eks")
def _gateway(self):
from claude_code.conftest import _build_control_gateway
return _build_control_gateway(self.RESOLVED)
def test_management_calls_go_to_the_resolved_host_and_key(self) -> None:
control = self._gateway().transport.control
assert control.base_url == self.RESOLVED.base_url
assert control.master_key == self.RESOLVED.api_key
def test_both_planes_share_the_one_address_the_cells_use(self) -> None:
"""The deployment is fronted by a single address that routes
management and LLM paths itself, so a resolved proxy pins both."""
transport = self._gateway().transport
assert transport.data.base_url == self.RESOLVED.base_url
assert transport.data.master_key == self.RESOLVED.api_key
assert transport.control.base_url == transport.data.base_url
def test_require_proxy_leaves_compat_result_untouched_on_success() -> None:
"""A successful resolution must NOT append a spurious fail entry.
Would have silently poisoned every compat cell's result rows."""
compat = _CompatResultStub()
require_proxy(
compat,
env={
PRIMARY_BASE_URL_ENV: "http://localhost:4000",
PRIMARY_API_KEY_ENV: "sk-1234",
},
)
assert compat.calls == []

View file

@ -1,164 +0,0 @@
"""Unit tests for the Claude Code PR-Gate Version Resolver.
The resolver picks the newest `@anthropic-ai/claude-code` version whose
publish timestamp is at least 3 days old. The 3-day window is a security
review buffer: a malicious or broken Claude Code release that slipped
through the npm publish process gets at least 72 hours to be detected
before it can land in the LiteLLM PR gate.
The unit tests inject npm metadata directly (no network) and a fixed
`as_of` clock (no real time), so they run anywhere and never flake on
the wall clock or registry availability.
"""
from __future__ import annotations
from datetime import datetime, timedelta, timezone
import pytest
from claude_code.pr_gate_version_resolver import (
NoEligibleVersionError,
resolve_pr_gate_version,
)
def _t(iso: str) -> str:
"""Helper for readable ISO-8601 publish timestamps in fixtures."""
return iso
# A clock fixed at a moment well after every fixture publish time below.
NOW = datetime(2026, 4, 25, 12, 0, 0, tzinfo=timezone.utc)
def _metadata_with_times(times: dict) -> dict:
"""Shape an npm `packument`-like dict with the `time` field populated.
The npm registry response includes `time.created` / `time.modified`
keys alongside per-version timestamps; the resolver must skip those.
"""
return {
"name": "@anthropic-ai/claude-code",
"time": {
"created": _t("2024-01-01T00:00:00.000Z"),
"modified": _t("2026-04-25T00:00:00.000Z"),
**times,
},
}
def test_picks_newest_version_at_least_three_days_old():
metadata = _metadata_with_times(
{
"2.1.118": _t("2026-04-15T10:00:00.000Z"),
"2.1.119": _t("2026-04-21T10:00:00.000Z"), # 4d 2h old
"2.1.120": _t("2026-04-23T10:00:00.000Z"), # 2d 2h old — too new
"2.1.121": _t("2026-04-25T11:00:00.000Z"), # 1h old — too new
}
)
assert resolve_pr_gate_version(metadata=metadata, as_of=NOW) == "2.1.119"
def test_skips_created_and_modified_meta_keys():
"""`time` contains `created` / `modified` non-version entries — must be ignored."""
metadata = {
"name": "@anthropic-ai/claude-code",
"time": {
"created": _t("2024-01-01T00:00:00.000Z"),
"modified": _t("2026-04-25T00:00:00.000Z"),
"2.0.0": _t("2026-04-10T00:00:00.000Z"),
},
}
assert resolve_pr_gate_version(metadata=metadata, as_of=NOW) == "2.0.0"
def test_min_age_boundary_is_inclusive():
"""A version published exactly 3 days ago is eligible (>= cutoff)."""
three_days_ago = NOW - timedelta(days=3)
metadata = _metadata_with_times(
{
"2.1.0": three_days_ago.isoformat().replace("+00:00", "Z"),
}
)
assert resolve_pr_gate_version(metadata=metadata, as_of=NOW) == "2.1.0"
def test_raises_when_every_version_is_too_new():
metadata = _metadata_with_times(
{
"2.1.121": _t("2026-04-25T08:00:00.000Z"), # 4h old
"2.1.120": _t("2026-04-24T10:00:00.000Z"), # ~26h old
}
)
with pytest.raises(NoEligibleVersionError):
resolve_pr_gate_version(metadata=metadata, as_of=NOW)
def test_raises_when_metadata_has_no_versions():
metadata = {"name": "@anthropic-ai/claude-code", "time": {}}
with pytest.raises(NoEligibleVersionError):
resolve_pr_gate_version(metadata=metadata, as_of=NOW)
def test_picks_latest_publish_time_not_largest_semver():
"""If a patch is published to an old major after a newer release,
"newest" is by publish time, not semver string ordering."""
metadata = _metadata_with_times(
{
"1.9.99": _t("2026-04-22T10:00:00.000Z"), # patched recently — wins
"2.0.0": _t("2026-03-01T10:00:00.000Z"), # older publish
}
)
assert resolve_pr_gate_version(metadata=metadata, as_of=NOW) == "1.9.99"
def test_uses_custom_min_age():
metadata = _metadata_with_times(
{
"1.0.0": _t("2026-04-23T10:00:00.000Z"), # 2d 2h old
"0.9.0": _t("2026-04-10T10:00:00.000Z"), # 15d old
}
)
# min_age = 5 days disqualifies 1.0.0
out = resolve_pr_gate_version(
metadata=metadata, as_of=NOW, min_age=timedelta(days=5)
)
assert out == "0.9.0"
def test_excludes_prerelease_versions():
"""Pre-release tags (1.0.0-alpha.1, 2.0.0-rc.1, etc.) must never win,
even if their publish timestamp is the newest eligible one."""
metadata = _metadata_with_times(
{
"2.1.119": _t("2026-04-21T10:00:00.000Z"), # stable, 4d old
"2.2.0-alpha.1": _t("2026-04-22T10:00:00.000Z"), # newer publish
"2.2.0-rc.1": _t("2026-04-22T11:00:00.000Z"), # newest publish
"3.0.0-beta": _t("2026-04-22T12:00:00.000Z"), # newest publish
}
)
assert resolve_pr_gate_version(metadata=metadata, as_of=NOW) == "2.1.119"
def test_raises_when_only_prereleases_are_eligible():
metadata = _metadata_with_times(
{
"2.2.0-alpha.1": _t("2026-04-22T10:00:00.000Z"),
"2.2.0-rc.1": _t("2026-04-22T11:00:00.000Z"),
}
)
with pytest.raises(NoEligibleVersionError):
resolve_pr_gate_version(metadata=metadata, as_of=NOW)
def test_resolver_uses_fetcher_when_metadata_not_provided():
captured = {}
def fake_fetch(package_name: str) -> dict:
captured["package"] = package_name
return _metadata_with_times({"3.0.0": _t("2026-04-10T10:00:00.000Z")})
out = resolve_pr_gate_version(as_of=NOW, fetcher=fake_fetch)
assert out == "3.0.0"
assert captured["package"] == "@anthropic-ai/claude-code"

View file

@ -15,13 +15,14 @@ shared fixtures build on it.
import functools
import sys
from collections.abc import Generator, Iterator
from pathlib import Path
from typing import Iterator
import pytest
import requests
from e2e_config import CONTROL_PLANE_BASE_URL, PROXY_BASE_URL
from e2e_result_reporter import covers_from_item, format_e2e_result_line, result_from_pytest
from lifecycle import GatewayProvider, ResourceManager
@ -85,6 +86,30 @@ def pytest_runtest_call(item: pytest.Item) -> None:
item.session.stash[_E2E_TEST_RAN] = True
@pytest.hookimpl(wrapper=True, tryfirst=True)
def pytest_runtest_makereport(
item: pytest.Item, call: pytest.CallInfo[object]
) -> Generator[None, pytest.TestReport, pytest.TestReport]:
"""Emit one structured E2E_RESULT line per finished test for Loki/Grafana.
Status-history panels should aggregate by package (and optional covers), not
scrape pytest progress basenames. See e2e_result_reporter.py.
"""
report = yield
result = result_from_pytest(
nodeid=str(report.nodeid),
when=str(report.when),
failed=bool(report.failed),
skipped=bool(report.skipped),
passed=bool(report.passed),
duration_seconds=float(report.duration),
covers=covers_from_item(item),
)
if result is not None:
print(format_e2e_result_line(result), flush=True)
return report
def pytest_sessionfinish(session: pytest.Session, exitstatus: int) -> None:
"""Once the whole e2e session is done (all suites), truncate the spend logs so
the DB doesn't accumulate test rows. Sessions where no e2e test body ran leave

View file

@ -53,6 +53,11 @@ in `MODULE_ORDER`, in that order. Loki uses log-safe `module=` labels from
`LOKI_MODULE_LABELS` (`core_llms`, `management_ui`, etc.) so existing JSON and
Prometheus consumers keep their human-readable module names unchanged.
Live pass/fail is separate: each finished pytest node prints an `E2E_RESULT`
logfmt line (see `tests/e2e/e2e_result_reporter.py` and
`tests/e2e/grafana/status_history_panels.md`). Coverage answers "is there a
test for this cell?"; `E2E_RESULT` answers "did that run pass?"
The headline is overall coverage. The collector also lists markers that point at ids
not in the registry, so a typo or an unenumerated behavior surfaces instead of being
silently dropped.

View file

@ -0,0 +1,144 @@
"""Structured e2e result lines for Loki / Grafana status history.
Pytest progress lines are a bad dashboard source: they only expose file basenames,
break under quiet modes, and force status-history rows to explode with suite growth.
Each finished test emits one logfmt line:
E2E_RESULT package=logging file=test_langfuse_e2e.py outcome=failed
duration_ms=1234 node_id=logging/test_langfuse_e2e.py::TestX::test_y
covers=logging.langfuse.team.success
Grafana package status-history queries max(fail) by package over E2E_RESULT lines.
Drill-down uses node_id / covers in Explore, not status-history cardinality.
"""
from __future__ import annotations
from collections.abc import Iterable, Sequence
from dataclasses import dataclass
from pathlib import Path
from typing import Literal, Protocol, runtime_checkable
Outcome = Literal["passed", "failed", "error", "skipped"]
@dataclass(frozen=True, slots=True)
class E2EResult:
package: str
file: str
outcome: Outcome
duration_ms: int
node_id: str
covers: tuple[str, ...]
@runtime_checkable
class _MarkerArgs(Protocol):
args: Sequence[object]
@runtime_checkable
class _ItemWithCovers(Protocol):
def iter_markers(self, name: str) -> Iterable[object]: ...
def package_from_nodeid(nodeid: str) -> str:
"""Top-level suite package under tests/e2e/, or 'root' for top-level files.
Pytest nodeids are relative to the invocation cwd. Repo-root runs look like
`tests/e2e/logging/...`; suite-cwd runs look like `logging/...`. Strip the
`tests/e2e` prefix so package is the suite dir either way.
"""
path_part = nodeid.split("::", 1)[0].replace("\\", "/")
parts = tuple(p for p in path_part.split("/") if p and p != ".")
if len(parts) >= 3 and parts[0] == "tests" and parts[1] == "e2e":
parts = parts[2:]
if len(parts) <= 1:
return "root"
return parts[0]
def file_from_nodeid(nodeid: str) -> str:
path_part = nodeid.split("::", 1)[0].replace("\\", "/")
return Path(path_part).name
def covers_from_item(item: object) -> tuple[str, ...]:
"""Read @pytest.mark.covers cell ids from a pytest Item."""
if not isinstance(item, _ItemWithCovers):
return ()
return tuple(
dict.fromkeys(
arg
for marker in item.iter_markers(name="covers")
if isinstance(marker, _MarkerArgs)
for arg in marker.args
if isinstance(arg, str) and arg
)
)
def outcome_from_report(when: str, failed: bool, skipped: bool, passed: bool) -> Outcome | None:
"""Map pytest TestReport fields to a terminal outcome. None if not final."""
if when == "setup" and skipped:
return "skipped"
if when == "setup" and failed:
return "error"
if when != "call":
return None
if skipped:
return "skipped"
if failed:
return "failed"
if passed:
return "passed"
return "failed"
def _logfmt_escape(value: str) -> str:
if value == "":
return '""'
needs_quote = any(ch.isspace() or ch in "\"=\\" for ch in value)
if not needs_quote:
return value
escaped = value.replace("\\", "\\\\").replace('"', '\\"')
return f'"{escaped}"'
def format_e2e_result_line(result: E2EResult) -> str:
covers = ",".join(result.covers)
fields = (
("package", result.package),
("file", result.file),
("outcome", result.outcome),
("duration_ms", str(result.duration_ms)),
("node_id", result.node_id),
("covers", covers),
)
body = " ".join(f"{key}={_logfmt_escape(value)}" for key, value in fields)
return f"E2E_RESULT {body}"
def result_from_pytest(
*,
nodeid: str,
when: str,
failed: bool,
skipped: bool,
passed: bool,
duration_seconds: float,
covers: tuple[str, ...] = (),
) -> E2EResult | None:
outcome = outcome_from_report(when=when, failed=failed, skipped=skipped, passed=passed)
if outcome is None:
return None
duration_ms = max(0, int(round(duration_seconds * 1000)))
return E2EResult(
package=package_from_nodeid(nodeid),
file=file_from_nodeid(nodeid),
outcome=outcome,
duration_ms=duration_ms,
node_id=nodeid,
covers=covers,
)

View file

@ -0,0 +1,66 @@
# Grafana: package status history for e2e
Dashboard: [LiteLLM E2E](https://berriai.grafana.net/d/mup2cfn/litellm-e2e) (`mup2cfn`).
The old **test suite status history** panel scraped pytest progress lines and
grouped by **file basename** (`test_foo.py`). That does not scale: multi-class
files collapse to one bit, and full `node_id` cardinality melts status-history.
## Emitter
After each test finishes, `tests/e2e/conftest.py` prints one logfmt line:
```
E2E_RESULT package=logging file=test_langfuse_e2e.py outcome=failed duration_ms=1500 node_id="logging/..." covers=cell.id
```
## Panel: package status history (replace panel 11)
**Type:** Status history
**Interval:** 15m (or 1h for multi-day ranges)
**Description:** Per top-level package under `tests/e2e/`: red if any test failed or errored in the bucket.
```logql
max by (package) (
max_over_time(
{service_name="litellm-e2e"}
|= "E2E_RESULT"
| logfmt
| outcome != ""
| label_format result=`{{ if or (eq .outcome "failed") (eq .outcome "error") }}1{{ else }}0{{ end }}`
| unwrap result
[$__interval]
)
)
```
Value mappings: `0` → Pass (green), `1` → Fail (red).
If `service_name` is missing on older scrapes, use:
```logql
{cluster="berrie-litellm-stage", pod=~"litellm-e2e-.+"}
```
instead of `{service_name="litellm-e2e"}`.
## Panel: failed tests (logs drill-down)
```logql
{service_name="litellm-e2e"} |= "E2E_RESULT" | logfmt | outcome=~"failed|error"
```
Show fields: `package`, `file`, `node_id`, `covers`, `duration_ms`.
## Panel (optional): filter by package variable
Dashboard variable `package` (custom or from label_values on E2E_RESULT):
```logql
{service_name="litellm-e2e"} |= "E2E_RESULT" | logfmt | package=`$package` | outcome=~"failed|error"
```
## Do not
- Put full `node_id` as the status-history series key (cardinality).
- Rely on `::S+ PASSED` progress regex as the primary signal once E2E_RESULT is live.

View file

@ -1,273 +0,0 @@
"""Unit coverage for the Gateway model-management surface (create_model /
delete_model) and the bounded spend read-back (spend_logs_window).
The batches conftest and several llm_translation tests register deployments at
runtime through gateway.create_model; when that method went missing, every batch
test errored at fixture setup (AttributeError) before a single request reached
the proxy. This pins the surface with a typed fake Transport so a rename or
signature drift fails here instead of in a live stage run.
spend_logs_window exists because the unpaginated /spend/logs whole-table read
grew past the e2e runner's memory limit on stage and OOMKilled every run; these
tests pin its /spend/logs/v2 pagination and that SpendLogsParams can no longer
express the unfiltered read.
"""
from dataclasses import dataclass, field
from datetime import datetime, timezone
import pytest
from pydantic import BaseModel, ValidationError
from batches.batch_client import BatchClient
from e2e_gateway import Gateway
from e2e_http import (
AuthHeaders,
FileUploadForm,
ProbeResult,
Result,
StreamingResponse,
Success,
UnknownApiError,
)
from models import (
LiteLLMParamsBody,
ModelDeleteBody,
ModelNewBody,
ModelNewResponse,
ModelsListResponse,
SpendLogsPage,
SpendLogsPageParams,
SpendLogsParams,
)
@dataclass
class _RecordingTransport:
"""Typed fake fulfilling the Transport protocol; records every post and
answers with a canned success so the test asserts on what was sent.
`get("/v1/models")` reports a created model as servable only after
`servable_after_gets` polls, so a test can drive the data-plane wait in
create_model."""
posts: list[tuple[str, BaseModel]] = field(default_factory=list)
servable_after_gets: int = 0
models_error: UnknownApiError | None = None
model_gets: int = 0
spend_total: int = 0
spend_gets: list[SpendLogsPageParams] = field(default_factory=list)
_created: list[str] = field(default_factory=list)
def post[R: BaseModel](
self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R]
) -> Result[R]:
self.posts.append((path, json))
if path == "/model/new" and isinstance(json, ModelNewBody):
self._created.append(json.model_name)
payload = (
{"model_id": "registered-id"} if response_type is ModelNewResponse else {}
)
return Success(data=response_type.model_validate(payload))
def stream(
self, path: str, *, headers: BaseModel, json: BaseModel
) -> StreamingResponse:
raise AssertionError("stream is not part of model management")
def send(
self,
path: str,
*,
headers: BaseModel,
json: BaseModel,
params: BaseModel | None = None,
stream: bool = False,
) -> StreamingResponse:
raise AssertionError("send is not part of model management")
def get[R: BaseModel](
self,
path: str,
*,
headers: BaseModel,
params: BaseModel,
response_type: type[R],
) -> Result[R]:
if path == "/v1/models" and response_type is ModelsListResponse:
self.model_gets += 1
if self.models_error is not None:
return self.models_error
visible = self._created if self.model_gets > self.servable_after_gets else []
return Success(
data=response_type.model_validate({"data": [{"id": name} for name in visible]})
)
if path == "/spend/logs/v2" and response_type is SpendLogsPage:
assert isinstance(params, SpendLogsPageParams)
self.spend_gets.append(params)
offset = (params.page - 1) * params.page_size
count = min(params.page_size, max(self.spend_total - offset, 0))
return Success(
data=response_type.model_validate(
{
"data": [{"request_id": f"req-{offset + i}"} for i in range(count)],
"total": self.spend_total,
"page": params.page,
"page_size": params.page_size,
"total_pages": (self.spend_total + params.page_size - 1) // params.page_size,
}
)
)
raise AssertionError(f"unexpected get: {path}")
def delete[R: BaseModel](
self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R]
) -> Result[R]:
raise AssertionError("delete is not part of model management")
def probe(self, path: str, *, params: BaseModel) -> ProbeResult:
raise AssertionError("probe is not part of model management")
def upload[R: BaseModel](
self,
path: str,
*,
headers: BaseModel,
form: FileUploadForm,
filename: str,
content: bytes,
params: BaseModel | None = None,
response_type: type[R],
) -> Result[R]:
raise AssertionError("upload is not part of model management")
def download(self, path: str, *, headers: BaseModel) -> StreamingResponse:
raise AssertionError("download is not part of model management")
def bearer(self, key: str) -> AuthHeaders:
return AuthHeaders(authorization=f"Bearer {key}")
@property
def master(self) -> AuthHeaders:
return self.bearer("sk-test-master")
def test_gateway_create_model_registers_deployment_and_returns_model_id() -> None:
transport = _RecordingTransport()
gateway = Gateway(transport=transport, poll_interval=0.0)
model_id = gateway.create_model(
"e2e-test-model", LiteLLMParamsBody(model="openai/gpt-4o-mini")
)
assert model_id == "registered-id"
path, body = transport.posts[0]
assert path == "/model/new"
assert isinstance(body, ModelNewBody)
assert body.model_name == "e2e-test-model"
# No pinned model_id: the proxy assigns a unique one, so a fixed-name model
# re-registered after a failed teardown can't collide on the id constraint.
assert body.model_info.id is None
assert body.model_info.mode is None
# It confirmed data-plane visibility before returning.
assert transport.model_gets >= 1
def test_gateway_create_model_waits_until_servable_on_the_data_plane() -> None:
# The model shows up on /v1/models only on the third poll (simulating the
# gateway's delayed DB reload in a split deployment); create_model must keep
# polling instead of returning after /model/new.
transport = _RecordingTransport(servable_after_gets=2)
gateway = Gateway(transport=transport, poll_interval=0.0)
gateway.create_model("e2e-late-model", LiteLLMParamsBody(model="openai/gpt-4o-mini"))
assert transport.model_gets == 3
def test_gateway_create_model_fails_loudly_when_never_servable() -> None:
transport = _RecordingTransport(servable_after_gets=10**9)
gateway = Gateway(transport=transport, poll_timeout=0.05, poll_interval=0.0)
with pytest.raises(AssertionError, match="never became servable"):
gateway.create_model("e2e-ghost-model", LiteLLMParamsBody(model="openai/gpt-4o-mini"))
def test_gateway_create_model_surfaces_the_last_data_plane_error() -> None:
transport = _RecordingTransport(
models_error=UnknownApiError(status_code=503, body="data plane down")
)
gateway = Gateway(transport=transport, poll_timeout=0.05, poll_interval=0.0)
with pytest.raises(AssertionError, match="data plane down") as excinfo:
gateway.create_model("e2e-flaky-model", LiteLLMParamsBody(model="openai/gpt-4o-mini"))
assert "503" in str(excinfo.value)
def test_batch_client_create_model_registers_a_batch_mode_deployment() -> None:
transport = _RecordingTransport()
client = BatchClient(gateway=Gateway(transport=transport, poll_interval=0.0))
model_id = client.create_model(
"e2e-batch-model", LiteLLMParamsBody(model="openai/gpt-4o-mini")
)
assert model_id == "registered-id"
path, body = transport.posts[0]
assert path == "/model/new"
assert isinstance(body, ModelNewBody)
assert body.model_info.mode == "batch"
def test_gateway_delete_model_posts_the_model_id() -> None:
transport = _RecordingTransport()
gateway = Gateway(transport=transport)
gateway.delete_model("registered-id")
path, body = transport.posts[0]
assert path == "/model/delete"
assert isinstance(body, ModelDeleteBody)
assert body.id == "registered-id"
WINDOW_START = datetime(2026, 7, 14, 12, 0, 0, tzinfo=timezone.utc)
WINDOW_END = datetime(2026, 7, 14, 14, 0, 0, tzinfo=timezone.utc)
def test_gateway_spend_logs_window_pages_through_every_row_in_the_window() -> None:
transport = _RecordingTransport(spend_total=250)
gateway = Gateway(transport=transport)
rows = gateway.spend_logs_window(start=WINDOW_START, end=WINDOW_END)
assert len(rows) == 250
assert len({row.request_id for row in rows}) == 250
assert [params.page for params in transport.spend_gets] == [1, 2, 3]
assert all(params.start_date == "2026-07-14 12:00:00" for params in transport.spend_gets)
assert all(params.end_date == "2026-07-14 14:00:00" for params in transport.spend_gets)
def test_gateway_spend_logs_window_stops_at_an_exact_page_boundary() -> None:
transport = _RecordingTransport(spend_total=200)
gateway = Gateway(transport=transport)
rows = gateway.spend_logs_window(start=WINDOW_START, end=WINDOW_END)
assert len(rows) == 200
assert [params.page for params in transport.spend_gets] == [1, 2]
def test_gateway_spend_logs_window_returns_empty_for_an_empty_window() -> None:
transport = _RecordingTransport(spend_total=0)
gateway = Gateway(transport=transport)
rows = gateway.spend_logs_window(start=WINDOW_START, end=WINDOW_END)
assert rows == []
assert [params.page for params in transport.spend_gets] == [1]
def test_spend_logs_params_rejects_the_unfiltered_whole_table_read() -> None:
with pytest.raises(ValidationError, match="spend_logs_window"):
SpendLogsParams()

View file

@ -1,46 +0,0 @@
"""Unit coverage for the lifecycle harness (lifecycle.run_case).
Cases register cleanups progressively during init() (create team, then user, then
key), so a failure partway through init() must still release whatever was already
created on the long-lived shared proxy. This guards that contract.
"""
from dataclasses import dataclass, field
from typing import Callable, List
import pytest
from lifecycle import run_case
@dataclass
class _PartialInitCase:
"""init() registers a cleanup, then raises before finishing - mirroring a real
case that creates a resource, registers its delete, then fails on the next
step."""
released: List[str] = field(default_factory=list)
_undo: List[Callable[[], None]] = field(default_factory=list)
def init(self) -> None:
self._undo.append(lambda: self.released.append("first"))
raise RuntimeError("init failed after registering the first resource")
def run(self) -> None:
raise AssertionError("run() must not execute when init() failed")
def teardown(self) -> None:
for undo in reversed(self._undo):
undo()
def test_run_case_releases_resources_when_init_fails_partway() -> None:
case = _PartialInitCase()
with pytest.raises(RuntimeError, match="init failed"):
run_case(case)
assert case.released == ["first"], (
"a resource registered before init() failed must still be released, or it "
"leaks on the long-lived shared proxy"
)

View file

@ -1,52 +0,0 @@
"""Unit coverage for SplitTransport path routing (is_control_plane_path).
Model-management calls (/model/new, /model/delete, /model/info) must go to the
control plane: the data-plane gateway does not serve management routes, so a
misrouted /model/new 404s and takes down every suite that registers deployments
at runtime (llm_translation, batches, access_control). /models must stay on the
data plane; it is the OpenAI-compatible list-models route, not a management
route.
"""
import pytest
from transport import is_control_plane_path
@pytest.mark.parametrize(
"path",
[
"/model/new",
"/model/delete",
"/model/update",
"/model/info",
"/key/generate",
"/budget/new",
"/spend/logs",
"/end_user/daily/activity",
"/user/daily/activity",
"/team/daily/activity",
"/tag/daily/activity",
],
)
def test_management_routes_go_to_the_control_plane(path: str) -> None:
assert is_control_plane_path(path), (
f"{path} is a management route; sending it to the data plane 404s"
)
@pytest.mark.parametrize(
"path",
[
"/models",
"/v1/models",
"/chat/completions",
"/v1/messages",
"/embeddings",
"/anthropic/v1/messages",
],
)
def test_llm_routes_stay_on_the_data_plane(path: str) -> None:
assert not is_control_plane_path(path), (
f"{path} is an LLM route; it must go to the data plane"
)

View file

@ -88,6 +88,60 @@ async def test_dual_cache_async_set_cache_injects_default_in_memory_ttl():
assert expiry <= after + 60
@pytest.mark.asyncio
async def test_dual_cache_redis_backfill_injects_default_in_memory_ttl():
"""
A Redis-hit backfill into the in-memory tier must honor
default_in_memory_ttl the same way the write paths do. Without it, the
backfilled entry falls to InMemoryCache's own default_ttl (600s), so a
replica that primed a management object (e.g. a virtual key's auth blob)
from Redis keeps serving it for 10 minutes after the object was updated
and invalidated, instead of re-reading within the configured TTL.
"""
in_memory_cache = InMemoryCache(default_ttl=600)
redis_cache = MagicMock()
redis_cache.async_get_cache = AsyncMock(return_value="redis_value")
dual_cache = DualCache(
in_memory_cache=in_memory_cache,
redis_cache=redis_cache,
default_in_memory_ttl=60,
)
before = time.time()
result = await dual_cache.async_get_cache(key="backfill_key")
after = time.time()
assert result == "redis_value"
expiry = in_memory_cache.ttl_dict["backfill_key"]
assert expiry >= before + 60
assert expiry <= after + 60
@pytest.mark.asyncio
async def test_dual_cache_batch_redis_backfill_injects_default_in_memory_ttl():
"""async_batch_get_cache's Redis-to-memory backfill must honor
default_in_memory_ttl, same as the single-key path."""
in_memory_cache = InMemoryCache(default_ttl=600)
mock_redis = MagicMock(spec=RedisCache)
mock_redis.async_batch_get_cache = AsyncMock(
return_value={"batch_backfill_key": "redis_value"}
)
dual_cache = DualCache(
in_memory_cache=in_memory_cache,
redis_cache=mock_redis,
default_in_memory_ttl=60,
)
before = time.time()
result = await dual_cache.async_batch_get_cache(keys=["batch_backfill_key"])
after = time.time()
assert result == ["redis_value"]
expiry = in_memory_cache.ttl_dict["batch_backfill_key"]
assert expiry >= before + 60
assert expiry <= after + 60
@pytest.mark.asyncio
async def test_dual_cache_async_set_cache_respects_explicit_ttl():
"""

View file

@ -522,6 +522,49 @@ async def test_get_key_object_should_raise_if_reconnect_fails_on_db_connection_e
assert mock_prisma_client.get_data.await_count == 1
def _fake_redis_cache():
fake_redis = MagicMock()
fake_redis.async_get_cache = AsyncMock(return_value=None)
fake_redis.async_set_cache = AsyncMock()
fake_redis.async_set_cache_pipeline = AsyncMock()
fake_redis.async_delete_cache = AsyncMock()
return fake_redis
class TestAuthCacheRedisWritePolicy:
"""Redis auth-cache entries may only be written from fresh DB loads.
With ``enable_redis_auth_cache`` and multiple replicas, a pod that re-publishes
a cache-derived key object to Redis can resurrect a stale auth blob after
``/key/update`` or ``/key/delete`` already deleted it, so limit changes never
propagate fleet-wide while traffic keeps refreshing the stale entry's TTL.
"""
@pytest.mark.asyncio
async def test_get_key_object_db_load_publishes_to_redis(self):
mock_prisma_client = MagicMock()
mock_prisma_client.get_data = AsyncMock(
return_value=UserAPIKeyAuth(token="hashed-token-db")
)
fake_redis = _fake_redis_cache()
cache = UserApiKeyCache()
cache.redis_cache = fake_redis
key_obj = await get_key_object(
hashed_token="hashed-token-db",
prisma_client=mock_prisma_client,
user_api_key_cache=cache,
)
assert key_obj.token == "hashed-token-db"
fake_redis.async_set_cache.assert_awaited_once()
assert (
fake_redis.async_set_cache.await_args.kwargs.get("key")
or fake_redis.async_set_cache.await_args.args[0]
) == "hashed-token-db"
def test_get_cli_jwt_auth_token_default_expiration(valid_sso_user_defined_values):
"""Test generating CLI JWT token with default 24-hour expiration"""
token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values)

View file

@ -1,3 +1,4 @@
import asyncio
import json
import os
import sys
@ -4161,6 +4162,99 @@ async def test_auth_path_caches_team_object_under_canonical_team_id_key():
assert cache.get_cache(key=None) is None
@pytest.mark.asyncio
async def test_auth_does_not_rewrite_cached_key_object_back_into_cache():
"""A cache-hit auth must not write the token back into the cache.
Re-writing on every auth let a replica holding a stale in-memory token
republish it to shared Redis with a fresh TTL on each request, so
/key/update and /key/delete never propagated across replicas or regional
Redis while the key kept calling (stale auth re-cache feedback loop).
Only the DB-load paths (IdentityStore._resolve_key / get_key_object) may
populate the cache.
"""
from fastapi import Request
from starlette.datastructures import URL
import litellm.proxy.proxy_server as _proxy_server_mod
from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.proxy_server import hash_token
api_key = "sk-lit-cached-key-no-rewrite"
hashed_key = hash_token(api_key)
key_cache = UserApiKeyCache()
stale_token = UserAPIKeyAuth(
api_key=api_key,
token=hashed_key,
metadata={"model_rpm_limit": {"gpt-5.4-mini": 3}},
last_refreshed_at=1000.0,
)
await key_cache.async_set_cache(
key=hashed_key, value=stale_token, model_type=UserAPIKeyAuth
)
fetch_from_db = AsyncMock(
side_effect=AssertionError("cache-hit auth must not touch the DB")
)
proxy_logging_obj = MagicMock()
proxy_logging_obj.internal_usage_cache = MagicMock()
proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
attrs = {
"prisma_client": MagicMock(),
"user_api_key_cache": key_cache,
"proxy_logging_obj": proxy_logging_obj,
"master_key": "sk-test-master",
"general_settings": {},
"llm_model_list": [],
"llm_router": None,
"open_telemetry_logger": None,
"model_max_budget_limiter": MagicMock(),
"user_custom_auth": None,
"jwt_handler": None,
"litellm_proxy_admin_name": "admin",
}
originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs}
try:
for k, v in attrs.items():
setattr(_proxy_server_mod, k, v)
request = Request(scope={"type": "http"})
request._url = URL(url="/chat/completions")
with patch(
"litellm.proxy.auth.resolvers.store._fetch_key_object_from_db_with_reconnect",
fetch_from_db,
):
result = await _user_api_key_auth_builder(
request=request,
api_key=f"Bearer {api_key}",
azure_api_key_header="",
anthropic_api_key_header=None,
google_ai_studio_api_key_header=None,
azure_apim_header=None,
request_data={},
)
pending = [t for t in asyncio.all_tasks() if t is not asyncio.current_task()]
if pending:
await asyncio.wait(pending, timeout=5)
assert result.token == hashed_key
fetch_from_db.assert_not_called()
cached_after = await key_cache.async_get_cache(
key=hashed_key, model_type=UserAPIKeyAuth
)
assert cached_after is not None
assert cached_after.last_refreshed_at == 1000.0
assert cached_after.metadata == {"model_rpm_limit": {"gpt-5.4-mini": 3}}
finally:
for k, v in originals.items():
setattr(_proxy_server_mod, k, v)
class TestCheckKeyModelBudgetWithFallback:
"""`_check_key_model_budget_with_fallback` must reroute a request to the
first configured `budget_fallbacks` entry still within its own budget,

View file

@ -1,9 +1,10 @@
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
from litellm.proxy._types import LiteLLM_UserTable
from litellm.proxy.management_endpoints.scim.scim_v2 import patch_user
from litellm.proxy.management_endpoints.scim.scim_v2 import _apply_patch_ops, patch_user
from litellm.types.proxy.management_endpoints.scim_v2 import (
SCIMPatchOp,
SCIMPatchOperation,
@ -329,3 +330,138 @@ async def test_patch_user_multiple_fields_without_path():
assert update_data["user_alias"] == "New Display Name"
assert "" not in metadata # Ensure no empty string key
assert result.active is False
def _user_with_metadata(metadata):
return LiteLLM_UserTable(
user_id="user-mva",
user_email="mva@example.com",
user_alias=None,
teams=[],
metadata=metadata,
)
def test_apply_patch_ops_replace_entitlements_writes_canonical_key():
"""A PATCH on path=entitlements must persist under scim_entitlements, not
fall through to the generic handler's raw path key"""
patch_ops = SCIMPatchOp(
Operations=[
SCIMPatchOperation(
op="replace",
path="entitlements",
value=[{"value": "jira-software", "display": "Jira Software"}],
)
]
)
update_data, _ = _apply_patch_ops(
existing_user=_user_with_metadata({}), patch_ops=patch_ops
)
metadata = update_data["metadata"]
assert metadata["scim_entitlements"] == [
{"value": "jira-software", "display": "Jira Software"}
]
assert "entitlements" not in metadata
def test_apply_patch_ops_add_roles_appends_to_existing():
patch_ops = SCIMPatchOp(
Operations=[
SCIMPatchOperation(op="add", path="roles", value=[{"value": "admin"}])
]
)
update_data, _ = _apply_patch_ops(
existing_user=_user_with_metadata({"scim_roles": [{"value": "viewer"}]}),
patch_ops=patch_ops,
)
assert update_data["metadata"]["scim_roles"] == [
{"value": "viewer"},
{"value": "admin"},
]
def test_apply_patch_ops_remove_entitlements_clears_canonical_key():
patch_ops = SCIMPatchOp(
Operations=[SCIMPatchOperation(op="remove", path="entitlements")]
)
update_data, _ = _apply_patch_ops(
existing_user=_user_with_metadata(
{"scim_entitlements": [{"value": "jira-software"}]}
),
patch_ops=patch_ops,
)
assert "scim_entitlements" not in update_data["metadata"]
def test_apply_patch_ops_pathless_value_dict_handles_roles():
patch_ops = SCIMPatchOp(
Operations=[
SCIMPatchOperation(
op="replace",
value={"roles": [{"value": "engineering-admin", "primary": True}]},
)
]
)
update_data, _ = _apply_patch_ops(
existing_user=_user_with_metadata({}), patch_ops=patch_ops
)
assert update_data["metadata"]["scim_roles"] == [
{"value": "engineering-admin", "primary": True}
]
def test_apply_patch_ops_invalid_entitlements_value_raises_400():
patch_ops = SCIMPatchOp(
Operations=[
SCIMPatchOperation(
op="replace", path="entitlements", value=[{"display": "no value"}]
)
]
)
with pytest.raises(HTTPException) as exc_info:
_apply_patch_ops(existing_user=_user_with_metadata({}), patch_ops=patch_ops)
assert exc_info.value.status_code == 400
def test_apply_patch_ops_add_without_value_raises_400_naming_value_member():
patch_ops = SCIMPatchOp(
Operations=[SCIMPatchOperation(op="add", path="entitlements")]
)
with pytest.raises(HTTPException) as exc_info:
_apply_patch_ops(existing_user=_user_with_metadata({}), patch_ops=patch_ops)
assert exc_info.value.status_code == 400
assert "value" in str(exc_info.value.detail)
def test_apply_patch_ops_filtered_path_raises_400_instead_of_junk_metadata():
"""A filtered path must fail loudly rather than fall through to the generic
handler, which would write a junk metadata key while reporting success"""
patch_ops = SCIMPatchOp(
Operations=[
SCIMPatchOperation(
op="remove", path='roles[value eq "engineering-admin"]'
)
]
)
with pytest.raises(HTTPException) as exc_info:
_apply_patch_ops(
existing_user=_user_with_metadata(
{"scim_roles": [{"value": "engineering-admin"}]}
),
patch_ops=patch_ops,
)
assert exc_info.value.status_code == 400

View file

@ -15,6 +15,7 @@ from litellm.proxy.management_endpoints.scim.scim_transformations import (
from litellm.types.proxy.management_endpoints.scim_v2 import (
SCIM_ENTERPRISE_USER_SCHEMA,
SCIMEnterpriseUser,
SCIMMultiValuedAttribute,
SCIMPatchOperation,
SCIMUser,
)
@ -179,6 +180,74 @@ class TestScimTransformations:
assert scim_user.enterprise_user.department == "Platform"
assert SCIM_ENTERPRISE_USER_SCHEMA in scim_user.schemas
@pytest.mark.asyncio
async def test_transform_user_with_entitlements_and_roles_metadata(
self, mock_prisma_client
):
mock_client, mock_find_unique = mock_prisma_client
mock_find_unique.return_value = None
user = LiteLLM_UserTable(
user_id="user-entitled",
user_email="entitled@example.com",
user_alias=None,
teams=[],
created_at=None,
updated_at=None,
metadata={
"scim_entitlements": [
{"value": "jira-software", "display": "Jira Software"}
],
"scim_roles": [{"value": "engineering-admin", "primary": True}],
},
)
with patch("litellm.proxy.proxy_server.prisma_client", mock_client):
scim_user = await ScimTransformations.transform_litellm_user_to_scim_user(
user
)
assert scim_user.entitlements is not None
assert scim_user.entitlements[0].value == "jira-software"
assert scim_user.entitlements[0].display == "Jira Software"
assert scim_user.roles is not None
assert scim_user.roles[0].value == "engineering-admin"
assert scim_user.roles[0].primary is True
@pytest.mark.asyncio
async def test_transform_user_with_malformed_directory_metadata_fails_soft(
self, mock_prisma_client
):
"""Metadata is writable outside the SCIM surface; a corrupted value on one
user must omit the attribute, not fail the whole directory response"""
mock_client, mock_find_unique = mock_prisma_client
mock_find_unique.return_value = None
user = LiteLLM_UserTable(
user_id="user-corrupt",
user_email="corrupt@example.com",
user_alias=None,
teams=[],
created_at=None,
updated_at=None,
metadata={
"scim_entitlements": [{"display": 123}],
"scim_roles": {"value": "not-a-list"},
"scim_enterprise": {"manager": 42},
},
)
with patch("litellm.proxy.proxy_server.prisma_client", mock_client):
scim_user = await ScimTransformations.transform_litellm_user_to_scim_user(
user
)
assert scim_user.id == "user-corrupt"
assert scim_user.entitlements is None
assert scim_user.roles is None
assert scim_user.enterprise_user is None
assert SCIM_ENTERPRISE_USER_SCHEMA not in scim_user.schemas
@pytest.mark.asyncio
async def test_transform_user_without_enterprise_metadata_omits_schema(
self, mock_user, mock_prisma_client
@ -223,6 +292,27 @@ class TestScimTransformations:
dumped_ent = with_enterprise.model_dump(by_alias=True)
assert dumped_ent[SCIM_ENTERPRISE_USER_SCHEMA]["costCenter"] == "CC-42"
def test_scim_user_serialization_omits_absent_entitlements_and_roles(self):
without_attrs = SCIMUser(
schemas=["urn:ietf:params:scim:schemas:core:2.0:User"],
id="user-1",
userName="user@example.com",
)
dumped = without_attrs.model_dump(by_alias=True)
assert "entitlements" not in dumped
assert "roles" not in dumped
with_attrs = SCIMUser(
schemas=["urn:ietf:params:scim:schemas:core:2.0:User"],
id="user-2",
userName="entitled@example.com",
entitlements=[SCIMMultiValuedAttribute(value="jira-software")],
roles=[SCIMMultiValuedAttribute(value="engineering-admin")],
)
dumped_attrs = with_attrs.model_dump(by_alias=True)
assert dumped_attrs["entitlements"][0]["value"] == "jira-software"
assert dumped_attrs["roles"][0]["value"] == "engineering-admin"
@pytest.mark.asyncio
async def test_transform_litellm_team_to_scim_group(
self, mock_team, mock_prisma_client

View file

@ -172,6 +172,70 @@ async def test_create_user_ingests_enterprise_extension(mocker, monkeypatch):
}
@pytest.mark.asyncio
async def test_create_user_ingests_entitlements_and_roles(mocker, monkeypatch):
"""A SCIM create payload carrying entitlements and roles should land in the
created user's metadata under scim_entitlements and scim_roles"""
scim_user = SCIMUser.model_validate(
{
"schemas": ["urn:ietf:params:scim:schemas:core:2.0:User"],
"userName": "entitled-user",
"name": {"familyName": "User", "givenName": "Entitled"},
"emails": [{"value": "entitled@example.com"}],
"entitlements": [
{
"value": "jira-software",
"display": "Jira Software",
"type": "app",
"primary": True,
},
"bare-entitlement",
],
"roles": [{"value": "engineering-admin", "type": "role"}],
}
)
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.db = mocker.MagicMock()
mock_prisma_client.db.litellm_usertable = mocker.MagicMock()
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None)
monkeypatch.setattr("litellm.default_internal_user_params", None, raising=False)
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception",
AsyncMock(return_value=mock_prisma_client),
)
new_user_mock = mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2.new_user",
AsyncMock(return_value=NewUserRequest(user_id="entitled-user")),
)
mocker.patch(
"litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_user_to_scim_user",
AsyncMock(return_value=scim_user),
)
await create_user(user=scim_user)
created_metadata = new_user_mock.call_args.kwargs["data"].metadata
assert created_metadata["scim_entitlements"] == [
{
"value": "jira-software",
"display": "Jira Software",
"type": "app",
"primary": True,
},
{"value": "bare-entitlement"},
]
assert created_metadata["scim_roles"] == [
{"value": "engineering-admin", "type": "role"}
]
@pytest.mark.asyncio
async def test_create_user_uses_default_internal_user_params_role(mocker, monkeypatch):
"""If role is set in default_internal_user_params, new user should use that role"""

View file

@ -555,6 +555,85 @@ async def test_new_team_with_mcp_tool_permissions(mock_db_client, mock_admin_aut
assert created_permission_data["mcp_servers"] == ["server_a", "server_b"]
@pytest.mark.parametrize(
"user_role,user_id,flag_value,expected",
[
(LitellmUserRoles.PROXY_ADMIN, "admin-1", True, False),
(LitellmUserRoles.PROXY_ADMIN, "admin-1", False, True),
(LitellmUserRoles.PROXY_ADMIN, "admin-1", None, True),
(LitellmUserRoles.INTERNAL_USER, "user-1", True, True),
(LitellmUserRoles.ORG_ADMIN, "org-admin-1", True, True),
(LitellmUserRoles.PROXY_ADMIN, None, False, False),
],
)
def test_should_auto_add_team_creator(user_role, user_id, flag_value, expected):
from litellm.proxy.management_endpoints.team_endpoints import (
_should_auto_add_team_creator,
)
general_settings = (
{} if flag_value is None else {"disable_auto_add_proxy_admin_to_teams": flag_value}
)
auth = UserAPIKeyAuth(user_role=user_role, user_id=user_id)
assert _should_auto_add_team_creator(auth, general_settings) is expected
@pytest.mark.asyncio
@pytest.mark.parametrize(
"disable_flag,expect_creator_added", [(True, False), (False, True)]
)
async def test_new_team_disable_auto_add_proxy_admin_flag(
mock_db_client, disable_flag, expect_creator_added
):
"""
When general_settings.disable_auto_add_proxy_admin_to_teams is True, a proxy
admin calling /team/new must NOT be auto-added to the team's members. When
the flag is off, the creator is auto-added as a team admin (default
behavior, regression guard for LIT-3739).
"""
mock_db_client.jsonify_team_object = lambda db_data: db_data
mock_db_client.get_data = AsyncMock(return_value=None)
mock_db_client.update_data = AsyncMock(return_value=MagicMock())
mock_db_client.db = MagicMock()
team_create_result = MagicMock(team_id="team-789")
team_create_result.model_dump.return_value = {"team_id": "team-789"}
mock_db_client.db.litellm_teamtable = MagicMock()
mock_db_client.db.litellm_teamtable.create = AsyncMock(
return_value=team_create_result
)
mock_db_client.db.litellm_teamtable.count = AsyncMock(return_value=0)
mock_db_client.db.litellm_usertable = MagicMock()
mock_db_client.db.litellm_usertable.update = AsyncMock(return_value=MagicMock())
from fastapi import Request
from litellm.proxy._types import NewTeamRequest
from litellm.proxy.management_endpoints.team_endpoints import new_team
admin_auth = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-user-1"
)
with patch(
"litellm.proxy.proxy_server.general_settings",
{"disable_auto_add_proxy_admin_to_teams": disable_flag},
), patch(
"litellm.proxy.management_endpoints.team_endpoints._add_team_members_to_team",
new_callable=AsyncMock,
) as mock_add_members:
await new_team(
data=NewTeamRequest(team_alias="flag-test-team"),
http_request=MagicMock(spec=Request),
user_api_key_dict=admin_auth,
)
mock_add_members.assert_called_once()
member_add_request = mock_add_members.call_args.kwargs["data"]
member_user_ids = [m.user_id for m in member_add_request.member]
assert ("admin-user-1" in member_user_ids) is expect_creator_added
@pytest.mark.asyncio
async def test_team_update_object_permissions_existing_permission(monkeypatch):
"""

View file

@ -4568,6 +4568,128 @@ async def test_update_cache_pipeline_honors_user_api_key_cache_ttl():
setattr(litellm.proxy.proxy_server, "user_api_key_cache", original_cache)
@pytest.mark.asyncio
async def test_spend_tracking_never_writes_the_auth_object_back():
"""Spend tracking must never write the auth object back into the cache.
Writing the mutated auth object back after every priced request let a
stale copy be re-published with a fresh TTL: to shared Redis it defeated
/key/update and /key/delete across replicas, and even a local-only write
could race an invalidation and resurrect a revoked key on this worker.
Spend is tracked through the spend:key:* counters, so the auth object is
only ever written by the DB-load paths.
"""
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
original_cache = litellm.proxy.proxy_server.user_api_key_cache
cache = UserApiKeyCache()
setattr(litellm.proxy.proxy_server, "user_api_key_cache", cache)
try:
hashed_token = "spend-tracking-no-writeback-token"
await cache.async_set_cache(
key=hashed_token,
value=UserAPIKeyAuth(token=hashed_token, spend=1.0),
model_type=UserAPIKeyAuth,
)
with (
patch.object(
cache, "async_set_cache_pipeline", new=AsyncMock()
) as mock_pipeline,
patch.object(cache, "async_set_cache", new=AsyncMock()) as mock_set,
):
await litellm.proxy.proxy_server.update_cache(
token=hashed_token,
user_id=None,
end_user_id=None,
team_id=None,
response_cost=5.0,
parent_otel_span=None,
)
pending = [
t for t in asyncio.all_tasks() if t is not asyncio.current_task()
]
if pending:
await asyncio.wait(pending, timeout=5)
key_pipeline_writes = [
call
for call in mock_pipeline.call_args_list
if any(k == hashed_token for k, _ in call.kwargs["cache_list"])
]
assert key_pipeline_writes == []
mock_set.assert_not_called()
finally:
setattr(litellm.proxy.proxy_server, "user_api_key_cache", original_cache)
@pytest.mark.asyncio
async def test_update_cache_global_proxy_spend_scalar_stays_shared():
"""
The proxy-wide spend estimate must keep flowing to Redis when the spend
writeback goes per-pod: the global max_budget check reads the
``{litellm_proxy_admin_name}:spend`` cache entry between authoritative DB
reloads, so keeping it pod-local would let traffic spread across replicas
exceed the proxy budget by roughly a factor of the replica count within a
cache TTL. Sharing this scalar is safe because it carries no limits or
permissions, so it cannot resurrect an invalidated auth blob.
"""
from litellm.caching.caching import DualCache
admin_name = litellm.proxy.proxy_server.litellm_proxy_admin_name
global_key = "{}:spend".format(admin_name)
async def fake_get(key, **kwargs):
if key == "user-lit":
return {"user_id": "user-lit", "spend": 1.0}
if key == global_key:
return 10.0
return None
original_cache = litellm.proxy.proxy_server.user_api_key_cache
cache = DualCache(default_in_memory_ttl=300)
setattr(litellm.proxy.proxy_server, "user_api_key_cache", cache)
try:
with patch.object(
cache, "async_get_cache", new=AsyncMock(side_effect=fake_get)
):
with patch.object(
cache, "async_set_cache_pipeline", new=AsyncMock()
) as mock_set_cache:
await litellm.proxy.proxy_server.update_cache(
token=None,
user_id="user-lit",
end_user_id=None,
team_id=None,
response_cost=5.0,
parent_otel_span=None,
)
pending = [
t for t in asyncio.all_tasks() if t is not asyncio.current_task()
]
if pending:
await asyncio.wait(pending, timeout=5)
calls = mock_set_cache.await_args_list
local_keys = [
k
for c in calls
if c.kwargs.get("local_only") is True
for k, _ in c.kwargs["cache_list"]
]
shared_keys = [
k
for c in calls
if c.kwargs.get("local_only") is not True
for k, _ in c.kwargs["cache_list"]
]
assert "user-lit" in local_keys
assert global_key not in local_keys
assert shared_keys == [global_key]
finally:
setattr(litellm.proxy.proxy_server, "user_api_key_cache", original_cache)
@pytest.mark.asyncio
async def test_init_sso_settings_in_db():
"""
@ -6265,6 +6387,32 @@ async def test_update_general_settings_store_model_in_db_false():
assert ps.general_settings["store_model_in_db"] is False
@pytest.mark.asyncio
@pytest.mark.parametrize(
"db_value,expected",
[(True, True), (False, False), ("true", True), ("false", False), (None, None)],
)
async def test_update_general_settings_disable_auto_add_proxy_admin_to_teams(db_value, expected):
"""
Verify _update_general_settings propagates disable_auto_add_proxy_admin_to_teams
from the DB config into the live general_settings dict, so a UI toggle via
/config/field/update takes effect on the next config poll instead of
requiring a proxy restart.
"""
from litellm.proxy.proxy_server import ProxyConfig
proxy_config = ProxyConfig()
with patch("litellm.proxy.proxy_server.general_settings", {}):
await proxy_config._update_general_settings(
db_general_settings={"disable_auto_add_proxy_admin_to_teams": db_value}
)
import litellm.proxy.proxy_server as ps
assert ps.general_settings["disable_auto_add_proxy_admin_to_teams"] is expected
@pytest.mark.asyncio
async def test_update_general_settings_store_model_in_db_string_normalization():
"""

View file

@ -124,3 +124,18 @@ def test_proxy_exception_str_returns_message():
"param": "key",
"code": "401",
}
def test_key_request_router_settings_keeps_enable_tag_filtering():
"""``router_settings`` on key requests validates through
``UpdateRouterConfig``; a field missing from that model is silently
dropped at parse time, so a key's "Enable Tag Filtering" toggle would
never reach the DB even though the team path (plain dict) kept it."""
from litellm.proxy._types import GenerateKeyRequest
req = GenerateKeyRequest(router_settings={"enable_tag_filtering": True, "num_retries": 2})
assert req.router_settings is not None
dumped = req.router_settings.model_dump(exclude_none=True)
assert dumped["enable_tag_filtering"] is True
assert dumped["num_retries"] == 2

View file

@ -819,3 +819,71 @@ async def test_route_request_realtime_transcription_session_resolves_credentials
)
assert mock_handler.call_args.kwargs["api_key"] == "transcription-key"
@pytest.mark.asyncio
async def test_route_request_merges_enable_tag_filtering_from_override():
"""Key/team router_settings carry enable_tag_filtering; the override
whitelist must forward it to the router call or the team's tag-routing
toggle saved in the UI is silently ignored at request time."""
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "Hello"}],
"router_settings_override": {
"enable_tag_filtering": True,
},
}
llm_router = MagicMock()
llm_router.acompletion.return_value = "success"
response = await route_request(data, llm_router, None, "acompletion")
assert response == "success"
call_kwargs = llm_router.acompletion.call_args[1]
assert call_kwargs["enable_tag_filtering"] is True
@pytest.mark.asyncio
async def test_route_request_strips_client_supplied_enable_tag_filtering():
"""enable_tag_filtering influences deployment selection and is only
trusted when it comes from key/team router_settings via
router_settings_override. A caller putting it in the request body must
not reach the router with it."""
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "Hello"}],
"enable_tag_filtering": True,
}
llm_router = MagicMock()
llm_router.acompletion.return_value = "ok"
await route_request(data, llm_router, None, "acompletion")
call_kwargs = llm_router.acompletion.call_args[1]
assert "enable_tag_filtering" not in call_kwargs
assert "enable_tag_filtering" not in data
@pytest.mark.asyncio
async def test_route_request_override_enable_tag_filtering_beats_body_value():
"""A client-sent enable_tag_filtering must not shadow the key/team
setting: the body copy is stripped first, so the override value is the
one the router sees."""
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "Hello"}],
"enable_tag_filtering": False,
"router_settings_override": {
"enable_tag_filtering": True,
},
}
llm_router = MagicMock()
llm_router.acompletion.return_value = "ok"
await route_request(data, llm_router, None, "acompletion")
call_kwargs = llm_router.acompletion.call_args[1]
assert call_kwargs["enable_tag_filtering"] is True

View file

@ -1019,3 +1019,99 @@ async def test_negation_removes_tag_regex_deployment_falls_to_ban_only():
mock_response="hi",
)
assert response._hidden_params["model_id"] == "openai-deployment"
@pytest.mark.asyncio()
async def test_request_level_enable_tag_filtering_applies_when_global_off():
"""
A request carrying enable_tag_filtering=True (set by the proxy from key/team
router_settings) must activate tag filtering even when the router-level flag
is off. Without this, a team's "Enable Tag Filtering" toggle saved in the UI
is silently ignored at request time.
"""
router = litellm.Router(
model_list=[
{
"model_name": "gpt-4",
"litellm_params": {
"model": "gpt-4o",
"api_base": "https://exampleopenaiendpoint-production.up.railway.app/",
"tags": ["teamA"],
},
"model_info": {"id": "team-a-deployment"},
},
{
"model_name": "gpt-4",
"litellm_params": {
"model": "gpt-4o-mini",
"api_base": "https://exampleopenaiendpoint-production.up.railway.app/",
"tags": ["teamB"],
},
"model_info": {"id": "team-b-deployment"},
},
],
enable_tag_filtering=False,
)
for _ in range(5):
response = await router.acompletion(
model="gpt-4",
messages=[{"role": "user", "content": "hi"}],
metadata={"tags": ["teamA"]},
enable_tag_filtering=True,
mock_response="hi",
)
assert response._hidden_params["model_id"] == "team-a-deployment"
for _ in range(5):
response = await router.acompletion(
model="gpt-4",
messages=[{"role": "user", "content": "hi"}],
metadata={"tags": ["teamB"]},
enable_tag_filtering=True,
mock_response="hi",
)
assert response._hidden_params["model_id"] == "team-b-deployment"
@pytest.mark.asyncio()
async def test_request_level_enable_tag_filtering_false_cannot_disable_global():
"""
A request-level enable_tag_filtering=False must not bypass a router-level
True: tag filtering can be an operator-level restriction on which
deployments a caller may reach, so per-request settings may only scope
down, never escape the global policy.
"""
router = litellm.Router(
model_list=[
{
"model_name": "gpt-4",
"litellm_params": {
"model": "gpt-4o",
"api_base": "https://exampleopenaiendpoint-production.up.railway.app/",
"tags": ["teamA"],
},
"model_info": {"id": "team-a-deployment"},
},
{
"model_name": "gpt-4",
"litellm_params": {
"model": "gpt-4o-mini",
"api_base": "https://exampleopenaiendpoint-production.up.railway.app/",
"tags": ["teamB"],
},
"model_info": {"id": "team-b-deployment"},
},
],
enable_tag_filtering=True,
)
for _ in range(5):
response = await router.acompletion(
model="gpt-4",
messages=[{"role": "user", "content": "hi"}],
metadata={"tags": ["teamA"]},
enable_tag_filtering=False,
mock_response="hi",
)
assert response._hidden_params["model_id"] == "team-a-deployment"

View file

@ -21,6 +21,7 @@ import {
getSemanticConfigError,
} from "./build_complexity_router_config";
import { buildAutoRouterTestTargets, AutoRouterTestTarget } from "./build_auto_router_test_targets";
import { getSemanticRouterError } from "./build_semantic_router_validation";
import AutoRouterConnectionTest from "./auto_router_connection_test";
import NotificationManager from "../molecules/notifications_manager";
@ -164,23 +165,13 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({ form, handleOk, acc
};
const submitSemanticRouter = (name: string) => {
if (!form.getFieldValue("auto_router_default_model")) {
NotificationManager.fromBackend("Please select a Default Model");
return;
}
if (!routerConfig || !routerConfig.routes || routerConfig.routes.length === 0) {
NotificationManager.fromBackend("Please configure at least one route for the auto router");
return;
}
const invalidRoutes = routerConfig.routes.filter(
(route: any) => !route.name || !route.description || route.utterances.length === 0,
);
if (invalidRoutes.length > 0) {
NotificationManager.fromBackend(
"Please ensure all routes have a target model, description, and at least one utterance",
);
const validationError = getSemanticRouterError({
defaultModel: form.getFieldValue("auto_router_default_model"),
embeddingModel: form.getFieldValue("auto_router_embedding_model"),
routerConfig,
});
if (validationError) {
NotificationManager.fromBackend(validationError);
return;
}
@ -358,18 +349,18 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({ form, handleOk, acc
</Form.Item>
<Form.Item
rules={[{ required: true, message: "Embedding model is required" }]}
label="Embedding Model"
name="auto_router_embedding_model"
tooltip="Optional: embedding model to use for semantic routing decisions"
tooltip="Embedding model to use for semantic routing decisions"
labelCol={{ span: 10 }}
labelAlign="left"
>
<AntdSelect
placeholder="Select an embedding model (optional)"
placeholder="Select an embedding model"
options={modelGroupOptions}
style={{ width: "100%" }}
showSearch
allowClear
/>
</Form.Item>
</>

View file

@ -0,0 +1,67 @@
import { getSemanticRouterError, SemanticRouterConfig } from "./build_semantic_router_validation";
const validRouterConfig: SemanticRouterConfig = {
routes: [{ name: "gpt-4o", description: "general chat", utterances: ["hello there"] }],
};
describe("getSemanticRouterError", () => {
it("requires an embedding model once the default model and routes are configured", () => {
expect(
getSemanticRouterError({
defaultModel: "gpt-4o",
embeddingModel: undefined,
routerConfig: validRouterConfig,
}),
).toBe("Please select an Embedding Model");
});
it("treats an empty embedding model string as missing", () => {
expect(
getSemanticRouterError({
defaultModel: "gpt-4o",
embeddingModel: "",
routerConfig: validRouterConfig,
}),
).toBe("Please select an Embedding Model");
});
it("passes when an embedding model is selected", () => {
expect(
getSemanticRouterError({
defaultModel: "gpt-4o",
embeddingModel: "text-embedding-3-large",
routerConfig: validRouterConfig,
}),
).toBeNull();
});
it("flags a missing default model before checking the embedding model", () => {
expect(
getSemanticRouterError({
defaultModel: undefined,
embeddingModel: undefined,
routerConfig: validRouterConfig,
}),
).toBe("Please select a Default Model");
});
it("flags missing routes before checking the embedding model", () => {
expect(
getSemanticRouterError({
defaultModel: "gpt-4o",
embeddingModel: undefined,
routerConfig: { routes: [] },
}),
).toBe("Please configure at least one route for the auto router");
});
it("validates route completeness after the embedding model is set", () => {
expect(
getSemanticRouterError({
defaultModel: "gpt-4o",
embeddingModel: "text-embedding-3-large",
routerConfig: { routes: [{ name: "gpt-4o", description: "", utterances: [] }] },
}),
).toBe("Please ensure all routes have a target model, description, and at least one utterance");
});
});

View file

@ -0,0 +1,29 @@
export interface SemanticRouterRoute {
name?: string;
description?: string;
utterances?: unknown[];
}
export interface SemanticRouterConfig {
routes?: SemanticRouterRoute[];
}
export interface SemanticRouterValidationParams {
defaultModel: string | undefined;
embeddingModel: string | undefined;
routerConfig: SemanticRouterConfig | null | undefined;
}
export const getSemanticRouterError = ({
defaultModel,
embeddingModel,
routerConfig,
}: SemanticRouterValidationParams): string | null => {
if (!defaultModel) return "Please select a Default Model";
if (!routerConfig?.routes || routerConfig.routes.length === 0)
return "Please configure at least one route for the auto router";
if (!embeddingModel) return "Please select an Embedding Model";
if (routerConfig.routes.some((route) => !route.name || !route.description || (route.utterances?.length ?? 0) === 0))
return "Please ensure all routes have a target model, description, and at least one utterance";
return null;
};

View file

@ -367,15 +367,18 @@ const EditAutoRouterModal: React.FC<EditAutoRouterModalProps> = ({
</Form.Item>
{/* Embedding Model */}
<Form.Item label="Embedding Model" name="auto_router_embedding_model">
<Form.Item
label="Embedding Model"
name="auto_router_embedding_model"
rules={[{ required: true, message: "Embedding model is required" }]}
>
<AntdSelect
placeholder="Select an embedding model (optional)"
placeholder="Select an embedding model"
onChange={(value) => {
setShowCustomEmbeddingModel(value === "custom");
}}
options={[...modelOptions, { value: "custom", label: "Enter custom model name" }]}
showSearch={true}
allowClear
/>
</Form.Item>
</>

View file

@ -22504,6 +22504,11 @@ export interface components {
* @description connect to a postgres db - needed for generating temporary keys + tracking spend / key
*/
database_url?: string | null;
/**
* Disable Auto Add Proxy Admin To Teams
* @description By default, the user calling /team/new is automatically added to the new team as a team admin. If True, proxy admins are no longer auto-added; members explicitly listed in members_with_roles are unaffected. Default is False.
*/
disable_auto_add_proxy_admin_to_teams?: boolean | null;
/**
* Disable Budget Reservation
* @description If True, disables the optimistic per-request budget reservation introduced in v1.84.0. WARNING: This weakens hard budget enforcement. Without the reservation, a burst of concurrent requests from a single key can each pass the read-time spend check before any of them is charged, allowing a configured budget to be exceeded under high concurrency. Budgets are still evaluated on every request at read time, so an already-exhausted budget is still rejected. Enable only if your deployment is experiencing phantom BudgetExceededError responses caused by leaked reservations (see GitHub issue #27639). A proxy-level WARNING is logged on every request while this flag is active as a reminder that hard enforcement is relaxed.
@ -32245,6 +32250,8 @@ export interface components {
}[] | null;
/** Cooldown Time */
cooldown_time?: number | null;
/** Enable Tag Filtering */
enable_tag_filtering?: boolean | null;
/** Fallbacks */
fallbacks?: {
[key: string]: unknown;

6
uv.lock generated
View file

@ -10,7 +10,7 @@ resolution-markers = [
]
[options]
exclude-newer = "2026-07-13T19:39:19.414377Z"
exclude-newer = "2026-07-13T21:03:01.672393Z"
exclude-newer-span = "P3D"
[manifest]
@ -4122,12 +4122,12 @@ proxy-dev = [
[[package]]
name = "litellm-enterprise"
version = "0.1.50"
version = "0.1.51"
source = { editable = "enterprise" }
[[package]]
name = "litellm-proxy-extras"
version = "0.4.77"
version = "0.4.78"
source = { editable = "litellm-proxy-extras" }
[[package]]