diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 26e9b02aace..caf9250c4f6 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -12423,6 +12423,12 @@ async def audio_transcriptions( call_type="transcription", ) + if data.get("stream") is True and ProxyLogging.has_post_call_guardrails(data, llm_router): + raise HTTPException( + detail="Streaming transcription does not support output guardrails. Use stream=false.", + status_code=status.HTTP_400_BAD_REQUEST, + ) + ## ROUTE TO CORRECT ENDPOINT ## llm_call: Final = await route_request( data=data, diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index b8cc30ad8a7..d6978d04590 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -2874,6 +2874,17 @@ class ProxyLogging: def has_post_call_response_headers_callbacks() -> bool: return ProxyLogging._callback_capabilities().has_post_call_response_headers + @staticmethod + def has_post_call_guardrails(request_data: Mapping[str, object], llm_router: Router | None) -> bool: + guardrail_data: Final = _check_and_merge_model_level_guardrails( + data=dict(request_data), llm_router=llm_router, trust_client_model_info=False + ) + guardrails, _ = _partition_post_call_callbacks() + return bool(pipeline_managed_guardrail_names(guardrail_data, "post_call")) or any( + guardrail.should_run_guardrail(data=guardrail_data, event_type=GuardrailEventHooks.post_call) + for guardrail in guardrails + ) + @staticmethod def has_streaming_callbacks() -> bool: caps: Final = ProxyLogging._callback_capabilities() diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_audio.py b/tests/test_litellm/proxy/proxy_server/test_routes_audio.py index a756454405e..03cfe3fa67e 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_audio.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_audio.py @@ -10,13 +10,21 @@ Pins (PR2): from __future__ import annotations import io +from collections.abc import Callable +from contextlib import AbstractContextManager +from typing import Final from unittest.mock import AsyncMock, MagicMock import httpx import pytest +from fastapi.testclient import TestClient +import litellm +from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.proxy import proxy_server +from litellm.types.guardrails import GuardrailEventHooks from litellm.types.llms.openai import HttpxBinaryResponseContent +from litellm.types.proxy.policy_engine.pipeline_types import GuardrailPipeline, PipelineStep @pytest.fixture @@ -305,3 +313,46 @@ def test_audio_transcription_stream_returns_sse(client, auth_as, patched_transcr assert response.headers["content-type"].startswith("text/event-stream") assert response.text == 'data: {"type":"transcript.text.done","text":"hello world"}\n\n' assert patched_transcription_stream.closed is True + + +@pytest.mark.usefixtures("patched_transcription_stream") +@pytest.mark.parametrize( + "configuration,expected_status", + [("default", 400), ("model", 400), ("policy", 400), ("pre_call", 200), ("disabled", 200)], +) +def test_streaming_transcription_rejects_applicable_output_guardrails( + client: TestClient, + auth_as: Callable[[], AbstractContextManager[None]], + monkeypatch: pytest.MonkeyPatch, + configuration: str, + expected_status: int, +) -> None: + guardrail: Final = CustomGuardrail( + guardrail_name="transcription-output", + event_hook=GuardrailEventHooks.pre_call if configuration == "pre_call" else GuardrailEventHooks.post_call, + default_on=configuration in ("default", "pre_call"), + ) + monkeypatch.setattr(litellm, "callbacks", [guardrail]) + proxy_server.llm_router.get_model_list.return_value = ( + [{"litellm_params": {"guardrails": ["transcription-output"]}}] if configuration == "model" else [] + ) + if configuration == "policy": + pipeline: Final = GuardrailPipeline(mode="post_call", steps=[PipelineStep(guardrail="transcription-output")]) + proxy_server.proxy_logging_obj.pre_call_hook.side_effect = lambda **kwargs: { + **kwargs["data"], + "metadata": {"_guardrail_pipelines": [("transcription-policy", pipeline)]}, + } + + with auth_as(): + response: Final = client.post( + "/v1/audio/transcriptions", + files={"file": ("sample.wav", b"audio", "audio/wav")}, + data={"model": "gpt-transcribe", "stream": "true"}, + ) + + assert response.status_code == expected_status + if expected_status == 400: + assert "stream=false" in response.json()["error"]["message"] + assert "hello world" not in response.text + else: + assert '"text":"hello world"' in response.text