mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(bedrock): preserve stream param and decode SSE for bedrock mantle streaming (#32141)
* fix(bedrock): preserve stream param and decode SSE for bedrock mantle streaming * refactor(bedrock): pass self positionally in mantle messages streaming delegation
This commit is contained in:
parent
733c01902f
commit
90440d75ae
3 changed files with 310 additions and 13 deletions
|
|
@ -7,13 +7,14 @@ The bedrock-mantle endpoint uses the Anthropic Messages API format but is served
|
|||
at a different endpoint (bedrock-mantle.{region}.api.aws) with AWS SigV4 auth.
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any, List, Optional
|
||||
from typing import TYPE_CHECKING, Any, AsyncIterator, Iterator, List, Optional
|
||||
|
||||
from litellm.llms.bedrock.chat.invoke_transformations.anthropic_claude3_transformation import (
|
||||
AmazonAnthropicClaudeConfig,
|
||||
)
|
||||
from litellm.llms.bedrock.common_utils import build_mantle_messages_url
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
|
@ -91,10 +92,14 @@ class AmazonMantleConfig(AmazonAnthropicClaudeConfig):
|
|||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
# The parent strips "model" from the body (Invoke API puts it in URL).
|
||||
# The mantle endpoint (Messages API) requires "model" in the body.
|
||||
request["model"] = model_id
|
||||
return request
|
||||
# The parent strips "model" and "stream" from the body (Invoke API puts
|
||||
# the model in the URL and streams via a dedicated endpoint). The mantle
|
||||
# endpoint (Messages API) requires both in the body.
|
||||
return self._restore_mantle_body_fields(
|
||||
request=request,
|
||||
model_id=model_id,
|
||||
optional_params=optional_params,
|
||||
)
|
||||
|
||||
async def async_transform_request(
|
||||
self,
|
||||
|
|
@ -114,5 +119,31 @@ class AmazonMantleConfig(AmazonAnthropicClaudeConfig):
|
|||
headers=headers,
|
||||
)
|
||||
await self._async_convert_document_url_sources_to_base64(request)
|
||||
request["model"] = model_id
|
||||
return request
|
||||
return self._restore_mantle_body_fields(
|
||||
request=request,
|
||||
model_id=model_id,
|
||||
optional_params=optional_params,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _restore_mantle_body_fields(request: dict, model_id: str, optional_params: dict) -> dict:
|
||||
stream_fields: dict = {"stream": True} if optional_params.get("stream") is True else {}
|
||||
return {**request, "model": model_id, **stream_fields}
|
||||
|
||||
@property
|
||||
def has_custom_stream_wrapper(self) -> bool:
|
||||
return False
|
||||
|
||||
def get_model_response_iterator(
|
||||
self,
|
||||
streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse,
|
||||
sync_stream: bool,
|
||||
json_mode: Optional[bool] = False,
|
||||
) -> Any:
|
||||
from litellm.llms.anthropic.chat.handler import ModelResponseIterator
|
||||
|
||||
return ModelResponseIterator(
|
||||
streaming_response=streaming_response,
|
||||
sync_stream=sync_stream,
|
||||
json_mode=json_mode,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -6,8 +6,13 @@ AmazonAnthropicClaudeMessagesConfig. Overrides only the URL and model-prefix
|
|||
stripping that are specific to the bedrock-mantle endpoint.
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple
|
||||
from typing import TYPE_CHECKING, Any, AsyncIterator, Dict, List, Optional, Tuple
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
|
||||
AnthropicMessagesConfig,
|
||||
)
|
||||
from litellm.llms.bedrock.common_utils import build_mantle_messages_url
|
||||
from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation import (
|
||||
AmazonAnthropicClaudeMessagesConfig,
|
||||
|
|
@ -89,8 +94,26 @@ class AmazonMantleMessagesConfig(AmazonAnthropicClaudeMessagesConfig):
|
|||
headers=headers,
|
||||
)
|
||||
|
||||
# Parent (AmazonAnthropicClaudeMessagesConfig) removes "model" from the
|
||||
# body (Bedrock Invoke puts model in the URL). The mantle endpoint
|
||||
# (Messages API) requires "model" in the request body.
|
||||
request["model"] = model_id
|
||||
return request
|
||||
# Parent (AmazonAnthropicClaudeMessagesConfig) removes "model" and
|
||||
# "stream" from the body (Bedrock Invoke puts the model in the URL and
|
||||
# streams via a dedicated endpoint). The mantle endpoint (Messages API)
|
||||
# requires both in the request body.
|
||||
stream_fields: dict[str, bool] = (
|
||||
{"stream": True} if anthropic_messages_optional_request_params.get("stream") is True else {}
|
||||
)
|
||||
return {**request, "model": model_id, **stream_fields}
|
||||
|
||||
def get_async_streaming_response_iterator(
|
||||
self,
|
||||
model: str,
|
||||
httpx_response: httpx.Response,
|
||||
request_body: dict,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
) -> AsyncIterator:
|
||||
return AnthropicMessagesConfig.get_async_streaming_response_iterator(
|
||||
self,
|
||||
model=model,
|
||||
httpx_response=httpx_response,
|
||||
request_body=request_body,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -204,6 +204,88 @@ def test_mantle_transform_request_strips_prefix_and_adds_model():
|
|||
)
|
||||
assert request["model"] == "anthropic.claude-mythos-preview"
|
||||
assert "mantle/" not in request["model"]
|
||||
assert "stream" not in request
|
||||
|
||||
|
||||
def test_mantle_transform_request_keeps_stream_in_body():
|
||||
config = AmazonMantleConfig()
|
||||
request = config.transform_request(
|
||||
model="mantle/anthropic.claude-mythos-preview",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
optional_params={"max_tokens": 100, "stream": True},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
assert request["stream"] is True
|
||||
assert request["model"] == "anthropic.claude-mythos-preview"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mantle_async_transform_request_keeps_stream_in_body():
|
||||
config = AmazonMantleConfig()
|
||||
request = await config.async_transform_request(
|
||||
model="mantle/anthropic.claude-mythos-preview",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
optional_params={"max_tokens": 100, "stream": True},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
assert request["stream"] is True
|
||||
assert request["model"] == "anthropic.claude-mythos-preview"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mantle_async_transform_request_omits_stream_when_not_streaming():
|
||||
config = AmazonMantleConfig()
|
||||
request = await config.async_transform_request(
|
||||
model="mantle/anthropic.claude-mythos-preview",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
optional_params={"max_tokens": 100},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
assert "stream" not in request
|
||||
|
||||
|
||||
def test_mantle_messages_transform_request_keeps_stream_in_body():
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
config = AmazonMantleMessagesConfig()
|
||||
request = config.transform_anthropic_messages_request(
|
||||
model="mantle/anthropic.claude-mythos-preview",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
anthropic_messages_optional_request_params={"max_tokens": 100, "stream": True},
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
headers={},
|
||||
)
|
||||
assert request["stream"] is True
|
||||
assert request["model"] == "anthropic.claude-mythos-preview"
|
||||
|
||||
|
||||
def test_mantle_messages_transform_request_omits_stream_when_not_streaming():
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
config = AmazonMantleMessagesConfig()
|
||||
request = config.transform_anthropic_messages_request(
|
||||
model="mantle/anthropic.claude-mythos-preview",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
anthropic_messages_optional_request_params={"max_tokens": 100},
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
headers={},
|
||||
)
|
||||
assert "stream" not in request
|
||||
|
||||
|
||||
def test_mantle_chat_streaming_uses_anthropic_sse_iterator():
|
||||
from litellm.llms.anthropic.chat.handler import ModelResponseIterator
|
||||
|
||||
config = AmazonMantleConfig()
|
||||
assert config.has_custom_stream_wrapper is False
|
||||
iterator = config.get_model_response_iterator(
|
||||
streaming_response=iter([]),
|
||||
sync_stream=True,
|
||||
)
|
||||
assert isinstance(iterator, ModelResponseIterator)
|
||||
|
||||
|
||||
def test_mantle_validate_environment_sets_workspace_header():
|
||||
|
|
@ -347,3 +429,164 @@ async def test_mantle_anthropic_messages_routes_to_vpc_api_base():
|
|||
assert len(urls) == 1
|
||||
assert urls[0] == f"{_VPC_ENDPOINT}/anthropic/v1/messages"
|
||||
assert "api.aws" not in urls[0]
|
||||
|
||||
|
||||
_ANTHROPIC_SSE_EVENTS = (
|
||||
(
|
||||
"message_start",
|
||||
{
|
||||
"type": "message_start",
|
||||
"message": {
|
||||
"id": "msg_stream_test",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "anthropic.claude-mythos-preview",
|
||||
"content": [],
|
||||
"stop_reason": None,
|
||||
"stop_sequence": None,
|
||||
"usage": {"input_tokens": 10, "output_tokens": 1},
|
||||
},
|
||||
},
|
||||
),
|
||||
(
|
||||
"content_block_start",
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": 0,
|
||||
"content_block": {"type": "text", "text": ""},
|
||||
},
|
||||
),
|
||||
(
|
||||
"content_block_delta",
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": "pong"},
|
||||
},
|
||||
),
|
||||
("content_block_stop", {"type": "content_block_stop", "index": 0}),
|
||||
(
|
||||
"message_delta",
|
||||
{
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": "end_turn", "stop_sequence": None},
|
||||
"usage": {"output_tokens": 2},
|
||||
},
|
||||
),
|
||||
("message_stop", {"type": "message_stop"}),
|
||||
)
|
||||
|
||||
|
||||
def _anthropic_sse_bytes() -> bytes:
|
||||
return "".join(
|
||||
f"event: {event}\ndata: {json.dumps(payload)}\n\n"
|
||||
for event, payload in _ANTHROPIC_SSE_EVENTS
|
||||
).encode()
|
||||
|
||||
|
||||
def _anthropic_sse_response(url: str) -> httpx.Response:
|
||||
return httpx.Response(
|
||||
status_code=200,
|
||||
content=_anthropic_sse_bytes(),
|
||||
headers={"content-type": "text/event-stream"},
|
||||
request=httpx.Request("POST", url),
|
||||
)
|
||||
|
||||
|
||||
def test_mantle_completion_streaming_sends_stream_and_decodes_sse():
|
||||
import litellm
|
||||
|
||||
requests = []
|
||||
|
||||
def mock_post(self, url, data=None, headers=None, **kwargs):
|
||||
requests.append(_capture_request(url=url, headers=headers or {}, data=data))
|
||||
return _anthropic_sse_response(url)
|
||||
|
||||
with patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post", mock_post):
|
||||
response = litellm.completion(
|
||||
model="bedrock/mantle/anthropic.claude-mythos-preview",
|
||||
messages=[{"role": "user", "content": "ping"}],
|
||||
max_tokens=10,
|
||||
stream=True,
|
||||
aws_access_key_id="fake-key",
|
||||
aws_secret_access_key="fake-secret",
|
||||
aws_region_name="us-east-1",
|
||||
)
|
||||
chunks = list(response)
|
||||
|
||||
assert len(requests) == 1
|
||||
assert requests[0]["body"]["stream"] is True
|
||||
content = "".join(chunk.choices[0].delta.content or "" for chunk in chunks)
|
||||
assert content == "pong"
|
||||
assert chunks[-1].choices[0].finish_reason == "stop"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mantle_acompletion_streaming_sends_stream_and_decodes_sse():
|
||||
import litellm
|
||||
|
||||
requests = []
|
||||
|
||||
async def mock_post(self, url, data=None, headers=None, **kwargs):
|
||||
requests.append(_capture_request(url=url, headers=headers or {}, data=data))
|
||||
return _anthropic_sse_response(url)
|
||||
|
||||
try:
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
new=mock_post,
|
||||
):
|
||||
response = await litellm.acompletion(
|
||||
model="bedrock/mantle/anthropic.claude-mythos-preview",
|
||||
messages=[{"role": "user", "content": "ping"}],
|
||||
max_tokens=10,
|
||||
stream=True,
|
||||
aws_access_key_id="fake-key",
|
||||
aws_secret_access_key="fake-secret",
|
||||
aws_region_name="us-east-1",
|
||||
)
|
||||
chunks = [chunk async for chunk in response]
|
||||
finally:
|
||||
await litellm.close_litellm_async_clients()
|
||||
|
||||
assert len(requests) == 1
|
||||
assert requests[0]["body"]["stream"] is True
|
||||
content = "".join(chunk.choices[0].delta.content or "" for chunk in chunks)
|
||||
assert content == "pong"
|
||||
assert chunks[-1].choices[0].finish_reason == "stop"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mantle_anthropic_messages_streaming_sends_stream_and_passes_through_sse():
|
||||
import litellm
|
||||
|
||||
requests = []
|
||||
|
||||
async def mock_post(self, url, data=None, headers=None, **kwargs):
|
||||
requests.append(_capture_request(url=url, headers=headers or {}, data=data))
|
||||
return _anthropic_sse_response(str(url))
|
||||
|
||||
try:
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
new=mock_post,
|
||||
):
|
||||
response = await litellm.anthropic_messages(
|
||||
model="bedrock/mantle/anthropic.claude-mythos-preview",
|
||||
messages=[{"role": "user", "content": "ping"}],
|
||||
max_tokens=10,
|
||||
stream=True,
|
||||
aws_access_key_id="fake-key",
|
||||
aws_secret_access_key="fake-secret",
|
||||
aws_region_name="us-east-1",
|
||||
)
|
||||
raw = b"".join([chunk async for chunk in response])
|
||||
finally:
|
||||
await litellm.close_litellm_async_clients()
|
||||
|
||||
assert len(requests) == 1
|
||||
assert requests[0]["body"]["stream"] is True
|
||||
text = raw.decode()
|
||||
assert "event: message_start" in text
|
||||
assert '"text": "pong"' in text
|
||||
assert "event: message_stop" in text
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue