mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(live): enforce delegated model grants and managed budgets
This commit is contained in:
parent
fbda29e1b7
commit
153ece2769
5 changed files with 1005 additions and 80 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -299,6 +299,7 @@ class TestProcessResponse:
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("local_model_cost_map")
|
||||
class TestProcessEmbedContentResponseUsage:
|
||||
"""Gemini Embedding 2 embedContent usageMetadata must drive spend.
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue