diff --git a/litellm/main.py b/litellm/main.py index aa1a632b0e7..49ee53e1c38 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -7808,7 +7808,10 @@ def adapter_completion(*, adapter_id: str, **kwargs) -> BaseModel | AdapterCompl def moderation(input: str, model: str | None = None, api_key: str | None = None, **kwargs) -> OpenAIModerationResponse: - litellm.ClinePassConfig.validate_moderation(model=model, custom_llm_provider=kwargs.get("custom_llm_provider")) + custom_llm_provider: Final[object] = kwargs.get("custom_llm_provider") + litellm.ClinePassConfig.validate_moderation( + model=model, custom_llm_provider=custom_llm_provider if isinstance(custom_llm_provider, str) else None + ) # only supports open ai for now api_key = api_key or litellm.api_key or litellm.openai_key or get_secret_str("OPENAI_API_KEY") diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 475acd5b967..02e43bc8bb8 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -2961,11 +2961,18 @@ class ManagedResponsesWebSocketHandler: call_kwargs: Final = self._build_base_call_kwargs(msg_obj) call_kwargs["stream"] = True - # A frame that repeats the connection's public alias (model_group) must - # reuse the router-resolved self.model; passing the alias raw to - # litellm.aresponses fails in get_llm_provider. A genuinely different - # provider-prefixed per-frame model is still honored. requested_model: Final[str | None] = _optional_str(call_kwargs.pop("model", None)) + authorized_models: Final = (self.model, self.model_group, f"{self.custom_llm_provider}/{self.model}") + if ( + self.user_api_key_dict is not None + and requested_model is not None + and requested_model not in authorized_models + ): + await self._send_error( + "Changing models requires a new authorized WebSocket connection", + error_type="invalid_request_error", + ) + return model: Final[str] = ( self.model if requested_model is None or requested_model == self.model_group else requested_model ) diff --git a/tests/unit/llms/clinepass/chat/test_clinepass_chat_transformation.py b/tests/unit/llms/clinepass/chat/test_clinepass_chat_transformation.py index 85115d754a2..2bb24e6e541 100644 --- a/tests/unit/llms/clinepass/chat/test_clinepass_chat_transformation.py +++ b/tests/unit/llms/clinepass/chat/test_clinepass_chat_transformation.py @@ -23,6 +23,7 @@ from litellm.llms.clinepass.chat.transformation import ( _unwrap_response_envelope, ) from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler +from litellm.proxy._types import UserAPIKeyAuth from litellm.responses.main import _aresponses_websocket from litellm.types.utils import LlmProviders from litellm.utils import ProviderConfigManager @@ -127,9 +128,10 @@ async def test_async_chat_sends_only_clinepass_credentials(monkeypatch, unrelate @pytest.mark.parametrize("credential_source", ["missing", "environment", "explicit"]) @pytest.mark.parametrize("connection_provider", ["clinepass", "mistral"]) +@pytest.mark.parametrize("changed_model", [None, "openai/gpt-4o", "clinepass/unauthorized-model"]) @pytest.mark.asyncio async def test_managed_responses_websocket_sends_only_clinepass_credentials( - monkeypatch, unrelated_credentials, credential_source, connection_provider + monkeypatch, unrelated_credentials, credential_source, connection_provider, changed_model ): if credential_source == "missing": monkeypatch.delenv("CLINEPASS_API_KEY", raising=False) @@ -142,7 +144,29 @@ async def test_managed_responses_websocket_sends_only_clinepass_credentials( ) sent = [] received = [] - lifecycle = iter(({"type": "websocket.connect"}, {"type": "websocket.disconnect", "code": 1000})) + foreign_requests = [] + lifecycle = iter( + ( + {"type": "websocket.connect"}, + *( + ( + { + "type": "websocket.receive", + "text": json.dumps({"type": "response.create", "model": changed_model, "input": "hi"}), + }, + ) + if changed_model is not None + else () + ), + {"type": "websocket.disconnect", "code": 1000}, + ) + ) + + async def block_foreign_request(self, request, *args, **kwargs): + foreign_requests.append(str(request.url)) + raise AssertionError("Unexpected provider HTTP request") + + monkeypatch.setattr(httpx.AsyncClient, "send", block_foreign_request) async def receive(): return next(lifecycle) @@ -179,6 +203,11 @@ async def test_managed_responses_websocket_sends_only_clinepass_credentials( model=f"{connection_provider}/deepseek-v4-flash", websocket=websocket, api_key=connection_key, + user_api_key_dict=( + UserAPIKeyAuth(models=[f"{connection_provider}/deepseek-v4-flash"]) + if changed_model is not None + else None + ), first_message=json.dumps( {"type": "response.create", "model": "clinepass/deepseek-v4-flash", "input": "ping"} ), @@ -200,12 +229,26 @@ async def test_managed_responses_websocket_sends_only_clinepass_credentials( if credential_source != "missing" else None ) - assert sent == [ - ("https://api.cline.bot/api/v1/chat/completions", f"Bearer {expected_key}" if expected_key else None) - ] + assert sent == ( + [] + if changed_model is not None and connection_provider == "mistral" + else [("https://api.cline.bot/api/v1/chat/completions", f"Bearer {expected_key}" if expected_key else None)] + ) assert result is None - assert "response.completed" in [event["type"] for event in received] - assert "error" not in [event["type"] for event in received] + errors: Final = [event["error"] for event in received if event["type"] == "error"] + assert errors == ( + [ + { + "type": "invalid_request_error", + "message": "Changing models requires a new authorized WebSocket connection", + } + ] + * (2 if connection_provider == "mistral" else 1) + if changed_model is not None + else [] + ) + assert foreign_requests == [] + assert ("response.completed" in [event["type"] for event in received]) == bool(sent) # --------------------------------------------------------------------------