mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-08 22:21:35 +00:00
fix(proxy): refuse OpenAI websocket passthrough on every enforced model allowlist and propagate the DB opt-in
This commit is contained in:
parent
a2d5215a4f
commit
7351911b53
6 changed files with 461 additions and 449 deletions
|
|
@ -4155,6 +4155,62 @@ async def _granted_model_lists(
|
|||
)
|
||||
|
||||
|
||||
async def enforced_model_allowlists(
|
||||
valid_token: UserAPIKeyAuth,
|
||||
prisma_client: PrismaClient | None,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
) -> tuple[Sequence[str], ...]:
|
||||
"""One model allowlist per level that ``common_checks`` enforces on a request from this identity."""
|
||||
key_models: Final = _resolve_key_models_for_auth_check(valid_token=valid_token)
|
||||
if prisma_client is None:
|
||||
return (key_models,)
|
||||
team_object: Final = (
|
||||
None
|
||||
if valid_token.team_id is None
|
||||
else await get_team_object(
|
||||
team_id=valid_token.team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
)
|
||||
user_object: Final = (
|
||||
None
|
||||
if team_object is not None
|
||||
else await get_user_object(
|
||||
user_id=valid_token.user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
)
|
||||
project_object: Final = (
|
||||
None
|
||||
if valid_token.project_id is None
|
||||
else await get_project_object(
|
||||
project_id=valid_token.project_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
)
|
||||
return (
|
||||
key_models,
|
||||
team_object.models if team_object is not None else (),
|
||||
await _team_member_granted_models(
|
||||
valid_token=valid_token,
|
||||
team_object=team_object,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
),
|
||||
user_object.models if user_object is not None else (),
|
||||
project_object.models if project_object is not None else (),
|
||||
)
|
||||
|
||||
|
||||
async def collect_matched_model_access_groups(
|
||||
model: str | Sequence[str] | None,
|
||||
valid_token: UserAPIKeyAuth | None,
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@ import inspect
|
|||
import json
|
||||
import os
|
||||
import re
|
||||
from collections.abc import AsyncGenerator, Callable, Mapping
|
||||
from collections.abc import AsyncGenerator, Callable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Annotated, Final, Protocol, cast
|
||||
|
|
@ -36,6 +36,7 @@ from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
|||
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
|
||||
from litellm.passthrough.main import AsyncPassthroughStreamingResponse
|
||||
from litellm.proxy._types import *
|
||||
from litellm.proxy.auth.auth_checks import enforced_model_allowlists
|
||||
from litellm.proxy.auth.handle_jwt import JWTHandler
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
from litellm.proxy.auth.user_api_key_auth import (
|
||||
|
|
@ -2341,9 +2342,8 @@ _OPENAI_WS_ALL_MODEL_ACCESS: Final = frozenset(
|
|||
)
|
||||
|
||||
|
||||
def _key_has_model_restrictions(user_api_key_dict: UserAPIKeyAuth) -> bool:
|
||||
scoped_models: Final = (*user_api_key_dict.models, *user_api_key_dict.team_models)
|
||||
return any(str(model) not in _OPENAI_WS_ALL_MODEL_ACCESS for model in scoped_models)
|
||||
def _has_model_restrictions(model_allowlists: tuple[Sequence[str], ...]) -> bool:
|
||||
return any(str(model) not in _OPENAI_WS_ALL_MODEL_ACCESS for allowlist in model_allowlists for model in allowlist)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
@ -2376,12 +2376,18 @@ def _is_openai_websocket_passthrough_enabled(general_settings: Mapping[str, obje
|
|||
return setting is True
|
||||
|
||||
|
||||
def _openai_websocket_refusal(
|
||||
user_api_key_dict: UserAPIKeyAuth, general_settings: Mapping[str, object]
|
||||
class _OpenAIWebsocketModelAllowlists(Protocol):
|
||||
async def __call__(self, valid_token: UserAPIKeyAuth, /) -> tuple[Sequence[str], ...]: ...
|
||||
|
||||
|
||||
async def _openai_websocket_refusal(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
general_settings: Mapping[str, object],
|
||||
model_allowlists: _OpenAIWebsocketModelAllowlists,
|
||||
) -> _OpenAIWebsocketRefusal | None:
|
||||
if not _is_openai_websocket_passthrough_enabled(general_settings):
|
||||
return _OPENAI_WS_DISABLED_REFUSAL
|
||||
if _key_has_model_restrictions(user_api_key_dict):
|
||||
if _has_model_restrictions(await model_allowlists(user_api_key_dict)):
|
||||
return _OPENAI_WS_MODEL_RESTRICTED_REFUSAL
|
||||
return None
|
||||
|
||||
|
|
@ -2410,6 +2416,20 @@ def _openai_websocket_relay() -> _OpenAIWebsocketRelay:
|
|||
return websocket_passthrough_request
|
||||
|
||||
|
||||
def _proxy_model_allowlists() -> _OpenAIWebsocketModelAllowlists:
|
||||
from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache
|
||||
|
||||
async def resolve(valid_token: UserAPIKeyAuth, /) -> tuple[Sequence[str], ...]:
|
||||
return await enforced_model_allowlists(
|
||||
valid_token=valid_token,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
return resolve
|
||||
|
||||
|
||||
@router.websocket("/openai_passthrough/{endpoint:path}")
|
||||
@router.websocket("/openai/{endpoint:path}")
|
||||
async def openai_websocket_proxy_route(
|
||||
|
|
@ -2418,6 +2438,7 @@ async def openai_websocket_proxy_route(
|
|||
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth_websocket)],
|
||||
general_settings: Annotated[Mapping[str, object], Depends(_proxy_general_settings)],
|
||||
relay: Annotated[_OpenAIWebsocketRelay, Depends(_openai_websocket_relay)],
|
||||
model_allowlists: Annotated[_OpenAIWebsocketModelAllowlists, Depends(_proxy_model_allowlists)],
|
||||
) -> None:
|
||||
"""WebSocket passthrough for OpenAI prefixes (realtime / responses.connect)."""
|
||||
requested_subprotocols: Final = tuple(
|
||||
|
|
@ -2427,7 +2448,7 @@ async def openai_websocket_proxy_route(
|
|||
)
|
||||
negotiated_subprotocol: Final = requested_subprotocols[0] if requested_subprotocols else None
|
||||
|
||||
refusal: Final = _openai_websocket_refusal(user_api_key_dict, general_settings)
|
||||
refusal: Final = await _openai_websocket_refusal(user_api_key_dict, general_settings, model_allowlists)
|
||||
if refusal is not None:
|
||||
await websocket.accept(subprotocol=negotiated_subprotocol)
|
||||
await websocket.send_text(
|
||||
|
|
|
|||
|
|
@ -6793,6 +6793,11 @@ class ProxyConfig:
|
|||
else:
|
||||
general_settings["apply_user_budget_to_team_keys"] = db_value if db_value is None else bool(db_value)
|
||||
|
||||
if "enable_openai_websocket_passthrough" not in self._yaml_general_settings_keys:
|
||||
general_settings["enable_openai_websocket_passthrough"] = _general_settings.get(
|
||||
"enable_openai_websocket_passthrough"
|
||||
)
|
||||
|
||||
## STORE MODEL IN DB ##
|
||||
if "store_model_in_db" in _general_settings:
|
||||
value = _general_settings["store_model_in_db"]
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -1,7 +1,7 @@
|
|||
"""OpenAI passthrough WebSocket route: registration, opt-in gating, and refusals."""
|
||||
|
||||
import json
|
||||
from collections.abc import Mapping
|
||||
from collections.abc import Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType, SimpleNamespace
|
||||
from typing import Final
|
||||
|
|
@ -15,10 +15,13 @@ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
|
|||
_OPENAI_WS_DISABLED_REFUSAL,
|
||||
_OPENAI_WS_MODEL_RESTRICTED_REFUSAL,
|
||||
_openai_websocket_refusal,
|
||||
_proxy_model_allowlists,
|
||||
openai_websocket_proxy_route,
|
||||
router,
|
||||
)
|
||||
|
||||
Scopes = tuple[Sequence[str], ...]
|
||||
|
||||
ENABLED: Final = MappingProxyType({"enable_openai_websocket_passthrough": True})
|
||||
DISABLED_SETTINGS: Final = (
|
||||
MappingProxyType({}),
|
||||
|
|
@ -96,21 +99,39 @@ class _FakeRelay:
|
|||
)
|
||||
|
||||
|
||||
class _FakeModelAllowlists:
|
||||
def __init__(self, scopes: Scopes) -> None:
|
||||
self.scopes = scopes
|
||||
self.calls: list[UserAPIKeyAuth] = []
|
||||
|
||||
async def __call__(self, valid_token: UserAPIKeyAuth, /) -> Scopes:
|
||||
self.calls.append(valid_token)
|
||||
return self.scopes
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Served:
|
||||
relay: _FakeRelay
|
||||
allowlists: _FakeModelAllowlists
|
||||
|
||||
|
||||
async def _serve(
|
||||
websocket: _FakeWebSocket,
|
||||
endpoint: str,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
general_settings: Mapping[str, object],
|
||||
) -> _FakeRelay:
|
||||
relay = _FakeRelay()
|
||||
scopes: Scopes = (),
|
||||
) -> _Served:
|
||||
served = _Served(relay=_FakeRelay(), allowlists=_FakeModelAllowlists(scopes))
|
||||
await openai_websocket_proxy_route(
|
||||
websocket=websocket,
|
||||
endpoint=endpoint,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
general_settings=general_settings,
|
||||
relay=relay,
|
||||
relay=served.relay,
|
||||
model_allowlists=served.allowlists,
|
||||
)
|
||||
return relay
|
||||
return served
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -120,9 +141,9 @@ async def test_openai_websocket_forwards_query_and_keeps_provider_auth(prefix, m
|
|||
websocket = _FakeWebSocket(f"/{prefix}/v1/realtime", "model=gpt-4o-realtime-preview")
|
||||
|
||||
with patch(GET_CREDENTIALS, return_value="sk-provider"):
|
||||
relay = await _serve(websocket, "v1/realtime", UserAPIKeyAuth(), ENABLED)
|
||||
served = await _serve(websocket, "v1/realtime", UserAPIKeyAuth(), ENABLED)
|
||||
|
||||
assert relay.calls == [
|
||||
assert served.relay.calls == [
|
||||
_RelayCall(
|
||||
target="wss://api.openai.com/v1/realtime?model=gpt-4o-realtime-preview",
|
||||
custom_headers=MappingProxyType({"Authorization": "Bearer sk-provider"}),
|
||||
|
|
@ -145,10 +166,10 @@ async def test_openai_websocket_accepts_first_client_subprotocol():
|
|||
)
|
||||
|
||||
with patch(GET_CREDENTIALS, return_value="sk-provider"):
|
||||
relay = await _serve(websocket, "v1/realtime", UserAPIKeyAuth(), ENABLED)
|
||||
served = await _serve(websocket, "v1/realtime", UserAPIKeyAuth(), ENABLED)
|
||||
|
||||
assert websocket.accepts == ["realtime"]
|
||||
assert [call.accept_websocket for call in relay.calls] == [False]
|
||||
assert [call.accept_websocket for call in served.relay.calls] == [False]
|
||||
assert websocket.closed is None
|
||||
|
||||
|
||||
|
|
@ -157,13 +178,13 @@ async def test_openai_websocket_closes_cleanly_when_provider_credentials_missing
|
|||
websocket = _FakeWebSocket("/openai/v1/realtime", "model=gpt-4o-realtime-preview")
|
||||
|
||||
with patch(GET_CREDENTIALS, return_value=None):
|
||||
relay = await _serve(websocket, "v1/realtime", UserAPIKeyAuth(), ENABLED)
|
||||
served = await _serve(websocket, "v1/realtime", UserAPIKeyAuth(), ENABLED)
|
||||
|
||||
assert websocket.closed is not None
|
||||
assert websocket.closed[0] == 1011
|
||||
assert "OPENAI_API_KEY" in websocket.closed[1]
|
||||
assert websocket.accepts == []
|
||||
assert relay.calls == []
|
||||
assert served.relay.calls == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -172,23 +193,26 @@ async def test_openai_websocket_closes_cleanly_when_provider_credentials_missing
|
|||
async def test_openai_websocket_refused_unless_explicitly_enabled(prefix, general_settings):
|
||||
websocket = _FakeWebSocket(f"/{prefix}/v1/realtime", "model=gpt-4o-realtime-preview")
|
||||
|
||||
relay = await _serve(websocket, "v1/realtime", UserAPIKeyAuth(), general_settings)
|
||||
served = await _serve(websocket, "v1/realtime", UserAPIKeyAuth(), general_settings)
|
||||
|
||||
assert "enable_openai_websocket_passthrough" in websocket.error_message()
|
||||
assert websocket.accepts == [None]
|
||||
assert websocket.closed == (1008, _OPENAI_WS_DISABLED_REFUSAL.close_reason)
|
||||
assert relay.calls == []
|
||||
assert served.relay.calls == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("general_settings", DISABLED_SETTINGS)
|
||||
def test_openai_websocket_refusal_is_disabled_for_falsy_settings(general_settings):
|
||||
assert _openai_websocket_refusal(UserAPIKeyAuth(), general_settings) is _OPENAI_WS_DISABLED_REFUSAL
|
||||
async def test_openai_websocket_refusal_is_disabled_for_falsy_settings(general_settings):
|
||||
refusal = await _openai_websocket_refusal(UserAPIKeyAuth(), general_settings, _FakeModelAllowlists(()))
|
||||
assert refusal is _OPENAI_WS_DISABLED_REFUSAL
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("value", [True, "true", "True"])
|
||||
def test_openai_websocket_refusal_is_none_for_truthy_settings(value):
|
||||
async def test_openai_websocket_refusal_is_none_for_truthy_settings(value):
|
||||
settings = MappingProxyType({"enable_openai_websocket_passthrough": value})
|
||||
assert _openai_websocket_refusal(UserAPIKeyAuth(), settings) is None
|
||||
assert await _openai_websocket_refusal(UserAPIKeyAuth(), settings, _FakeModelAllowlists(())) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -199,51 +223,73 @@ async def test_openai_websocket_refusal_echoes_requested_subprotocol():
|
|||
subprotocols="realtime, openai-beta.realtime-v1",
|
||||
)
|
||||
|
||||
relay = await _serve(websocket, "v1/realtime", UserAPIKeyAuth(), MappingProxyType({}))
|
||||
served = await _serve(websocket, "v1/realtime", UserAPIKeyAuth(), MappingProxyType({}))
|
||||
|
||||
assert websocket.accepts == ["realtime"]
|
||||
assert websocket.closed == (1008, _OPENAI_WS_DISABLED_REFUSAL.close_reason)
|
||||
assert relay.calls == []
|
||||
assert served.relay.calls == []
|
||||
|
||||
|
||||
RESTRICTED_KEYS: Final = (
|
||||
UserAPIKeyAuth(models=["gpt-4o"]),
|
||||
UserAPIKeyAuth(team_models=["gpt-4o-realtime-preview"]),
|
||||
UserAPIKeyAuth(models=["all-team-models"], team_models=["gpt-4o"]),
|
||||
RESTRICTED_SCOPES: Final[tuple[Scopes, ...]] = (
|
||||
(("gpt-4o",),),
|
||||
((), ("gpt-4o-realtime-preview",)),
|
||||
(("all-team-models",), ("gpt-4o",)),
|
||||
((), ("all-proxy-models",), ("gpt-4o",)),
|
||||
((), (), (), ("gpt-4o",)),
|
||||
(("*",), (), (), (), ("gpt-4o",)),
|
||||
)
|
||||
UNRESTRICTED_KEYS: Final = (
|
||||
UserAPIKeyAuth(),
|
||||
UserAPIKeyAuth(models=["all-proxy-models"]),
|
||||
UserAPIKeyAuth(models=["*"]),
|
||||
UserAPIKeyAuth(models=["all-team-models"], team_models=["all-proxy-models"]),
|
||||
UNRESTRICTED_SCOPES: Final[tuple[Scopes, ...]] = (
|
||||
(),
|
||||
((),),
|
||||
(("all-proxy-models",),),
|
||||
(("*",),),
|
||||
(("all-team-models",), ("all-proxy-models",)),
|
||||
((), (), (), (), ()),
|
||||
(("*",), ("all-proxy-models",), ("all-team-models",), (), ()),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("user_api_key_dict", RESTRICTED_KEYS)
|
||||
async def test_openai_websocket_rejects_model_restricted_keys(user_api_key_dict):
|
||||
@pytest.mark.parametrize("scopes", RESTRICTED_SCOPES)
|
||||
async def test_openai_websocket_rejects_model_restricted_identities(scopes):
|
||||
websocket = _FakeWebSocket("/openai/v1/realtime", "model=gpt-4o-realtime-preview")
|
||||
user_api_key_dict = UserAPIKeyAuth(token="hashed-fake", user_id="user-fake", team_id="team-fake")
|
||||
|
||||
relay = await _serve(websocket, "v1/realtime", user_api_key_dict, ENABLED)
|
||||
served = await _serve(websocket, "v1/realtime", user_api_key_dict, ENABLED, scopes)
|
||||
|
||||
assert "model restrictions" in websocket.error_message()
|
||||
assert websocket.closed == (1008, _OPENAI_WS_MODEL_RESTRICTED_REFUSAL.close_reason)
|
||||
assert relay.calls == []
|
||||
|
||||
|
||||
@pytest.mark.parametrize("user_api_key_dict", RESTRICTED_KEYS)
|
||||
def test_openai_websocket_refusal_prefers_disabled_over_model_restriction(user_api_key_dict):
|
||||
assert _openai_websocket_refusal(user_api_key_dict, MappingProxyType({})) is _OPENAI_WS_DISABLED_REFUSAL
|
||||
assert served.relay.calls == []
|
||||
assert served.allowlists.calls == [user_api_key_dict]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("user_api_key_dict", UNRESTRICTED_KEYS)
|
||||
async def test_openai_websocket_allows_unrestricted_keys(user_api_key_dict):
|
||||
@pytest.mark.parametrize("scopes", RESTRICTED_SCOPES)
|
||||
async def test_openai_websocket_disabled_refusal_skips_allowlist_lookups(scopes):
|
||||
allowlists = _FakeModelAllowlists(scopes)
|
||||
|
||||
refusal = await _openai_websocket_refusal(UserAPIKeyAuth(), MappingProxyType({}), allowlists)
|
||||
|
||||
assert refusal is _OPENAI_WS_DISABLED_REFUSAL
|
||||
assert allowlists.calls == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("scopes", UNRESTRICTED_SCOPES)
|
||||
async def test_openai_websocket_allows_unrestricted_identities(scopes):
|
||||
websocket = _FakeWebSocket("/openai/v1/responses", "")
|
||||
|
||||
with patch(GET_CREDENTIALS, return_value="sk-provider"):
|
||||
relay = await _serve(websocket, "v1/responses", user_api_key_dict, ENABLED)
|
||||
served = await _serve(websocket, "v1/responses", UserAPIKeyAuth(), ENABLED, scopes)
|
||||
|
||||
assert len(relay.calls) == 1
|
||||
assert len(served.relay.calls) == 1
|
||||
assert websocket.sent == []
|
||||
assert websocket.closed is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_proxy_model_allowlists_reads_the_key_scope_without_a_database():
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", None):
|
||||
scopes = await _proxy_model_allowlists()(UserAPIKeyAuth(models=["gpt-4o"]))
|
||||
|
||||
assert tuple(tuple(scope) for scope in scopes) == (("gpt-4o",),)
|
||||
|
|
|
|||
|
|
@ -11568,14 +11568,10 @@ async def test_key_window_spend_row_is_enqueued_with_the_actual_cost():
|
|||
|
||||
reset_at = datetime.now(timezone.utc) + timedelta(days=10)
|
||||
key_obj = MagicMock()
|
||||
key_obj.budget_limits = [
|
||||
{"budget_duration": "30d", "max_budget": 100.0, "reset_at": reset_at.isoformat()}
|
||||
]
|
||||
key_obj.budget_limits = [{"budget_duration": "30d", "max_budget": 100.0, "reset_at": reset_at.isoformat()}]
|
||||
|
||||
with _window_spend_enqueue_env({"hashed-token": key_obj}) as queue:
|
||||
await increment_spend_counters(
|
||||
token="hashed-token", team_id=None, user_id=None, response_cost=0.25
|
||||
)
|
||||
await increment_spend_counters(token="hashed-token", team_id=None, user_id=None, response_cost=0.25)
|
||||
enqueued = await _drain(queue)
|
||||
|
||||
assert len(enqueued) == 1
|
||||
|
|
@ -11594,14 +11590,10 @@ async def test_team_window_spend_row_is_enqueued():
|
|||
|
||||
reset_at = datetime.now(timezone.utc) + timedelta(days=3)
|
||||
team_obj = MagicMock()
|
||||
team_obj.budget_limits = [
|
||||
{"budget_duration": "7d", "max_budget": 50.0, "reset_at": reset_at.isoformat()}
|
||||
]
|
||||
team_obj.budget_limits = [{"budget_duration": "7d", "max_budget": 50.0, "reset_at": reset_at.isoformat()}]
|
||||
|
||||
with _window_spend_enqueue_env({"team_id:team-1": team_obj}) as queue:
|
||||
await increment_spend_counters(
|
||||
token=None, team_id="team-1", user_id=None, response_cost=1.5
|
||||
)
|
||||
await increment_spend_counters(token=None, team_id="team-1", user_id=None, response_cost=1.5)
|
||||
enqueued = await _drain(queue)
|
||||
|
||||
assert len(enqueued) == 1
|
||||
|
|
@ -11620,9 +11612,7 @@ async def test_window_spend_row_is_enqueued_even_when_the_counter_was_reserved()
|
|||
|
||||
reset_at = datetime.now(timezone.utc) + timedelta(days=10)
|
||||
key_obj = MagicMock()
|
||||
key_obj.budget_limits = [
|
||||
{"budget_duration": "30d", "max_budget": 100.0, "reset_at": reset_at.isoformat()}
|
||||
]
|
||||
key_obj.budget_limits = [{"budget_duration": "30d", "max_budget": 100.0, "reset_at": reset_at.isoformat()}]
|
||||
reservation = {
|
||||
"entries": [
|
||||
{"counter_key": "spend:key:hashed-token", "reserved": 1.0},
|
||||
|
|
@ -11660,9 +11650,7 @@ async def test_sliding_window_without_reset_at_is_not_enqueued():
|
|||
key_obj.budget_limits = [{"budget_duration": "30d", "max_budget": 100.0}]
|
||||
|
||||
with _window_spend_enqueue_env({"hashed-token": key_obj}) as queue:
|
||||
await increment_spend_counters(
|
||||
token="hashed-token", team_id=None, user_id=None, response_cost=0.25
|
||||
)
|
||||
await increment_spend_counters(token="hashed-token", team_id=None, user_id=None, response_cost=0.25)
|
||||
enqueued = await _drain(queue)
|
||||
|
||||
assert enqueued == []
|
||||
|
|
@ -11680,9 +11668,7 @@ async def test_each_configured_window_gets_its_own_row_enqueue():
|
|||
]
|
||||
|
||||
with _window_spend_enqueue_env({"hashed-token": key_obj}) as queue:
|
||||
await increment_spend_counters(
|
||||
token="hashed-token", team_id=None, user_id=None, response_cost=0.25
|
||||
)
|
||||
await increment_spend_counters(token="hashed-token", team_id=None, user_id=None, response_cost=0.25)
|
||||
enqueued = await _drain(queue)
|
||||
|
||||
assert sorted(item["window_duration"] for item in enqueued) == ["1d", "30d"]
|
||||
|
|
@ -11697,9 +11683,7 @@ async def test_no_window_spend_row_enqueued_without_budget_limits():
|
|||
key_obj.budget_limits = None
|
||||
|
||||
with _window_spend_enqueue_env({"hashed-token": key_obj}) as queue:
|
||||
await increment_spend_counters(
|
||||
token="hashed-token", team_id=None, user_id=None, response_cost=0.25
|
||||
)
|
||||
await increment_spend_counters(token="hashed-token", team_id=None, user_id=None, response_cost=0.25)
|
||||
enqueued = await _drain(queue)
|
||||
|
||||
assert enqueued == []
|
||||
|
|
@ -11713,9 +11697,7 @@ async def test_window_spend_row_carries_the_request_start_time():
|
|||
|
||||
reset_at = datetime.now(timezone.utc) + timedelta(days=10)
|
||||
key_obj = MagicMock()
|
||||
key_obj.budget_limits = [
|
||||
{"budget_duration": "30d", "max_budget": 100.0, "reset_at": reset_at.isoformat()}
|
||||
]
|
||||
key_obj.budget_limits = [{"budget_duration": "30d", "max_budget": 100.0, "reset_at": reset_at.isoformat()}]
|
||||
|
||||
with _window_spend_enqueue_env({"hashed-token": key_obj}) as queue:
|
||||
await increment_spend_counters(
|
||||
|
|
@ -11736,9 +11718,7 @@ async def test_team_window_spend_row_carries_the_request_start_time():
|
|||
|
||||
reset_at = datetime.now(timezone.utc) + timedelta(days=3)
|
||||
team_obj = MagicMock()
|
||||
team_obj.budget_limits = [
|
||||
{"budget_duration": "7d", "max_budget": 50.0, "reset_at": reset_at.isoformat()}
|
||||
]
|
||||
team_obj.budget_limits = [{"budget_duration": "7d", "max_budget": 50.0, "reset_at": reset_at.isoformat()}]
|
||||
|
||||
with _window_spend_enqueue_env({"team_id:team-1": team_obj}) as queue:
|
||||
await increment_spend_counters(
|
||||
|
|
@ -12050,7 +12030,6 @@ async def test_init_guardrails_in_db_snapshots_and_reconciles_under_guardrail_re
|
|||
assert not GUARDRAIL_RECONCILE_LOCK.locked()
|
||||
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_init_prompts_in_db_reloads_rows_patched_on_another_worker(monkeypatch):
|
||||
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
|
||||
|
|
@ -12094,7 +12073,9 @@ async def test_init_prompts_in_db_reloads_rows_patched_on_another_worker(monkeyp
|
|||
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
|
||||
assert served_content() == "Begin every reply with AHOY"
|
||||
|
||||
prisma_client.db.litellm_prompttable.find_many = AsyncMock(return_value=[db_row("Begin every reply with HOWDY")])
|
||||
prisma_client.db.litellm_prompttable.find_many = AsyncMock(
|
||||
return_value=[db_row("Begin every reply with HOWDY")]
|
||||
)
|
||||
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
|
||||
|
||||
assert served_content() == "Begin every reply with HOWDY"
|
||||
|
|
@ -12547,3 +12528,40 @@ def test_disabling_docs_does_not_disable_other_routes(monkeypatch):
|
|||
|
||||
assert client.get("/redoc").status_code == 404
|
||||
assert client.get("/health/liveliness").status_code == 200
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"db_general_settings, expected",
|
||||
[
|
||||
({"enable_openai_websocket_passthrough": True}, True),
|
||||
({"enable_openai_websocket_passthrough": False}, False),
|
||||
({}, None),
|
||||
],
|
||||
)
|
||||
async def test_update_general_settings_propagates_openai_websocket_passthrough(db_general_settings, expected):
|
||||
from litellm.proxy.proxy_server import ProxyConfig
|
||||
|
||||
proxy_config = ProxyConfig()
|
||||
|
||||
with patch("litellm.proxy.proxy_server.general_settings", {"enable_openai_websocket_passthrough": True}):
|
||||
await proxy_config._update_general_settings(db_general_settings=db_general_settings)
|
||||
|
||||
import litellm.proxy.proxy_server as ps
|
||||
|
||||
assert ps.general_settings["enable_openai_websocket_passthrough"] is expected
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_general_settings_keeps_yaml_openai_websocket_passthrough():
|
||||
from litellm.proxy.proxy_server import ProxyConfig
|
||||
|
||||
proxy_config = ProxyConfig()
|
||||
proxy_config._yaml_general_settings_keys = {"enable_openai_websocket_passthrough"}
|
||||
|
||||
with patch("litellm.proxy.proxy_server.general_settings", {"enable_openai_websocket_passthrough": False}):
|
||||
await proxy_config._update_general_settings(db_general_settings={"enable_openai_websocket_passthrough": True})
|
||||
|
||||
import litellm.proxy.proxy_server as ps
|
||||
|
||||
assert ps.general_settings["enable_openai_websocket_passthrough"] is False
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue