mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge branch 'main' into fix/bedrock-converse-redacted-thinking-replay-43009
This commit is contained in:
commit
4bee9838f1
12 changed files with 234 additions and 107 deletions
|
|
@ -57,7 +57,7 @@
|
|||
"mcp-servers-2025-12-04": null,
|
||||
"output-128k-2025-02-19": null,
|
||||
"structured-output-2024-03-01": null,
|
||||
"per-turn-control-2026-07-01": null,
|
||||
"per-turn-control-2026-07-01": "per-turn-control-2026-07-01",
|
||||
"prompt-caching-scope-2026-01-05": "prompt-caching-scope-2026-01-05",
|
||||
"skills-2025-10-02": "skills-2025-10-02",
|
||||
"structured-outputs-2025-11-13": "structured-outputs-2025-11-13",
|
||||
|
|
|
|||
|
|
@ -954,7 +954,7 @@ def _map_bedrock_exception(
|
|||
llm_provider="bedrock",
|
||||
response=getattr(original_exception, "response", None),
|
||||
)
|
||||
elif "Could not process image" in error_str:
|
||||
elif "Could not process image" in error_str and getattr(original_exception, "status_code", 500) == 500:
|
||||
raise litellm.InternalServerError(
|
||||
message=f"BedrockException - {error_str}",
|
||||
model=model,
|
||||
|
|
|
|||
|
|
@ -685,7 +685,7 @@ class ChunkProcessor:
|
|||
|
||||
def _flush_thinking_block() -> None:
|
||||
nonlocal current_thinking_text_parts, current_signature
|
||||
if len(current_thinking_text_parts) > 0 and current_signature:
|
||||
if current_signature:
|
||||
thinking_blocks.append(
|
||||
ChatCompletionThinkingBlock(
|
||||
type="thinking",
|
||||
|
|
|
|||
|
|
@ -109,8 +109,6 @@ class VertexAIPartnerModels(VertexBase):
|
|||
client=None,
|
||||
):
|
||||
try:
|
||||
import vertexai
|
||||
|
||||
from litellm.llms.anthropic.chat import AnthropicChatCompletion
|
||||
from litellm.llms.codestral.completion.handler import (
|
||||
CodestralTextCompletion,
|
||||
|
|
@ -119,14 +117,9 @@ class VertexAIPartnerModels(VertexBase):
|
|||
except Exception as e:
|
||||
raise VertexAIError(
|
||||
status_code=400,
|
||||
message=f"""vertexai import failed please run `pip install -U "google-cloud-aiplatform>=1.38"`. Got error: {e}""",
|
||||
message=f"Failed to import a partner model handler. Got error: {e}",
|
||||
)
|
||||
|
||||
if not (hasattr(vertexai, "preview") or hasattr(vertexai.preview, "language_models")):
|
||||
raise VertexAIError(
|
||||
status_code=400,
|
||||
message="""Upgrade vertex ai. Run `pip install "google-cloud-aiplatform>=1.38"`""",
|
||||
)
|
||||
try:
|
||||
access_token, project_id = self._ensure_access_token(
|
||||
credentials=vertex_credentials,
|
||||
|
|
|
|||
|
|
@ -73,6 +73,11 @@ def _output_items_with_id(items: tuple[Any, ...], item_type: str, item_id: str |
|
|||
)
|
||||
|
||||
|
||||
def _delta_has_signed_thinking_block(delta: object) -> bool:
|
||||
blocks: Final = getattr(delta, "thinking_blocks", None) or ()
|
||||
return any(isinstance(b, dict) and (b.get("signature") or b.get("data")) for b in blocks)
|
||||
|
||||
|
||||
class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
||||
"""
|
||||
Async iterator for processing streaming responses from the Responses API.
|
||||
|
|
@ -936,7 +941,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
self.sent_output_item_added_event = True
|
||||
|
||||
# Reasoning-first
|
||||
if hasattr(delta, "reasoning_content") and delta.reasoning_content:
|
||||
if (hasattr(delta, "reasoning_content") and delta.reasoning_content) or _delta_has_signed_thinking_block(delta):
|
||||
self._reasoning_active = True
|
||||
if self._cached_reasoning_item_id is None:
|
||||
self._cached_reasoning_item_id = f"rs_{uuid.uuid4()}"
|
||||
|
|
|
|||
|
|
@ -1,93 +0,0 @@
|
|||
import json
|
||||
from collections.abc import AsyncIterator
|
||||
from types import SimpleNamespace
|
||||
from typing import Any, cast
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
from litellm.main import vertex_gemma_chat_completion
|
||||
from litellm.types.llms.openai import OutputTextDeltaEvent, ResponseCompletedEvent, ResponsesAPIStreamingResponse
|
||||
|
||||
_VERTEX_URL = "https://example.invalid/v1/projects/test/locations/us-central1/endpoints/test:predict"
|
||||
_MESSAGES = [{"role": "user", "content": "Reply exactly READY"}]
|
||||
_FAKE_CREDENTIALS = "gemma-test-credentials"
|
||||
|
||||
|
||||
def _vertex_response():
|
||||
return {
|
||||
"predictions": {
|
||||
"id": "chatcmpl-stream-test",
|
||||
"created": 1759863903,
|
||||
"model": "google/gemma-3-12b-it",
|
||||
"object": "chat.completion",
|
||||
"choices": [{"index": 0, "finish_reason": "stop", "message": {"role": "assistant", "content": "READY"}}],
|
||||
"usage": {"prompt_tokens": 14, "completion_tokens": 1, "total_tokens": 15},
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _cached_access_token():
|
||||
"""Serve a fake token from the handler's credential cache so no auth round-trip runs."""
|
||||
cache = vertex_gemma_chat_completion._credentials_project_mapping
|
||||
key = (_FAKE_CREDENTIALS, "test")
|
||||
cache[key] = (SimpleNamespace(token="fake-token", expired=False), "test")
|
||||
yield
|
||||
cache.pop(key, None)
|
||||
|
||||
|
||||
def test_sync_gemma_stream():
|
||||
captured: dict[str, Any] = {}
|
||||
|
||||
def handle(request: httpx.Request) -> httpx.Response:
|
||||
captured["body"] = json.loads(request.content)
|
||||
return httpx.Response(200, json=_vertex_response())
|
||||
|
||||
stream = litellm.completion(
|
||||
model="vertex_ai/gemma/test-model",
|
||||
messages=_MESSAGES,
|
||||
stream=True,
|
||||
api_base=_VERTEX_URL,
|
||||
vertex_project="test",
|
||||
vertex_location="us-central1",
|
||||
vertex_credentials=_FAKE_CREDENTIALS,
|
||||
client=httpx.Client(transport=httpx.MockTransport(handle)),
|
||||
)
|
||||
|
||||
assert isinstance(stream, CustomStreamWrapper)
|
||||
chunks = list(stream)
|
||||
|
||||
assert "stream" not in captured["body"]["instances"][0]
|
||||
assert len(chunks) == 2
|
||||
assert chunks[0].choices[0].delta.content == "READY"
|
||||
assert chunks[1].choices[0].finish_reason == "stop"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_gemma_responses_stream():
|
||||
captured: dict[str, Any] = {}
|
||||
|
||||
def handle(request: httpx.Request) -> httpx.Response:
|
||||
captured["body"] = json.loads(request.content)
|
||||
return httpx.Response(200, json=_vertex_response())
|
||||
|
||||
response = await litellm.aresponses(
|
||||
model="vertex_ai/gemma/test-model",
|
||||
input="Reply exactly READY",
|
||||
stream=True,
|
||||
api_base=_VERTEX_URL,
|
||||
vertex_project="test",
|
||||
vertex_location="us-central1",
|
||||
vertex_credentials=_FAKE_CREDENTIALS,
|
||||
client=httpx.AsyncClient(transport=httpx.MockTransport(handle)),
|
||||
)
|
||||
events = [event async for event in cast(AsyncIterator[ResponsesAPIStreamingResponse], response)]
|
||||
|
||||
assert "stream" not in captured["body"]["instances"][0]
|
||||
assert "READY" in "".join(event.delta for event in events if isinstance(event, OutputTextDeltaEvent))
|
||||
assert isinstance(events[-1], ResponseCompletedEvent)
|
||||
assert events[-1].response.usage is not None
|
||||
assert events[-1].response.usage.total_tokens == 15
|
||||
|
|
@ -1280,7 +1280,7 @@ def test_bedrock_500_preserves_provider_response_headers():
|
|||
"bedrock",
|
||||
400,
|
||||
'{"message":"Could not process image"}',
|
||||
litellm.InternalServerError,
|
||||
litellm.BadRequestError,
|
||||
),
|
||||
],
|
||||
)
|
||||
|
|
@ -1313,6 +1313,41 @@ def test_bedrock_classified_errors_preserve_provider_response_headers(
|
|||
assert exc_info.value.response.headers["x-amzn-requestid"] == "req-classified"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"status_code, expected_exception",
|
||||
[
|
||||
(400, litellm.BadRequestError),
|
||||
(503, litellm.ServiceUnavailableError),
|
||||
(500, litellm.InternalServerError),
|
||||
],
|
||||
)
|
||||
def test_bedrock_unprocessable_image_keeps_provider_status_code(status_code, expected_exception):
|
||||
"""An unprocessable image maps to the status Bedrock sent, so the 400 it returns stays a client error."""
|
||||
provider_message = '{"message":"The model returned the following errors: Could not process image"}'
|
||||
provider_response = httpx.Response(
|
||||
status_code=status_code,
|
||||
text=provider_message,
|
||||
request=httpx.Request("POST", "https://bedrock-runtime.us-east-1.amazonaws.com/"),
|
||||
)
|
||||
original_exception = BedrockError(
|
||||
status_code=status_code,
|
||||
message=provider_message,
|
||||
headers=provider_response.headers,
|
||||
response=provider_response,
|
||||
)
|
||||
|
||||
with pytest.raises(expected_exception) as exc_info:
|
||||
exception_type(
|
||||
model="anthropic.claude-haiku-4-5-20251001-v1:0",
|
||||
original_exception=original_exception,
|
||||
custom_llm_provider="bedrock",
|
||||
completion_kwargs={},
|
||||
extra_kwargs={},
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == status_code
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"status_code, provider_message",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -236,6 +236,31 @@ def test_get_combined_thinking_content_preserves_interleaved_blocks():
|
|||
assert result[2]["signature"] == "sig_block2"
|
||||
|
||||
|
||||
def test_get_combined_thinking_content_keeps_signed_block_without_thinking_text():
|
||||
chunks: Final = [
|
||||
ModelResponseStream(
|
||||
id="chatcmpl-123",
|
||||
object="chat.completion.chunk",
|
||||
created=1234567890,
|
||||
model="claude-sonnet-4-20250514",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
index=0,
|
||||
delta=Delta(thinking_blocks=[{"type": "thinking", "thinking": "", "signature": "sig_only"}]),
|
||||
finish_reason=None,
|
||||
)
|
||||
],
|
||||
)
|
||||
]
|
||||
|
||||
result: Final = ChunkProcessor(chunks=chunks).get_combined_thinking_content(chunks)
|
||||
|
||||
assert result is not None
|
||||
assert [(block["type"], block["thinking"], block["signature"]) for block in result] == [
|
||||
("thinking", "", "sig_only")
|
||||
]
|
||||
|
||||
|
||||
def test_cache_read_input_tokens_retained():
|
||||
chunk1 = ModelResponseStream(
|
||||
id="chatcmpl-95aabb85-c39f-443d-ae96-0370c404d70c",
|
||||
|
|
|
|||
|
|
@ -95,13 +95,19 @@ def test_added_per_turn_control_beta_survives_the_anthropic_allowlist():
|
|||
assert PER_TURN_CONTROL in _betas(filtered)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("provider", ["bedrock", "bedrock_converse", "vertex_ai", "azure_ai", "databricks"])
|
||||
@pytest.mark.parametrize("provider", ["bedrock", "bedrock_converse", "vertex_ai", "databricks"])
|
||||
def test_per_turn_control_beta_is_dropped_for_providers_without_it(provider):
|
||||
filtered = update_headers_with_filtered_beta(headers={"anthropic-beta": PER_TURN_CONTROL}, provider=provider)
|
||||
|
||||
assert "anthropic-beta" not in filtered
|
||||
|
||||
|
||||
def test_per_turn_control_beta_is_forwarded_for_azure_ai():
|
||||
filtered = update_headers_with_filtered_beta(headers={"anthropic-beta": PER_TURN_CONTROL}, provider="azure_ai")
|
||||
|
||||
assert _betas(filtered) == {PER_TURN_CONTROL}
|
||||
|
||||
|
||||
def test_json_provider_passthrough_adds_per_turn_control_beta():
|
||||
config = JSONProviderAnthropicMessagesConfig(
|
||||
SimpleProviderConfig(
|
||||
|
|
|
|||
|
|
@ -127,6 +127,44 @@ class TestPartnerModelsCredentialReuse:
|
|||
|
||||
assert mock_load.call_count == 1
|
||||
|
||||
def test_completion_works_without_the_vertexai_sdk(self):
|
||||
"""completion() reaches the HTTP handler when `import vertexai` raises ImportError."""
|
||||
partner = VertexAIPartnerModels()
|
||||
|
||||
with (
|
||||
patch.dict(sys.modules, {"vertexai": None}),
|
||||
patch.object(
|
||||
partner,
|
||||
"_ensure_access_token",
|
||||
return_value=("cached-token", "test-project"),
|
||||
),
|
||||
patch(
|
||||
"litellm.llms.vertex_ai.vertex_ai_partner_models.main.base_llm_http_handler"
|
||||
) as mock_handler,
|
||||
):
|
||||
mock_handler.completion.return_value = "response"
|
||||
|
||||
result = partner.completion(
|
||||
model="meta/llama-3.1-405b-instruct-maas",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
model_response=MagicMock(),
|
||||
print_verbose=lambda *a, **kw: None,
|
||||
encoding=MagicMock(),
|
||||
logging_obj=MagicMock(),
|
||||
api_base=None,
|
||||
optional_params={},
|
||||
custom_prompt_dict={},
|
||||
headers=None,
|
||||
timeout=30.0,
|
||||
litellm_params={},
|
||||
vertex_project="test-project",
|
||||
vertex_location="us-central1",
|
||||
vertex_credentials=None,
|
||||
)
|
||||
|
||||
assert result == "response"
|
||||
mock_handler.completion.assert_called_once()
|
||||
|
||||
|
||||
class TestGemmaModelsCredentialReuse:
|
||||
def test_completion_uses_self_ensure_access_token(self):
|
||||
|
|
|
|||
|
|
@ -1303,3 +1303,81 @@ class TestVertexGemmaCompletion:
|
|||
mock_async_post.assert_awaited_once()
|
||||
assert mock_async_post.call_args.kwargs["client"] is None
|
||||
assert response.choices[0].message.content == "default async handler fallback"
|
||||
|
||||
|
||||
_GEMMA_VERTEX_URL = "https://example.invalid/v1/projects/test/locations/us-central1/endpoints/test:predict"
|
||||
_FAKE_GEMMA_CREDENTIALS = "gemma-test-credentials"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def _gemma_cached_access_token():
|
||||
"""Serve a fake token from the handler's credential cache so no auth round-trip runs."""
|
||||
from types import SimpleNamespace
|
||||
|
||||
from litellm.main import vertex_gemma_chat_completion
|
||||
|
||||
cache = vertex_gemma_chat_completion._credentials_project_mapping
|
||||
key = (_FAKE_GEMMA_CREDENTIALS, "test")
|
||||
cache[key] = (SimpleNamespace(token="fake-token", expired=False), "test")
|
||||
yield
|
||||
cache.pop(key, None)
|
||||
|
||||
|
||||
def test_sync_gemma_stream(_gemma_cached_access_token):
|
||||
import httpx
|
||||
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
|
||||
captured = {}
|
||||
|
||||
def handle(request):
|
||||
captured["body"] = json.loads(request.content)
|
||||
return httpx.Response(200, json=_make_gemma_vertex_response(content="READY"))
|
||||
|
||||
stream = litellm.completion(
|
||||
model="vertex_ai/gemma/test-model",
|
||||
messages=[{"role": "user", "content": "Reply exactly READY"}],
|
||||
stream=True,
|
||||
api_base=_GEMMA_VERTEX_URL,
|
||||
vertex_project="test",
|
||||
vertex_location="us-central1",
|
||||
vertex_credentials=_FAKE_GEMMA_CREDENTIALS,
|
||||
client=httpx.Client(transport=httpx.MockTransport(handle)),
|
||||
)
|
||||
|
||||
assert isinstance(stream, CustomStreamWrapper)
|
||||
chunks = list(stream)
|
||||
|
||||
assert "stream" not in captured["body"]["instances"][0]
|
||||
assert len(chunks) == 2
|
||||
assert chunks[0].choices[0].delta.content == "READY"
|
||||
assert chunks[1].choices[0].finish_reason == "stop"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_gemma_responses_stream(_gemma_cached_access_token):
|
||||
import httpx
|
||||
|
||||
captured = {}
|
||||
|
||||
def handle(request):
|
||||
captured["body"] = json.loads(request.content)
|
||||
return httpx.Response(200, json=_make_gemma_vertex_response(content="READY"))
|
||||
|
||||
response = await litellm.aresponses(
|
||||
model="vertex_ai/gemma/test-model",
|
||||
input="Reply exactly READY",
|
||||
stream=True,
|
||||
api_base=_GEMMA_VERTEX_URL,
|
||||
vertex_project="test",
|
||||
vertex_location="us-central1",
|
||||
vertex_credentials=_FAKE_GEMMA_CREDENTIALS,
|
||||
client=httpx.AsyncClient(transport=httpx.MockTransport(handle)),
|
||||
)
|
||||
events = [event async for event in cast(AsyncIterator[ResponsesAPIStreamingResponse], response)]
|
||||
|
||||
assert "stream" not in captured["body"]["instances"][0]
|
||||
assert "READY" in "".join(event.delta for event in events if isinstance(event, OutputTextDeltaEvent))
|
||||
assert isinstance(events[-1], ResponseCompletedEvent)
|
||||
assert events[-1].response.usage is not None
|
||||
assert events[-1].response.usage.total_tokens == 114
|
||||
|
|
|
|||
|
|
@ -978,6 +978,25 @@ def _reasoning_chunk(reasoning: str, finish_reason: str | None = None) -> ModelR
|
|||
)
|
||||
|
||||
|
||||
def _signature_only_thinking_chunk(signature: str) -> ModelResponseStream:
|
||||
return ModelResponseStream(
|
||||
id=CHAT_COMPLETION_ID,
|
||||
created=1748575031,
|
||||
model="claude-haiku-4-5",
|
||||
object="chat.completion.chunk",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
index=0,
|
||||
delta=Delta(
|
||||
role="assistant",
|
||||
thinking_blocks=[{"type": "thinking", "thinking": "", "signature": signature}],
|
||||
),
|
||||
finish_reason=None,
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
async def _collect_events(
|
||||
iterator: LiteLLMCompletionStreamingIterator, sync_mode: bool
|
||||
) -> list[BaseLiteLLMOpenAIResponseObject]:
|
||||
|
|
@ -1015,6 +1034,27 @@ async def test_tool_only_stream_emits_no_message_item_events(sync_mode: bool):
|
|||
assert any(getattr(event, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED for event in events)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
@pytest.mark.asyncio
|
||||
async def test_signature_only_thinking_streams_a_replayable_reasoning_item(sync_mode: bool):
|
||||
iterator: Final = _build_iterator([_signature_only_thinking_chunk("sig_only"), _chunk("4", finish_reason="stop")])
|
||||
|
||||
events: Final = await _collect_events(iterator, sync_mode)
|
||||
|
||||
added_item_types: Final = [
|
||||
event.item.type
|
||||
for event in events
|
||||
if getattr(event, "type", None) == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED
|
||||
]
|
||||
completed: Final = next(
|
||||
event for event in events if getattr(event, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED
|
||||
)
|
||||
reasoning_items: Final = [item for item in completed.response.output if getattr(item, "type", None) == "reasoning"]
|
||||
assert added_item_types[0] == "reasoning"
|
||||
assert len(reasoning_items) == 1
|
||||
assert json.loads(reasoning_items[0].encrypted_content)[0]["signature"] == "sig_only"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
@pytest.mark.asyncio
|
||||
async def test_reasoning_then_text_announces_message_item_before_text_events(sync_mode: bool):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue