mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-28 01:32:17 +00:00
Merge 4dc9572a22 into e26a6450c8
This commit is contained in:
commit
cd8ea40a4b
2 changed files with 61 additions and 5 deletions
|
|
@ -117,6 +117,12 @@ def _strip_total_tokens_from_anthropic_response(response: Any) -> None:
|
|||
usage.pop("total_tokens", None)
|
||||
|
||||
|
||||
def _create_base_llm_response_processor(
|
||||
data: dict, # mutable-ok: existing processor contract owns mutable request data
|
||||
) -> ProxyBaseLLMRequestProcessing:
|
||||
return ProxyBaseLLMRequestProcessing(data=data)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/v1/messages",
|
||||
tags=["[beta] Anthropic `/v1/messages`"],
|
||||
|
|
@ -146,7 +152,7 @@ async def anthropic_response(
|
|||
)
|
||||
|
||||
data: Final = await _read_request_body(request=request)
|
||||
base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data)
|
||||
base_llm_response_processor: Final = _create_base_llm_response_processor(data=data)
|
||||
try:
|
||||
result: Final = await base_llm_response_processor.base_process_llm_request(
|
||||
request=request,
|
||||
|
|
@ -208,12 +214,14 @@ async def anthropic_response(
|
|||
)
|
||||
|
||||
if data.get("stream", None) is not None and data["stream"] is True:
|
||||
# For streaming, use the standard SSE data generator
|
||||
async def _passthrough_stream_generator():
|
||||
yield _anthropic_response
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.fake_stream_iterator import (
|
||||
FakeAnthropicMessagesStreamIterator,
|
||||
)
|
||||
|
||||
blocked_stream: Final = FakeAnthropicMessagesStreamIterator(response=_anthropic_response)
|
||||
|
||||
selected_data_generator: Final = ProxyBaseLLMRequestProcessing.async_sse_data_generator(
|
||||
response=_passthrough_stream_generator(),
|
||||
response=blocked_stream,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_data=_data,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
|
|
|
|||
|
|
@ -126,6 +126,54 @@ class TestBlockedResponseUsage:
|
|||
assert response["usage"] == {"input_tokens": 12, "output_tokens": 5}
|
||||
mock_logging.post_call_failure_hook.assert_awaited_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_blocked_streaming_response_uses_anthropic_sse_events(self):
|
||||
"""A streaming guardrail block must emit valid Anthropic SSE frames."""
|
||||
import litellm.proxy.anthropic_endpoints.endpoints as ep
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
from litellm.integrations.custom_guardrail import ModifyResponseException
|
||||
|
||||
exc = ModifyResponseException(
|
||||
message="blocked by guardrail",
|
||||
model="claude-3-5-sonnet-20240620",
|
||||
request_data={"stream": True, "messages": [{"role": "user", "content": "hi"}]},
|
||||
guardrail_name="rubrik",
|
||||
original_response={"usage": {"input_tokens": 12, "output_tokens": 5}},
|
||||
)
|
||||
|
||||
processor = MagicMock()
|
||||
processor.base_process_llm_request = AsyncMock(side_effect=exc)
|
||||
|
||||
with (
|
||||
patch.object( # test-quality-ok: isolate the endpoint request body for this branch regression
|
||||
ep, "_read_request_body", new=AsyncMock(return_value={"stream": True})
|
||||
),
|
||||
patch.object( # test-quality-ok: inject a processor without mutating its class method
|
||||
ep, "_create_base_llm_response_processor", return_value=processor
|
||||
),
|
||||
patch.object(proxy_server, "proxy_logging_obj") as mock_logging, # test-quality-ok: isolate endpoint logging
|
||||
):
|
||||
mock_logging.post_call_failure_hook = AsyncMock()
|
||||
mock_logging.async_post_call_streaming_iterator_hook.side_effect = lambda **kwargs: kwargs["response"]
|
||||
response = await ep.anthropic_response(
|
||||
fastapi_response=MagicMock(),
|
||||
request=MagicMock(),
|
||||
user_api_key_dict=MagicMock(),
|
||||
)
|
||||
|
||||
frames = [chunk.decode() async for chunk in response.body_iterator]
|
||||
event_types = [frame.splitlines()[0].removeprefix("event: ") for frame in frames]
|
||||
assert event_types == [
|
||||
"message_start",
|
||||
"content_block_start",
|
||||
"content_block_delta",
|
||||
"content_block_stop",
|
||||
"message_delta",
|
||||
"message_stop",
|
||||
]
|
||||
assert all(frame.startswith(f"event: {event}\ndata: ") for frame, event in zip(frames, event_types))
|
||||
assert all(frame.endswith("\n\n") for frame in frames)
|
||||
|
||||
|
||||
class TestProxyExceptionAnthropicEnvelope:
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue