fix(proxy): refuse OpenAI websocket passthrough on every enforced model allowlist and propagate the DB opt-in

This commit is contained in:
mateo-berri 2026-09-04 18:39:56 -07:00
parent a2d5215a4f
commit 7351911b53
6 changed files with 461 additions and 449 deletions

View file

@ -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,

View file

@ -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(

View file

@ -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

View file

@ -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",),)

View file

@ -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