mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +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",
|
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 ##
|
## ROUTE TO CORRECT ENDPOINT ##
|
||||||
llm_call: Final = await route_request(
|
llm_call: Final = await route_request(
|
||||||
data=data,
|
data=data,
|
||||||
|
|
|
||||||
|
|
@ -2874,6 +2874,17 @@ class ProxyLogging:
|
||||||
def has_post_call_response_headers_callbacks() -> bool:
|
def has_post_call_response_headers_callbacks() -> bool:
|
||||||
return ProxyLogging._callback_capabilities().has_post_call_response_headers
|
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
|
@staticmethod
|
||||||
def has_streaming_callbacks() -> bool:
|
def has_streaming_callbacks() -> bool:
|
||||||
caps: Final = ProxyLogging._callback_capabilities()
|
caps: Final = ProxyLogging._callback_capabilities()
|
||||||
|
|
|
||||||
|
|
@ -10,13 +10,21 @@ Pins (PR2):
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import io
|
import io
|
||||||
|
from collections.abc import Callable
|
||||||
|
from contextlib import AbstractContextManager
|
||||||
|
from typing import Final
|
||||||
from unittest.mock import AsyncMock, MagicMock
|
from unittest.mock import AsyncMock, MagicMock
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
import pytest
|
import pytest
|
||||||
|
from fastapi.testclient import TestClient
|
||||||
|
|
||||||
|
import litellm
|
||||||
|
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||||
from litellm.proxy import proxy_server
|
from litellm.proxy import proxy_server
|
||||||
|
from litellm.types.guardrails import GuardrailEventHooks
|
||||||
from litellm.types.llms.openai import HttpxBinaryResponseContent
|
from litellm.types.llms.openai import HttpxBinaryResponseContent
|
||||||
|
from litellm.types.proxy.policy_engine.pipeline_types import GuardrailPipeline, PipelineStep
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@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.headers["content-type"].startswith("text/event-stream")
|
||||||
assert response.text == 'data: {"type":"transcript.text.done","text":"hello world"}\n\n'
|
assert response.text == 'data: {"type":"transcript.text.done","text":"hello world"}\n\n'
|
||||||
assert patched_transcription_stream.closed is True
|
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