Merge remote-tracking branch 'origin/main' into fix-session-limits

This commit is contained in:
AlisinaDevelo 2026-09-27 08:31:22 +02:00
commit 7bdfe7884c
10 changed files with 213 additions and 48 deletions

View file

@ -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",

View file

@ -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,

View file

@ -951,7 +951,7 @@ def _count_content_list(
)
def _format_function_definitions(tools):
def _format_function_definitions(tools: Sequence[object]) -> str:
"""Formats tool definitions in the format that OpenAI appears to use.
Based on https://github.com/forestwanglin/openai-java/blob/main/jtokkit/src/main/java/xyz/felh/openai/jtokkit/utils/TikTokenUtils.java
"""
@ -959,41 +959,57 @@ def _format_function_definitions(tools):
lines.append("namespace functions {")
lines.append("")
for tool in tools:
if not isinstance(tool, dict):
if not isinstance(tool, Mapping):
continue
function = tool.get("function")
if not isinstance(function, dict):
# Anthropic tool shape → OpenAI function dict for token counting.
params = tool.get("input_schema") or tool.get("parameters") or {}
if not isinstance(params, dict):
params = {}
function = {
"name": tool.get("name"),
"description": tool.get("description"),
"parameters": params,
}
function_name = function.get("name")
if not function_name:
# Skip malformed tools missing a name to avoid emitting
# ``type None = ...`` which would produce inaccurate token counts.
continue
if function_description := function.get("description"):
lines.append(f"// {function_description}")
parameters = function.get("parameters") or {}
if not isinstance(parameters, dict):
parameters = {}
properties = parameters.get("properties")
if properties and properties.keys():
lines.append(f"type {function_name} = (_: {{")
lines.append(_format_object_parameters(parameters, 0))
lines.append("}) => any;")
else:
lines.append(f"type {function_name} = () => any;")
lines.append("")
for function in _function_definitions_for_tool(cast(Mapping[str, object], tool)):
lines.extend(_format_single_function_definition(function))
lines.append("} // namespace functions")
return "\n".join(lines)
def _function_definitions_for_tool(tool: Mapping[str, object]) -> Iterable[Mapping[str, object]]:
function: Final = tool.get("function")
if isinstance(function, Mapping):
yield function
return
declarations: Final = tool.get("function_declarations") or tool.get("functionDeclarations")
if isinstance(declarations, list):
for declaration in declarations:
if isinstance(declaration, Mapping):
yield declaration
return
parameters: Final = tool.get("input_schema") or tool.get("parameters") or {}
normalized_parameters: Final = parameters if isinstance(parameters, Mapping) else {}
yield {
"name": tool.get("name"),
"description": tool.get("description"),
"parameters": normalized_parameters,
}
def _format_single_function_definition(function: Mapping[str, object]) -> tuple[str, ...]:
function_name: Final = function.get("name")
if not function_name:
return ()
function_description: Final = function.get("description")
parameters_value: Final = function.get("parameters") or {}
parameters: Final = parameters_value if isinstance(parameters_value, Mapping) else {}
properties: Final = parameters.get("properties")
if isinstance(properties, Mapping) and properties:
return (
*((f"// {function_description}",) if function_description else ()),
f"type {function_name} = (_: {{",
_format_object_parameters(parameters, 0),
"}) => any;",
"",
)
return (
*((f"// {function_description}",) if function_description else ()),
f"type {function_name} = () => any;",
"",
)
def _format_object_parameters(parameters, indent):
properties: Final = parameters.get("properties")
if not properties:

View file

@ -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,

View file

@ -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",
[

View file

@ -442,6 +442,78 @@ def test_token_counter_with_tools(message_count_pair):
), f"Expected {expected_tokens} tokens, got {counted_tokens}."
def test_token_counter_counts_gemini_function_declarations():
openai_tools: Final = [
{
"type": "function",
"function": {
"name": "lookup_weather",
"description": "Find current weather conditions for a location",
"parameters": {
"type": "object",
"properties": {
"location": {"type": "string", "description": "City and region"},
"units": {"type": "string", "enum": ["celsius", "fahrenheit"]},
},
"required": ["location"],
},
},
}
]
gemini_tools: Final = litellm.utils.get_optional_params(
model="gemini-2.5-pro",
custom_llm_provider="gemini",
tools=openai_tools,
)["tools"]
camel_case_tools: Final = [{"functionDeclarations": gemini_tools[0]["function_declarations"]}]
openai_tokens: Final = token_counter_new(
model="gemini-2.5-pro",
messages=[{"role": "user", "content": "What's the weather?"}],
tools=openai_tools,
)
gemini_tokens: Final = token_counter_new(
model="gemini-2.5-pro",
messages=[{"role": "user", "content": "What's the weather?"}],
tools=gemini_tools,
)
camel_case_tokens: Final = token_counter_new(
model="gemini-2.5-pro",
messages=[{"role": "user", "content": "What's the weather?"}],
tools=camel_case_tools,
)
assert openai_tokens == gemini_tokens == camel_case_tokens
def test_token_counter_skips_non_mapping_tools():
openai_tool: Final = {
"type": "function",
"function": {
"name": "lookup_weather",
"description": "Find current weather conditions for a location",
"parameters": {
"type": "object",
"properties": {"location": {"type": "string", "description": "City and region"}},
"required": ["location"],
},
},
}
messages: Final = [{"role": "user", "content": "What's the weather?"}]
valid_tokens: Final = token_counter_new(
model="gemini-2.5-pro",
messages=messages,
tools=[openai_tool],
)
mixed_tokens: Final = token_counter_new(
model="gemini-2.5-pro",
messages=messages,
tools=["bad", None, openai_tool],
)
assert mixed_tokens == valid_tokens
class NeedsToleranceUpdateError(Exception):
"""Custom exception to mark tests that have improved"""

View file

@ -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(

View file

@ -1,4 +1,4 @@
from typing import List
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
@ -1530,7 +1530,7 @@ class TestContextCachingEndpoints:
]
all_messages = short_cached_messages + non_cached_messages
large_tools = [
openai_large_tools: Final = [
{
"type": "function",
"function": {
@ -1548,6 +1548,11 @@ class TestContextCachingEndpoints:
}
for i in range(12)
]
large_tools: Final = litellm.utils.get_optional_params(
model="gemini-1.5-pro",
custom_llm_provider="gemini",
tools=openai_large_tools,
)["tools"]
optional_params = {
**self.sample_optional_params,

View file

@ -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):

View file

@ -1332,7 +1332,7 @@ def test_sync_gemma_stream(_gemma_cached_access_token):
def handle(request):
captured["body"] = json.loads(request.content)
return httpx.Response(200, json=_make_gemma_vertex_response(content="READY", total_tokens=15))
return httpx.Response(200, json=_make_gemma_vertex_response(content="READY"))
stream = litellm.completion(
model="vertex_ai/gemma/test-model",
@ -1362,7 +1362,7 @@ async def test_async_gemma_responses_stream(_gemma_cached_access_token):
def handle(request):
captured["body"] = json.loads(request.content)
return httpx.Response(200, json=_make_gemma_vertex_response(content="READY", total_tokens=15))
return httpx.Response(200, json=_make_gemma_vertex_response(content="READY"))
response = await litellm.aresponses(
model="vertex_ai/gemma/test-model",
@ -1380,4 +1380,4 @@ async def test_async_gemma_responses_stream(_gemma_cached_access_token):
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
assert events[-1].response.usage.total_tokens == 114