mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(transcription): reject streams that bypass output guardrails
This commit is contained in:
parent
998bde0963
commit
6cb6811e45
3 changed files with 68 additions and 0 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue