fix(live): enforce delegated model grants and managed budgets

This commit is contained in:
jibanez-staticduo 2026-09-17 11:10:57 +02:00
parent fbda29e1b7
commit 153ece2769
No known key found for this signature in database
5 changed files with 1005 additions and 80 deletions

View file

@ -114,8 +114,8 @@ Session controls require a session known to the proxy and owned by the authentic
Live duration uses cumulative `usage.seconds`; legacy Codex milliseconds remain supported. WebRTC initialization has a 15-second minimum credited against running duration, not added to it. Nested terminal Responses usage is charged separately using its backend model and deduplicated by response ID. A failed observation connection cannot establish complete usage. Managed delegation also depends on receiving its backend usage events; the upstream sideband does not replay events emitted before attachment
Managed Responses delegation authorizes its backend model as well as the voice model. Use client delegation when a key's per-model budgets or token/request limits require admission checks for each backend invocation. Restricted-model WebRTC keys must explicitly exclude `session.update` from frontend client events when using managed delegation, because that data channel connects directly to OpenAI and could otherwise change the backend model outside the proxy's checks
Managed Responses delegation authorizes its backend model as well as the voice model. Use client delegation when budgets or request/token limits apply to the key, user, team, project, organization, team member, end user, or a model access group, since each backend invocation needs its own admission check. Keys scoped to access groups, projects, users, organizations, or teams are treated as model-restricted even if the key's own model list is empty. Managed WebRTC sessions with model restrictions must explicitly exclude `session.update` and wildcard events from frontend client events. That data channel connects directly to OpenAI and could otherwise change the backend model outside the proxy's checks. Client delegation does not need this restriction: the delegation type cannot change after startup or on a fork. Sparse sideband updates may omit the backend model to retain its current value
For both HTTP and WebSocket forks, restricted-model keys must explicitly set `session.delegation` to `{ "type": "client" }` or `{ "type": "responses", "responses": { "model": "authorized-backend" } }`. Empty overrides cannot safely authorize an inherited backend: the session handle records startup configuration, while later updates may have changed the upstream model. Upstream rules still determine which delegation overrides a source session permits
For both HTTP and WebSocket forks of managed sessions, restricted-model keys must explicitly provide an authorized `session.delegation.responses.model`. Empty overrides cannot safely authorize an inherited managed backend: the session handle records startup configuration, while later updates may have changed the upstream model. Client-delegation forks can use empty overrides because the delegation type is immutable
See the official [Live overview](https://developers.openai.com/api/docs/guides/live), [Live API reference](https://developers.openai.com/api/reference/resources/live), [session management](https://developers.openai.com/api/docs/guides/live-conversations), [WebRTC guide](https://developers.openai.com/api/docs/guides/voice-webrtc?api=live), [WebSocket guide](https://developers.openai.com/api/docs/guides/voice-websockets?api=live), [server controls](https://developers.openai.com/api/docs/guides/voice-server-controls?api=live) and [SIP guide](https://developers.openai.com/api/docs/guides/voice-sip?api=live) for the upstream contract. The voice guides also contain Realtime tabs with different routes and formats

View file

@ -2149,6 +2149,7 @@ async def get_team_membership(
user_api_key_cache: UserApiKeyCache,
parent_otel_span: Span | None = None,
proxy_logging_obj: ProxyLogging | None = None,
raise_on_error: bool = False,
) -> Optional["LiteLLM_TeamMembership"]:
"""
Returns team membership object if user is member of team.
@ -2202,6 +2203,8 @@ async def get_team_membership(
user_id,
team_id,
)
if raise_on_error:
raise
return None
@ -4122,6 +4125,8 @@ async def _team_member_granted_models(
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging,
*,
strict_grant_lookup: bool = False,
) -> Sequence[str]:
"""The member's own ``allowed_models`` scope; empty when the member is not narrowed below the team."""
if team_object is None or valid_token.user_id is None:
@ -4133,6 +4138,7 @@ async def _team_member_granted_models(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
raise_on_error=strict_grant_lookup,
)
return () if team_membership is None else _member_allowed_models(team_membership)
@ -4143,6 +4149,8 @@ async def _org_granted_models(
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging,
*,
strict_grant_lookup: bool = False,
) -> Sequence[str]:
"""The org allowlist reached through the key, or through its team when the key names no org."""
org_id: Final = valid_token.org_id or (team_object.organization_id if team_object is not None else None)
@ -4158,6 +4166,8 @@ async def _org_granted_models(
)
except Exception as e: # noqa: BLE001 # fail-safe: attribution degrades to "no org grant", it must never break auth
verbose_proxy_logger.debug("access group attribution: org lookup failed: %s", e)
if strict_grant_lookup:
raise
return ()
return org_object.models if org_object is not None else ()
@ -4169,6 +4179,8 @@ async def _granted_model_lists(
prisma_client: PrismaClient,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging,
*,
strict_grant_lookup: bool = False,
) -> tuple[Sequence[str], ...]:
"""One model allowlist per level that participates in authorizing the request."""
return (
@ -4180,6 +4192,7 @@ async def _granted_model_lists(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
strict_grant_lookup=strict_grant_lookup,
),
project_object.models if project_object is not None else (),
await _org_granted_models(
@ -4188,6 +4201,7 @@ async def _granted_model_lists(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
strict_grant_lookup=strict_grant_lookup,
),
)
@ -4274,6 +4288,8 @@ async def collect_matched_model_access_groups(
prisma_client: PrismaClient | None,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging,
*,
strict_grant_lookup: bool = False,
) -> tuple[str, ...]:
"""
The budgeted model access groups that authorized this request, sorted and deduplicated.
@ -4289,7 +4305,9 @@ async def collect_matched_model_access_groups(
The whole walk is gated on the budget registry, because collecting every match costs a full scan
of each allowlist where the plain access check stops at the first hit. An empty registry means no
group carries a budget, so there is nothing to attribute and no work worth doing.
group carries a budget, so there is nothing to attribute and no work worth doing. The strict
lookup mode is reserved for enforcement paths that must not treat an unavailable inherited grant
as absent; the default remains fail-safe attribution for ordinary request telemetry.
"""
if model is None or valid_token is None or llm_router is None or prisma_client is None:
return ()
@ -4319,6 +4337,7 @@ async def collect_matched_model_access_groups(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
strict_grant_lookup=strict_grant_lookup,
)
for granted_model in granted_models
)

View file

@ -19,12 +19,32 @@ if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.realtime_streaming import RealTimeStreaming
from litellm.llms.chatgpt.live import LiveDeployment, LiveOperation, LiveTransport, live_session_path
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.models.budget import LiteLLM_BudgetTable
from litellm.models.team import LiteLLM_TeamTable
from litellm.models.team_membership import LiteLLM_TeamMembership
from litellm.proxy._types import (
LiteLLM_ProjectTableCachedObj,
LiteLLM_TeamTableCachedObj,
LitellmUserRoles,
UserAPIKeyAuth,
)
from litellm.proxy.auth.auth_checks import (
can_key_call_resolved_model, # pyright: ignore[reportUnknownVariableType] # legacy authorization accepts untyped deployment lists
can_org_access_model,
can_user_call_model,
collect_matched_model_access_groups,
get_org_object,
get_team_object,
get_user_object,
)
from litellm.proxy.auth.user_api_key_auth import get_websocket_api_key, user_api_key_auth
from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper
from litellm.proxy.common_utils.user_api_key_cache import (
NO_TEAM_MEMBERSHIP_SENTINEL,
model_access_group_cache_key,
team_membership_reservation_cache_key,
)
from litellm.proxy.hooks.parallel_request_limiter import (
_PROXY_MaxParallelRequestsHandler, # pyright: ignore[reportPrivateUsage] # limiter class is the existing hook identity
)
@ -38,6 +58,11 @@ from litellm.proxy.realtime_endpoints.call_supervision import CALL_SUPERVISORS,
from litellm.proxy.spend_tracking.budget_reservation import (
release_or_invalidate_budget_reservation, # pyright: ignore[reportUnknownVariableType] # budget helper accepts legacy reservation dicts
)
from litellm.repositories.budget_repository import BudgetRepository
from litellm.repositories.project_repository import ProjectRepository
from litellm.repositories.table_repositories import ModelAccessGroupBudgetRepository, TeamMembershipRepository
from litellm.repositories.team_repository import TeamRepository
from litellm.types.proxy.model_access_group_budget import ModelAccessGroupBudget
_routes: Final = APIRouter()
_JSON: Final = TypeAdapter[JsonValue](JsonValue)
@ -45,17 +70,45 @@ _EMPTY: Final[Mapping[str, JsonValue]] = MappingProxyType({})
_MAPPING: Final = TypeAdapter(Mapping[str, object])
_OBJECT: Final = TypeAdapter(Mapping[str, JsonValue])
_DEPLOYMENT: Final = TypeAdapter(LiveDeployment)
_MODEL_NAMES: Final = TypeAdapter(tuple[str, ...])
_OBJECT_VALUE: Final = TypeAdapter(object)
_SEQUENCE: Final = TypeAdapter(tuple[object, ...])
_PREFIX: Final = "live_litellm_"
def _json_value(value: object) -> JsonValue:
if isinstance(value, Mapping):
entries: Final = _MAPPING.validate_python(value)
return {key: _json_value(item) for key, item in entries.items()} # mutable-ok: JSON wire objects require dicts
if isinstance(value, (tuple, list)):
items: Final = TypeAdapter(tuple[object, ...]).validate_python(value)
return [_json_value(item) for item in items] # mutable-ok: JSON wire arrays require lists
return _JSON.validate_python(value)
root: Final[list[JsonValue]] = [None] # mutable-ok: iterative conversion fills JSON output containers
pending: Final[ # mutable-ok: work stack carries mutable JSON output containers
list[tuple[object, dict[str, JsonValue] | list[JsonValue], str | int, int]]
] = [(value, root, 0, 0)] # mutable-ok: traversal adds pending nodes
while pending:
source, parent, key, depth = pending.pop()
if depth > 256:
raise ValueError("Live JSON nesting exceeds the supported depth")
converted: JsonValue # rebind-ok: each visited input produces a new JSON value
if isinstance(source, Mapping):
entries: Mapping[str, object] = _MAPPING.validate_python(
source
) # rebind-ok: entries belong to the current node
converted = {name: None for name in entries} # mutable-ok: JSON wire objects require dicts
pending.extend((item, converted, name, depth + 1) for name, item in entries.items())
elif isinstance(source, (tuple, list)):
items: tuple[object, ...] = TypeAdapter(tuple[object, ...]).validate_python(
source
) # rebind-ok: items belong to the current node
array: list[JsonValue] = [None] * len(items) # mutable-ok: JSON output; # rebind-ok: per-node buffer
pending.extend((item, array, index, depth + 1) for index, item in enumerate(items))
converted = array
else:
converted = _JSON.validate_python(source)
match parent, key:
case dict(), str():
parent[key] = converted
case list(), int():
parent[key] = converted
case _:
raise TypeError("Invalid Live JSON conversion target")
return root[0]
def _object(value: object) -> Mapping[str, JsonValue]:
@ -182,6 +235,24 @@ def _session_model(body: Mapping[str, JsonValue], fallback: str | None = None) -
return model
async def _live_organization_id(auth: UserAPIKeyAuth) -> str | None:
if auth.org_id is not None or auth.team_id is None:
return auth.org_id
from litellm.proxy import proxy_server as server
try:
team_object: Final = await get_team_object(
team_id=auth.team_id,
prisma_client=server.prisma_client,
user_api_key_cache=server.user_api_key_cache,
parent_otel_span=auth.parent_otel_span,
proxy_logging_obj=server.proxy_logging_obj,
)
except Exception as exc:
raise HTTPException(503, "Could not verify Live team organization model access") from exc
return team_object.organization_id
async def _authorize(model: str, auth: UserAPIKeyAuth) -> None:
from litellm.proxy import proxy_server as server
@ -194,6 +265,38 @@ async def _authorize(model: str, auth: UserAPIKeyAuth) -> None:
llm_router=server.llm_router,
)
if auth.user_id is not None and auth.team_id is None and auth.user_role != LitellmUserRoles.PROXY_ADMIN:
try:
user_object: Final = await get_user_object(
user_id=auth.user_id,
prisma_client=server.prisma_client,
user_api_key_cache=server.user_api_key_cache,
user_id_upsert=False,
parent_otel_span=auth.parent_otel_span,
proxy_logging_obj=server.proxy_logging_obj,
)
except Exception as exc:
raise HTTPException(503, "Could not verify Live user model access") from exc
if user_object is None:
raise HTTPException(503, "Could not verify Live user model access")
await can_user_call_model(model=model, llm_router=server.llm_router, user_object=user_object)
organization_id: Final = await _live_organization_id(auth)
if organization_id is not None:
try:
org_object: Final = await get_org_object(
org_id=organization_id,
prisma_client=server.prisma_client,
user_api_key_cache=server.user_api_key_cache,
parent_otel_span=auth.parent_otel_span,
proxy_logging_obj=server.proxy_logging_obj,
)
except Exception as exc:
raise HTTPException(503, "Could not verify Live organization model access") from exc
if org_object is None:
raise HTTPException(503, "Could not verify Live organization model access")
can_org_access_model(model=model, org_object=org_object, llm_router=server.llm_router)
async def _deployment(model: str, processed: Mapping[str, object]) -> LiveDeployment:
from litellm.proxy import proxy_server as server
@ -379,54 +482,362 @@ def _session_policy(body: Mapping[str, JsonValue], source: LiveHandle | None) ->
def _managed_constraints(auth: UserAPIKeyAuth) -> bool:
values: Final = _object(
auth.model_dump(
include=MappingProxyType(
{
name: True
for name in (
"model_max_budget",
"user_model_max_budget",
"end_user_model_max_budget",
"rpm_limit_per_model",
"tpm_limit_per_model",
"rpm_limit",
"tpm_limit",
"team_rpm_limit",
"team_tpm_limit",
"user_rpm_limit",
"user_tpm_limit",
"team_metadata",
"metadata",
"organization_metadata",
"project_metadata",
)
}
)
if any(
value is not None
for value in (
auth.rpm_limit,
auth.tpm_limit,
auth.team_rpm_limit,
auth.team_tpm_limit,
auth.user_rpm_limit,
auth.user_tpm_limit,
auth.organization_rpm_limit,
auth.organization_tpm_limit,
auth.team_member_rpm_limit,
auth.team_member_tpm_limit,
auth.end_user_rpm_limit,
auth.end_user_tpm_limit,
auth.max_budget,
auth.team_max_budget,
auth.user_max_budget,
auth.end_user_max_budget,
auth.organization_max_budget,
)
):
return True
direct_maps: Final[tuple[object, ...]] = (
_OBJECT_VALUE.validate_python(getattr(auth, "model_max_budget", None)),
_OBJECT_VALUE.validate_python(getattr(auth, "user_model_max_budget", None)),
_OBJECT_VALUE.validate_python(getattr(auth, "end_user_model_max_budget", None)),
_OBJECT_VALUE.validate_python(getattr(auth, "rpm_limit_per_model", None)),
_OBJECT_VALUE.validate_python(getattr(auth, "tpm_limit_per_model", None)),
_OBJECT_VALUE.validate_python(getattr(auth, "budget_limits", None)),
)
if any(_nonempty_limit_value(value) for value in direct_maps):
return True
def constrained(value: JsonValue | Mapping[str, JsonValue]) -> bool:
if not isinstance(value, Mapping):
return False
return any(
bool(item)
if any(marker in key for marker in ("rpm_limit", "tpm_limit", "model_max_budget"))
else constrained(item)
for key, item in value.items()
pending: Final[list[object]] = [ # mutable-ok: explicit metadata traversal stack
value
for value in (
_OBJECT_VALUE.validate_python(getattr(auth, "team_metadata", None)),
_OBJECT_VALUE.validate_python(getattr(auth, "metadata", None)),
_OBJECT_VALUE.validate_python(getattr(auth, "organization_metadata", None)),
_OBJECT_VALUE.validate_python(getattr(auth, "project_metadata", None)),
)
if isinstance(value, (Mapping, list, tuple))
]
visited: Final[set[int]] = set() # mutable-ok: cycle guard for hook-provided metadata
while pending:
current: object = pending.pop() # rebind-ok: advance the explicit metadata traversal stack
if id(current) in visited:
continue
visited.add(id(current))
if len(visited) > 4096:
return True
if isinstance(current, Mapping):
entries: Mapping[str, object] = _MAPPING.validate_python(
current
) # rebind-ok: entries belong to the current metadata node
for key, item in entries.items():
if key in (
"rpm_limit",
"tpm_limit",
"max_budget",
"model_rpm_limit",
"model_tpm_limit",
"model_itpm_limit",
"model_otpm_limit",
"model_max_budget",
"budget_limits",
) and _nonempty_limit_value(item):
return True
if isinstance(item, (Mapping, list, tuple)):
nested: object = _OBJECT_VALUE.validate_python(item)
pending.append(nested)
elif isinstance(current, (list, tuple)):
sequence: object = _OBJECT_VALUE.validate_python(current)
pending.extend(_SEQUENCE.validate_python(sequence))
return False
return constrained(values)
def _nonempty_limit_value(value: object) -> bool:
if value is None:
return False
if isinstance(value, Mapping):
mapping: Final[object] = _OBJECT_VALUE.validate_python(value)
return bool(_MAPPING.validate_python(mapping))
if isinstance(value, (list, tuple)):
sequence: Final[object] = _OBJECT_VALUE.validate_python(value)
return bool(_SEQUENCE.validate_python(sequence))
return True
def _restricted_model_list(models: object) -> bool:
values: Final = _MODEL_NAMES.validate_python(models or ())
return bool(values) and "*" not in values and "all-proxy-models" not in values
def _restricted_models(auth: UserAPIKeyAuth) -> bool:
key_models: Final = TypeAdapter(tuple[str, ...]).validate_python(getattr(auth, "models", ()) or ())
team_models: Final = TypeAdapter(tuple[str, ...]).validate_python(getattr(auth, "team_models", ()) or ())
return any(
models and "*" not in models and "all-proxy-models" not in models for models in (key_models, team_models)
if _restricted_model_list(getattr(auth, "models", None)) or _restricted_model_list(
getattr(auth, "team_models", None)
):
return True
if auth.user_id is not None and auth.user_role != LitellmUserRoles.PROXY_ADMIN:
return True
return bool(
auth.access_group_ids or auth.matched_model_access_groups or auth.project_id or auth.org_id or auth.team_id
)
def _live_budget_configured(value: object, zero_is_limit: bool) -> bool:
if value is None:
return False
value_object: Final[object] = value
value_mapping: Final[Mapping[str, object] | None] = (
_MAPPING.validate_python(value) if isinstance(value, Mapping) else None
)
max_budget: Final[object] = _OBJECT_VALUE.validate_python(
value_mapping.get("max_budget") if value_mapping is not None else getattr(value_object, "max_budget", None)
)
if max_budget is not None:
if zero_is_limit or (isinstance(max_budget, (int, float)) and max_budget > 0):
return True
if not isinstance(max_budget, (int, float)):
return True
fields: Final[tuple[object, ...]] = tuple(
_OBJECT_VALUE.validate_python(
value_mapping.get(field) if value_mapping is not None else getattr(value_object, field, None)
)
for field in ("rpm_limit", "tpm_limit", "model_max_budget", "max_parallel_requests")
)
return any(_nonempty_limit_value(field) for field in fields)
def _live_budget_scope_present(auth: UserAPIKeyAuth, model: str | None, llm_router: object | None) -> bool:
explicit_model_scope: Final = _restricted_model_list(getattr(auth, "models", None)) or _restricted_model_list(
getattr(auth, "team_models", None)
)
model_group_lookup_scope: Final = bool(
model is not None and llm_router is not None and (explicit_model_scope or auth.org_id is not None)
)
return bool(
auth.team_id
or auth.project_id
or auth.access_group_ids
or auth.matched_model_access_groups
or model_group_lookup_scope
)
async def _live_team_membership(auth: UserAPIKeyAuth) -> object | None:
from litellm.proxy import proxy_server as server
if auth.team_id is None or auth.user_id is None:
return None
membership_key: Final = team_membership_reservation_cache_key(user_id=auth.user_id, team_id=auth.team_id)
membership_cached_raw: Final[object] = _OBJECT_VALUE.validate_python(
await server.user_api_key_cache.async_get_cache(key=membership_key)
)
membership_cached: Final = (
CacheCodec.deserialize(membership_cached_raw, model_type=LiteLLM_TeamMembership)
if membership_cached_raw is not None and membership_cached_raw != NO_TEAM_MEMBERSHIP_SENTINEL
else None
)
if membership_cached is not None or membership_cached_raw == NO_TEAM_MEMBERSHIP_SENTINEL:
return membership_cached
return await TeamMembershipRepository(server.prisma_client).table.find_unique(
where={ # mutable-ok: Prisma serializes query filters from concrete dictionaries
"user_id_team_id": { # mutable-ok: Prisma serializes nested filters from concrete dictionaries
"user_id": auth.user_id,
"team_id": auth.team_id,
}
},
include={"litellm_budget_table": True}, # mutable-ok: Prisma serializes concrete include dictionaries
)
async def _live_team(auth: UserAPIKeyAuth) -> LiteLLM_TeamTable | None:
from litellm.proxy import proxy_server as server
if auth.team_id is None:
return None
team_from_cache: Final = await server.user_api_key_cache.async_get_cache(
key=f"team_id:{auth.team_id}", model_type=LiteLLM_TeamTableCachedObj
)
if team_from_cache is not None:
return team_from_cache
return await TeamRepository(server.prisma_client).find_by_id(auth.team_id, id_field="team_id")
def _live_team_budget_configured(auth: UserAPIKeyAuth, team: LiteLLM_TeamTable | None) -> bool:
if team is None:
return False
team_budget_limits: Final = getattr(team, "budget_limits", None)
if _nonempty_limit_value(team_budget_limits):
return True
if any(getattr(team, field, None) is not None for field in ("rpm_limit", "tpm_limit", "max_budget")):
return True
if _nonempty_limit_value(getattr(team, "model_max_budget", None)):
return True
team_metadata_value: Final = getattr(team, "metadata", None)
team_metadata: Final = (
team_metadata_value if team_metadata_value is not None else getattr(auth, "team_metadata", None)
)
return _managed_constraints(
auth.model_copy(update=MappingProxyType({"team_metadata": team_metadata, "budget_limits": team_budget_limits}))
)
async def _live_default_budget(auth: UserAPIKeyAuth, team: LiteLLM_TeamTable | None) -> LiteLLM_BudgetTable | None:
metadata_source: Final = (
getattr(team, "metadata", None) if team is not None else getattr(auth, "team_metadata", None)
)
default_id: Final = _MAPPING.validate_python(metadata_source or _EMPTY).get("team_member_budget_id")
if not isinstance(default_id, str) or auth.team_id is None or auth.user_id is None:
return None
from litellm.proxy import proxy_server as server
default_cached: Final = await server.user_api_key_cache.async_get_cache(
key=f"team_member_default_budget:{default_id}", model_type=LiteLLM_BudgetTable
)
if default_cached is not None:
return default_cached
return await BudgetRepository(server.prisma_client).find_by_id(default_id, id_field="budget_id")
async def _live_project(auth: UserAPIKeyAuth) -> LiteLLM_ProjectTableCachedObj | None:
from litellm.proxy import proxy_server as server
if auth.project_id is None:
return None
project_from_cache: Final = await server.user_api_key_cache.async_get_cache(
key=f"project_id:{auth.project_id}", model_type=LiteLLM_ProjectTableCachedObj
)
if project_from_cache is not None:
return project_from_cache
project_row: Final = await ProjectRepository(server.prisma_client).table.find_unique(
where={"project_id": auth.project_id}, # mutable-ok: Prisma serializes concrete query dictionaries
include={"litellm_budget_table": True}, # mutable-ok: Prisma serializes concrete include dictionaries
)
if project_row is None:
return None
return LiteLLM_ProjectTableCachedObj.model_validate(project_row.model_dump())
async def _live_project_budget_configured(auth: UserAPIKeyAuth, project: LiteLLM_ProjectTableCachedObj | None) -> bool:
if project is None:
return False
from litellm.proxy import proxy_server as server
project_budget: Final = getattr(project, "litellm_budget_table", None)
project_budget_id: Final = getattr(project, "budget_id", None)
project_budget_from_db: Final = (
await BudgetRepository(server.prisma_client).find_by_id(project_budget_id, id_field="budget_id")
if project_budget is None and isinstance(project_budget_id, str)
else None
)
if _live_budget_configured(project_budget or project_budget_from_db, zero_is_limit=True):
return True
if _nonempty_limit_value(getattr(project, "model_rpm_limit", None)) or _nonempty_limit_value(
getattr(project, "model_tpm_limit", None)
):
return True
project_metadata_value: Final = getattr(project, "metadata", None)
project_metadata: Final = (
project_metadata_value if project_metadata_value is not None else getattr(auth, "project_metadata", None)
)
return _managed_constraints(auth.model_copy(update=MappingProxyType({"project_metadata": project_metadata})))
async def _live_model_group_budget_configured(
auth: UserAPIKeyAuth,
model: str | None,
team: LiteLLM_TeamTable | None,
project: LiteLLM_ProjectTableCachedObj | None,
) -> bool:
if model is None:
return False
from litellm.proxy import proxy_server as server
matched_groups: Final = await collect_matched_model_access_groups(
model=model,
valid_token=auth,
team_object=team,
project_object=project,
llm_router=server.llm_router,
prisma_client=server.prisma_client,
user_api_key_cache=server.user_api_key_cache,
proxy_logging_obj=server.proxy_logging_obj,
strict_grant_lookup=True,
)
if not matched_groups:
return False
cached_values: Final = await asyncio.gather(
*(
server.user_api_key_cache.async_get_cache(
key=model_access_group_cache_key(group), model_type=ModelAccessGroupBudget
)
for group in matched_groups
)
)
cached_groups: Final = tuple(zip(matched_groups, cached_values))
uncached_groups: Final = tuple(group for group, budget in cached_groups if budget is None)
named_group_rows: Final = (
await ModelAccessGroupBudgetRepository(server.prisma_client).table.find_many(
where={ # mutable-ok: Prisma serializes query filters from concrete dictionaries
"access_group_name": { # mutable-ok: Prisma serializes nested filters from concrete dictionaries
"in": uncached_groups,
}
},
include={"litellm_budget_table": True}, # mutable-ok: Prisma serializes concrete include dictionaries
)
if uncached_groups
else ()
)
return any(
_live_budget_configured(
budget
if budget is not None
else next(
(
getattr(row, "litellm_budget_table", None)
for row in named_group_rows
if getattr(row, "access_group_name", None) == group
),
None,
),
zero_is_limit=False,
)
for group, budget in cached_groups
)
async def _managed_member_budget(auth: UserAPIKeyAuth, model: str | None = None) -> bool:
from litellm.proxy import proxy_server as server
try:
if not _live_budget_scope_present(auth, model, server.llm_router):
return False
if server.prisma_client is None:
raise HTTPException(503, "Could not verify Live managed budgets")
membership: Final = await _live_team_membership(auth)
if _live_budget_configured(getattr(membership, "litellm_budget_table", None), zero_is_limit=True):
return True
team: Final = await _live_team(auth)
if _live_team_budget_configured(auth, team):
return True
default: Final = await _live_default_budget(auth, team)
if _live_budget_configured(default, zero_is_limit=False):
return True
project: Final = await _live_project(auth)
if await _live_project_budget_configured(auth, project):
return True
return await _live_model_group_budget_configured(auth, model, team, project)
except Exception as exc:
raise HTTPException(503, "Could not verify Live managed budgets") from exc
async def _authorize_delegation(body: Mapping[str, JsonValue], auth: UserAPIKeyAuth) -> None:
session: Final = body.get("session")
if not isinstance(session, Mapping):
@ -437,26 +848,32 @@ async def _authorize_delegation(body: Mapping[str, JsonValue], auth: UserAPIKeyA
responses: Final = delegation.get("responses")
if delegation.get("type") != "responses" and not isinstance(responses, dict):
return
if _managed_constraints(auth):
model: Final = responses.get("model") if isinstance(responses, dict) else None
if _managed_constraints(auth) or await _managed_member_budget(
auth,
model if isinstance(model, str) else None,
):
raise HTTPException(
400,
"Managed Live delegation cannot enforce configured backend model budgets or rate limits; use client delegation",
"Managed Live delegation cannot enforce configured budgets or rate limits; use client delegation",
)
if not isinstance(responses, dict):
if _restricted_models(auth):
if body.get("type") != "session.update" and _restricted_models(auth):
raise HTTPException(400, "Restricted keys require an explicit authorized delegation.responses.model")
return
model: Final = responses.get("model")
if isinstance(model, str):
await _authorize(model, auth)
elif body.get("type") != "session.update" and _restricted_models(auth):
raise HTTPException(400, "Restricted keys require an explicit authorized delegation.responses.model")
transport: Final = body.get("transport")
if isinstance(transport, dict) and transport.get("type") == "webrtc" and _restricted_models(auth):
client: Final = session.get("client")
channel: Final = client.get("data_channel") if isinstance(client, dict) else None
events: Final = channel.get("allowed_client_events") if isinstance(channel, dict) else None
if not isinstance(events, list) or "session.update" in events:
if not isinstance(events, list) or any(
not isinstance(event, str) or event.strip() == "session.update" or "*" in event for event in events
):
raise HTTPException(
400, "Restricted keys must explicitly exclude session.update from WebRTC allowed_client_events"
)
@ -465,7 +882,11 @@ async def _authorize_delegation(body: Mapping[str, JsonValue], auth: UserAPIKeyA
async def _authorize_fork_policy(
body: Mapping[str, JsonValue], source: LiveHandle | None, auth: UserAPIKeyAuth
) -> None:
if source is not None and (_restricted_models(auth) or _managed_constraints(auth)):
if (
source is not None
and source.policy.get("delegation")
and (_restricted_models(auth) or _managed_constraints(auth))
):
# Handles contain startup policy; later sideband or WebRTC updates can change the backend model.
session: Final = _policy_object(body.get("session", _EMPTY))
delegation: Final = _policy_object(session.get("delegation") or _EMPTY)

View file

@ -299,6 +299,7 @@ class TestProcessResponse:
)
@pytest.mark.usefixtures("local_model_cost_map")
class TestProcessEmbedContentResponseUsage:
"""Gemini Embedding 2 embedContent usageMetadata must drive spend.

View file

@ -8,8 +8,15 @@ import httpx
import pytest
from fastapi import FastAPI, HTTPException, Request
from fastapi.testclient import TestClient
from prisma.builder import QueryBuilder
from litellm.llms.chatgpt.live import LiveDeployment
from litellm.models.budget import LiteLLM_BudgetTable
from litellm.models.organization import LiteLLM_OrganizationTable
from litellm.models.project import LiteLLM_ProjectTable
from litellm.models.team import LiteLLM_TeamTable
from litellm.models.team_membership import LiteLLM_TeamMembership
from litellm.models.user import LiteLLM_UserTable
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.realtime_endpoints import live
@ -78,6 +85,30 @@ def test_handle_serializes_mappingproxy_without_losing_pinned_deployment():
assert live._pinned(live.decode_session(live.encode_session(original), original.owner)) == deployment
def test_json_value_iteratively_serializes_nested_mappingproxy_tuple_and_shared_subtree():
shared = MappingProxyType({"deep": (1, 2)})
value = MappingProxyType({"left": shared, "right": (shared,)})
assert live._json_value(value) == {"left": {"deep": [1, 2]}, "right": [{"deep": [1, 2]}]}
@pytest.mark.parametrize("value", [{1: "invalid"}, {"invalid": object()}])
def test_json_value_rejects_non_json_objects_and_keys(value):
with pytest.raises(ValueError, match="validation error"):
live._json_value(value)
def test_json_value_rejects_cycles_and_excessive_depth():
cycle = {}
cycle["self"] = cycle
with pytest.raises(ValueError, match="depth"):
live._json_value(cycle)
nested = json.loads('{"value":' * 257 + "null" + "}" * 257)
with pytest.raises(ValueError, match="depth"):
live._json_value(nested)
def test_only_protocol_session_ids_are_rewritten_and_application_values_survive():
event = {
"type": "session.started",
@ -322,9 +353,19 @@ async def test_websocket_delegation_model_update_is_authorized_before_forwarding
"limits",
[
{"rpm_limit": 1},
{"rpm_limit": 0},
{"model_max_budget": {"backend": 1}},
{"rpm_limit_per_model": {"backend": 0}},
{"tpm_limit_per_model": {"backend": 0}},
{"team_tpm_limit": 10},
{"organization_rpm_limit": 0},
{"organization_tpm_limit": 0},
{"team_member_rpm_limit": 0},
{"team_member_tpm_limit": 0},
{"end_user_rpm_limit": 0},
{"end_user_tpm_limit": 0},
{"team_metadata": {"model_rpm_limit": {"backend": 1}}},
{"metadata": {"scopes": [{"nested": {"model_tpm_limit": {"backend": 0}}}]}},
],
)
async def test_managed_delegation_fails_closed_for_unenforceable_constraints(limits):
@ -336,13 +377,173 @@ async def test_managed_delegation_fails_closed_for_unenforceable_constraints(lim
@pytest.mark.asyncio
async def test_restricted_webrtc_cannot_change_managed_model_outside_proxy(monkeypatch):
@pytest.mark.parametrize("delegation", [{"type": "responses"}, {"type": "responses", "responses": {}}])
async def test_restricted_session_update_can_retain_backend_delegation_model(
delegation,
):
body = {"type": "session.update", "session": {"delegation": delegation}}
result = await live._authorize_delegation(body, UserAPIKeyAuth(api_key="owner", models=["voice", "backend"]))
assert result is None
@pytest.mark.asyncio
@pytest.mark.parametrize("session", [{}, {"delegation": None}, {"delegation": {"type": "client"}}])
async def test_constrained_webrtc_client_delegation_allows_frontend_updates(
session,
):
result = await live._authorize_delegation(
{"session": session, "transport": {"type": "webrtc"}},
UserAPIKeyAuth(api_key="owner", models=["voice"], rpm_limit=10),
)
assert result is None
@pytest.mark.parametrize(
"limits",
[
{"max_budget": 1},
{"team_max_budget": 1},
{"user_max_budget": 1},
{"end_user_max_budget": 1},
{"organization_max_budget": 1},
{"budget_limits": [{"budget_duration": "1d", "max_budget": 1}]},
],
)
def test_managed_constraints_detect_scalar_and_window_budgets(limits):
assert live._managed_constraints(UserAPIKeyAuth(api_key="owner", **limits)) is True
@pytest.mark.asyncio
@pytest.mark.parametrize(
"member_limit, default_limit, blocked",
[(0, None, True), (1, None, True), (None, 1, True), (None, 0, False), (None, None, False)],
)
async def test_managed_delegation_checks_authoritative_member_and_default_budget(
monkeypatch, member_limit, default_limit, blocked
):
auth = UserAPIKeyAuth(
api_key="owner", team_id="team", user_id="user", team_metadata={"team_member_budget_id": "budget"}
)
from litellm.proxy import proxy_server
membership = AsyncMock(
return_value=SimpleNamespace(litellm_budget_table=LiteLLM_BudgetTable(max_budget=member_limit))
)
db = SimpleNamespace(
litellm_teammembership=SimpleNamespace(find_unique=membership),
litellm_teamtable=SimpleNamespace(find_unique=AsyncMock(return_value=None)),
litellm_budgettable=SimpleNamespace(
find_unique=AsyncMock(return_value=LiteLLM_BudgetTable(max_budget=default_limit))
),
)
monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db))
monkeypatch.setattr(
proxy_server, "user_api_key_cache", SimpleNamespace(async_get_cache=AsyncMock(return_value=None))
)
monkeypatch.setattr(proxy_server, "llm_router", None)
monkeypatch.setattr(live, "_authorize", AsyncMock())
body = {"session": {"delegation": {"type": "responses", "responses": {"model": "backend"}}}}
if blocked:
with pytest.raises(HTTPException) as rejected:
await live._authorize_delegation(body, auth)
assert rejected.value.status_code == 400
assert "client delegation" in rejected.value.detail
else:
await live._authorize_delegation(body, auth)
assert membership.await_args.kwargs["where"] == {"user_id_team_id": {"user_id": "user", "team_id": "team"}}
@pytest.mark.asyncio
async def test_managed_delegation_rejects_unverifiable_member_budget_but_allows_client(monkeypatch):
from litellm.proxy import proxy_server
db = SimpleNamespace(
litellm_teammembership=SimpleNamespace(find_unique=AsyncMock(side_effect=RuntimeError("Database unavailable")))
)
monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db))
monkeypatch.setattr(
proxy_server, "user_api_key_cache", SimpleNamespace(async_get_cache=AsyncMock(return_value=None))
)
auth = UserAPIKeyAuth(api_key="owner", team_id="team", user_id="user")
body = {"session": {"delegation": {"type": "responses", "responses": {"model": "backend"}}}}
with pytest.raises(HTTPException) as rejected:
await live._authorize_delegation(body, auth)
assert rejected.value.status_code == 503
await live._authorize_delegation({"session": {"delegation": {"type": "client"}}}, auth)
def test_managed_constraints_fails_closed_after_metadata_node_limit():
auth = UserAPIKeyAuth(api_key="owner")
auth.metadata = {"items": [{} for _ in range(4097)]}
assert live._managed_constraints(auth) is True
@pytest.mark.parametrize(
"limits",
[
{"model_max_budget": {}},
{"team_metadata": {"model_rpm_limit": {}}},
{"metadata": {"nested": [{"model_max_budget": {}}]}},
],
)
def test_empty_model_limit_maps_do_not_mark_delegation_as_managed(limits):
assert live._managed_constraints(UserAPIKeyAuth(api_key="owner", **limits)) is False
def test_managed_constraints_terminates_on_cyclic_metadata_without_a_limit():
metadata = {}
metadata["self"] = metadata
auth = UserAPIKeyAuth(api_key="owner")
auth.metadata = metadata
assert live._managed_constraints(auth) is False
@pytest.mark.asyncio
@pytest.mark.parametrize(
"scope",
[
{"access_group_ids": ["restricted-group"]},
{"project_id": "restricted-project"},
{"org_id": "restricted-org"},
{"team_id": "restricted-team"},
{"team_id": "restricted-team", "user_id": "restricted-member"},
],
)
async def test_restricted_scopes_cannot_delegate_without_an_explicit_backend_model(monkeypatch, scope):
monkeypatch.setattr(live, "_managed_member_budget", AsyncMock(return_value=False))
body = {"session": {"delegation": {"type": "responses", "responses": {}}}}
with pytest.raises(HTTPException) as rejected:
await live._authorize_delegation(body, UserAPIKeyAuth(api_key="owner", **scope))
assert rejected.value.status_code == 400
assert "explicit authorized delegation.responses.model" in rejected.value.detail
@pytest.mark.asyncio
@pytest.mark.parametrize(
"scope",
[
{"models": ["voice", "backend"]},
{"access_group_ids": ["restricted-group"]},
{"matched_model_access_groups": ["restricted-group"]},
{"project_id": "restricted-project"},
{"org_id": "restricted-org"},
{"team_id": "restricted-team"},
{"team_id": "restricted-team", "user_id": "restricted-member"},
],
)
async def test_restricted_webrtc_cannot_change_managed_model_outside_proxy(monkeypatch, scope):
monkeypatch.setattr(live, "_managed_member_budget", AsyncMock(return_value=False))
monkeypatch.setattr(live, "_authorize", AsyncMock())
body = {
"session": {"delegation": {"type": "responses", "responses": {"model": "backend"}}},
"transport": {"type": "webrtc", "sdp": "offer"},
}
auth = UserAPIKeyAuth(api_key="owner", models=["voice", "backend"])
auth = UserAPIKeyAuth(api_key="owner", **scope)
with pytest.raises(HTTPException) as rejected:
await live._authorize_delegation(body, auth)
assert rejected.value.status_code == 400
@ -465,9 +666,7 @@ async def test_inherited_managed_fork_cannot_bypass_new_key_constraints():
@pytest.mark.parametrize("protocol", ["http", "websocket"])
@pytest.mark.parametrize(
"startup_policy", [{}, {"delegation": {"type": "responses", "responses": {"model": "allowed"}}}]
)
@pytest.mark.parametrize("startup_policy", [{"delegation": {"type": "responses", "responses": {"model": "allowed"}}}])
@pytest.mark.parametrize("overrides", [{}, {"delegation": {"responses": {}}}])
def test_restricted_fork_never_trusts_startup_delegation(route_client, protocol, startup_policy, overrides):
from starlette.websockets import WebSocketDisconnect
@ -504,9 +703,9 @@ async def test_explicit_fork_backend_is_authorized_even_when_startup_policy_diff
assert rejected.value.status_code == 403
def test_restricted_fork_can_explicitly_select_client_delegation(route_client):
def test_restricted_client_fork_can_inherit_delegation(route_client):
route_client.auth.models = ["voice"]
body = {"session": {"delegation": {"type": "client"}}}
body = {"session": {}}
token = live.encode_session(handle())
route_client.transport.request.return_value = httpx.Response(
200, json={"session": {"id": "sess_fork"}, "transport": {"type": "webrtc", "sdp": "answer"}}
@ -517,26 +716,13 @@ def test_restricted_fork_can_explicitly_select_client_delegation(route_client):
route_client.transport.request.assert_awaited_once_with("POST", "live/sessions/sess_upstream/fork", body=body)
@pytest.mark.parametrize("protocol", ["http", "websocket"])
@pytest.mark.asyncio
@pytest.mark.parametrize("limits", [{"rpm_limit": 10}, {"tpm_limit": 100}, {"model_max_budget": {"backend": 1}}])
def test_fork_with_new_limits_cannot_trust_old_client_policy(route_client, protocol, limits):
from starlette.websockets import WebSocketDisconnect
# An unrestricted source could have switched to managed delegation after its handle was issued.
for key, value in limits.items():
setattr(route_client.auth, key, value)
token = live.encode_session(handle())
path = f"/v1/live/sessions/{token}/fork"
if protocol == "http":
response = route_client.client.post(path, json={"session": {}})
assert response.status_code == 400
else:
with route_client.client.websocket_connect(path, headers={"Authorization": "Bearer owner"}) as ws:
ws.send_json({"type": "session.start", "session": {}})
with pytest.raises(WebSocketDisconnect) as rejected:
ws.receive_json()
assert rejected.value.code == 1008
route_client.factory.assert_not_called()
async def test_client_fork_can_inherit_immutable_delegation_with_new_limits(
limits,
):
result = await live._authorize_fork_policy({"session": {}}, handle(), UserAPIKeyAuth(api_key="owner", **limits))
assert result is None
@pytest.mark.asyncio
@ -928,3 +1114,301 @@ async def test_live_observer_becomes_ready_without_session_started_event(monkeyp
for supervisor in supervisors:
await supervisor.close()
assert observer.closed
@pytest.fixture
def isolated_live_model_auth(monkeypatch):
from litellm.proxy import proxy_server
monkeypatch.setattr(live, "can_key_call_resolved_model", AsyncMock())
monkeypatch.setattr(proxy_server, "llm_model_list", [])
monkeypatch.setattr(proxy_server, "llm_router", None)
@pytest.mark.asyncio
async def test_authorize_enforces_authoritative_personal_user_models(monkeypatch, isolated_live_model_auth):
user_loader = AsyncMock(return_value=LiteLLM_UserTable(user_id="user-only", models=["voice"]))
monkeypatch.setattr(live, "get_user_object", user_loader)
auth = UserAPIKeyAuth(api_key="owner", models=[], user_id="user-only")
await live._authorize("voice", auth)
with pytest.raises(Exception, match="user can only access"):
await live._authorize("backend", auth)
assert user_loader.await_count == 2
@pytest.mark.asyncio
async def test_authorize_enforces_authoritative_organization_models(monkeypatch, isolated_live_model_auth):
org_loader = AsyncMock(
return_value=LiteLLM_OrganizationTable(
organization_id="org-only",
budget_id="budget",
created_by="admin",
updated_by="admin",
models=["voice"],
)
)
monkeypatch.setattr(live, "get_org_object", org_loader)
auth = UserAPIKeyAuth(api_key="owner", models=[], org_id="org-only")
await live._authorize("voice", auth)
with pytest.raises(Exception, match="org can only access"):
await live._authorize("backend", auth)
assert org_loader.await_count == 2
@pytest.mark.asyncio
@pytest.mark.parametrize(
("scope", "auth_kwargs", "loader_name"),
[
("user", {"user_id": "user-only"}, "get_user_object"),
("organization", {"org_id": "org-only"}, "get_org_object"),
],
)
async def test_authorize_fails_closed_when_principal_grant_lookup_fails(
monkeypatch, isolated_live_model_auth, scope, auth_kwargs, loader_name
):
loader = AsyncMock(side_effect=RuntimeError(f"{scope} lookup unavailable"))
monkeypatch.setattr(live, loader_name, loader)
with pytest.raises(HTTPException) as rejected:
await live._authorize("backend", UserAPIKeyAuth(api_key="owner", **auth_kwargs))
assert rejected.value.status_code == 503
assert "verify Live" in str(rejected.value.detail)
def test_restricted_models_marks_user_scoped_identity_as_restricted():
assert live._restricted_models(UserAPIKeyAuth(api_key="owner", user_id="user-only")) is True
@pytest.mark.asyncio
async def test_sparse_responses_update_without_model_remains_valid_for_user_scoped_identity():
result = await live._authorize_delegation(
{"type": "session.update", "session": {"delegation": {"type": "responses", "responses": {}}}},
UserAPIKeyAuth(api_key="owner", user_id="user-only"),
)
assert result is None
def test_managed_constraints_uses_exact_metadata_keys():
for key in ("max_budget_alert_emails", "model_max_budget_usage"):
assert live._managed_constraints(UserAPIKeyAuth(api_key="owner", metadata={key: {"backend": 1}})) is False
@pytest.mark.asyncio
async def test_managed_budget_reads_authoritative_member_budget(monkeypatch):
from litellm.proxy import proxy_server
membership = LiteLLM_TeamMembership(
user_id="member",
team_id="team",
litellm_budget_table=LiteLLM_BudgetTable(max_budget=1),
)
def find_unique(*, where, include):
QueryBuilder(method="find_unique", arguments={"where": where}).build_query()
return membership
db = SimpleNamespace(
litellm_teammembership=SimpleNamespace(find_unique=AsyncMock(side_effect=find_unique)),
litellm_teamtable=SimpleNamespace(find_unique=AsyncMock(return_value=None)),
)
cache = SimpleNamespace(async_get_cache=AsyncMock(return_value=None), async_set_cache=AsyncMock())
monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db))
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
assert await live._managed_member_budget(UserAPIKeyAuth(api_key="owner", team_id="team", user_id="member")) is True
@pytest.mark.asyncio
async def test_managed_budget_fails_closed_when_membership_repository_is_unreadable(monkeypatch):
from litellm.proxy import proxy_server
failure = RuntimeError("database unavailable")
db = SimpleNamespace(
litellm_teammembership=SimpleNamespace(find_unique=AsyncMock(side_effect=failure)),
)
cache = SimpleNamespace(async_get_cache=AsyncMock(return_value=None), async_set_cache=AsyncMock())
monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db))
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
with pytest.raises(HTTPException) as rejected:
await live._managed_member_budget(UserAPIKeyAuth(api_key="owner", team_id="team", user_id="member"))
assert rejected.value.status_code == 503
@pytest.mark.asyncio
async def test_managed_budget_checks_project_team_and_model_group_tables(monkeypatch):
from litellm.proxy import proxy_server
project = LiteLLM_ProjectTable(
project_id="project",
budget_id="project-budget",
litellm_budget_table=LiteLLM_BudgetTable(max_budget=1),
)
team = LiteLLM_TeamTable(team_id="team", budget_limits=[{"budget_duration": "1d", "max_budget": 1}])
group = SimpleNamespace(
access_group_name="group",
litellm_budget_table=SimpleNamespace(max_budget=1),
)
def find_project(*, where, include):
QueryBuilder(method="find_unique", arguments={"where": where}).build_query()
return project
def find_groups(*, where, include):
QueryBuilder(method="find_many", arguments={"where": where}).build_query()
return [group]
db = SimpleNamespace(
litellm_teamtable=SimpleNamespace(find_unique=AsyncMock(return_value=team)),
litellm_projecttable=SimpleNamespace(find_unique=AsyncMock(side_effect=find_project)),
litellm_modelaccessgroupbudgettable=SimpleNamespace(find_many=AsyncMock(side_effect=find_groups)),
)
cache = SimpleNamespace(async_get_cache=AsyncMock(return_value=None), async_set_cache=AsyncMock())
monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db))
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
monkeypatch.setattr(live, "collect_matched_model_access_groups", AsyncMock(return_value=("group",)))
assert await live._managed_member_budget(UserAPIKeyAuth(api_key="owner", project_id="project")) is True
assert await live._managed_member_budget(UserAPIKeyAuth(api_key="owner", team_id="team")) is True
assert (
await live._managed_member_budget(
UserAPIKeyAuth(api_key="owner", matched_model_access_groups=["voice-group"]), model="backend"
)
is True
)
@pytest.mark.asyncio
@pytest.mark.parametrize("backend_budget, blocked", [(None, False), (0, False), (1, True)])
async def test_managed_budget_uses_the_delegated_model_group(monkeypatch, backend_budget, blocked):
from litellm.proxy import proxy_server
rows = [
SimpleNamespace(access_group_name="voice-group", litellm_budget_table=SimpleNamespace(max_budget=1)),
SimpleNamespace(
access_group_name="backend-group", litellm_budget_table=SimpleNamespace(max_budget=backend_budget)
),
]
db = SimpleNamespace(
litellm_modelaccessgroupbudgettable=SimpleNamespace(find_many=AsyncMock(return_value=rows)),
)
cache = SimpleNamespace(async_get_cache=AsyncMock(return_value=None), async_set_cache=AsyncMock())
monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db))
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
monkeypatch.setattr(proxy_server, "llm_router", SimpleNamespace())
monkeypatch.setattr(live, "collect_matched_model_access_groups", AsyncMock(return_value=("backend-group",)))
auth = UserAPIKeyAuth(api_key="owner", models=["voice", "backend-group"])
assert await live._managed_member_budget(auth, model="backend") is blocked
@pytest.mark.asyncio
async def test_responses_delegation_fails_closed_when_inherited_org_lookup_fails(monkeypatch):
from litellm.proxy import proxy_server
from litellm.proxy.auth import auth_checks
team = LiteLLM_TeamTable(team_id="team", organization_id="org", models=["*"])
group = SimpleNamespace(
access_group_name="backend-group",
litellm_budget_table=SimpleNamespace(max_budget=1),
)
db = SimpleNamespace(
litellm_teamtable=SimpleNamespace(find_unique=AsyncMock(return_value=team)),
litellm_modelaccessgroupbudgettable=SimpleNamespace(find_many=AsyncMock(return_value=[group])),
)
cache = SimpleNamespace(async_get_cache=AsyncMock(return_value=None), async_set_cache=AsyncMock())
monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db))
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
monkeypatch.setattr(proxy_server, "proxy_logging_obj", None)
monkeypatch.setattr(
proxy_server,
"llm_router",
SimpleNamespace(get_model_access_groups=lambda model_name, team_id=None: {"backend-group"}),
)
monkeypatch.setattr(
auth_checks,
"get_org_object",
AsyncMock(side_effect=RuntimeError("organization lookup unavailable")),
)
monkeypatch.setattr(live, "_authorize", AsyncMock())
with pytest.raises(HTTPException) as rejected:
await live._authorize_delegation(
{"session": {"delegation": {"type": "responses", "responses": {"model": "backend"}}}},
UserAPIKeyAuth(api_key="owner", models=["*"], team_id="team"),
)
assert rejected.value.status_code == 503
@pytest.mark.asyncio
async def test_authorize_checks_organization_inherited_from_team(monkeypatch):
from litellm.proxy import proxy_server
team = SimpleNamespace(organization_id="org")
org = SimpleNamespace(models=["backend"])
team_loader = AsyncMock(return_value=team)
org_loader = AsyncMock(return_value=org)
org_check = Mock()
monkeypatch.setattr(proxy_server, "prisma_client", object())
monkeypatch.setattr(proxy_server, "user_api_key_cache", object())
monkeypatch.setattr(proxy_server, "proxy_logging_obj", None)
monkeypatch.setattr(proxy_server, "llm_model_list", [])
monkeypatch.setattr(proxy_server, "llm_router", object())
monkeypatch.setattr(live, "can_key_call_resolved_model", AsyncMock())
monkeypatch.setattr(live, "get_team_object", team_loader, raising=False)
monkeypatch.setattr(live, "get_org_object", org_loader)
monkeypatch.setattr(live, "can_org_access_model", org_check)
await live._authorize("backend", UserAPIKeyAuth(api_key="owner", team_id="team"))
team_loader.assert_awaited_once()
org_loader.assert_awaited_once()
org_check.assert_called_once_with(model="backend", org_object=org, llm_router=proxy_server.llm_router)
@pytest.mark.asyncio
async def test_managed_budget_fails_closed_when_member_scope_lookup_fails_after_snapshot(monkeypatch):
from litellm.proxy import proxy_server
team = LiteLLM_TeamTable(team_id="team", models=["*"])
membership = LiteLLM_TeamMembership(
user_id="member",
team_id="team",
litellm_budget_table=LiteLLM_BudgetTable(allowed_models=["backend-group"]),
)
group = SimpleNamespace(
access_group_name="backend-group",
litellm_budget_table=SimpleNamespace(max_budget=1),
)
db = SimpleNamespace(
litellm_teammembership=SimpleNamespace(
find_unique=AsyncMock(side_effect=[membership, RuntimeError("membership lookup unavailable")])
),
litellm_teamtable=SimpleNamespace(find_unique=AsyncMock(return_value=team)),
litellm_modelaccessgroupbudgettable=SimpleNamespace(find_many=AsyncMock(return_value=[group])),
)
cache = SimpleNamespace(async_get_cache=AsyncMock(return_value=None), async_set_cache=AsyncMock())
monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db))
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
monkeypatch.setattr(proxy_server, "proxy_logging_obj", None)
monkeypatch.setattr(
proxy_server,
"llm_router",
SimpleNamespace(get_model_access_groups=lambda model_name, team_id=None: {"backend-group"}),
)
monkeypatch.setattr(live, "_authorize", AsyncMock())
with pytest.raises(HTTPException) as rejected:
await live._authorize_delegation(
{"session": {"delegation": {"type": "responses", "responses": {"model": "backend"}}}},
UserAPIKeyAuth(api_key="owner", models=["*"], team_id="team", user_id="member"),
)
assert rejected.value.status_code == 503