fix(bedrock): harden citations, search_results mapping, and token counting
Some checks failed
Unit Tests: Proxy DB Operations / assert-shard-coverage (push) Has been cancelled
Unit Tests: Security / security (push) Has been cancelled
Unit Tests: Proxy DB Operations / auth-checks (push) Has been cancelled
Unit Tests: Proxy DB Operations / budgets (push) Has been cancelled
Unit Tests: Proxy DB Operations / custom-logging (push) Has been cancelled
Unit Tests: Proxy DB Operations / db-and-spend (push) Has been cancelled
Unit Tests: Proxy DB Operations / endpoints-and-responses (push) Has been cancelled
Unit Tests: Proxy DB Operations / guardrails-hooks (push) Has been cancelled
Unit Tests: Proxy DB Operations / jwt-and-keys (push) Has been cancelled
Unit Tests: Proxy DB Operations / key-generation (push) Has been cancelled
Unit Tests: Proxy DB Operations / logging-misc (push) Has been cancelled
Unit Tests: Proxy DB Operations / proxy-runtime (push) Has been cancelled
Unit Tests: Proxy DB Operations / proxy-server-core (push) Has been cancelled
Unit Tests: Proxy DB Operations / schema-migration (push) Has been cancelled
Unit Tests: Proxy DB Operations / proxy-utils (push) Has been cancelled

Resolve mypy issues in citation parsing, only attach url_citation annotations when citation text is stitched into content, fall back to tool content when search_results is empty, and count search_results text in token/TPM preflight paths.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Sameer Kankute 2026-05-28 22:25:42 +05:30
parent 65f8e21ebd
commit 0eb3658e9b
No known key found for this signature in database
5 changed files with 220 additions and 16 deletions

View file

@ -132,6 +132,30 @@ def strip_none_values_from_message(message: AllMessageValues) -> AllMessageValue
return cast(AllMessageValues, {k: v for k, v in message.items() if v is not None})
def extract_search_results_text(search_results: object) -> str:
"""
Extract model-visible text from OpenAI tool-message ``search_results``.
Used by token estimators and TPM limiters so large search result payloads
cannot bypass preflight checks via a small ``content`` field.
"""
if not isinstance(search_results, list):
return ""
texts = ""
for result in search_results:
if not isinstance(result, dict):
continue
content = result.get("content")
if not isinstance(content, list):
continue
for block in content:
if isinstance(block, dict):
text = block.get("text")
if isinstance(text, str):
texts += text
return texts
def convert_content_list_to_str(
message: Union[AllMessageValues, ChatCompletionResponseMessage],
) -> str:
@ -152,6 +176,7 @@ def convert_content_list_to_str(
elif message_content is not None and isinstance(message_content, str):
texts = message_content
texts += extract_search_results_text(message.get("search_results"))
return texts

View file

@ -3658,6 +3658,7 @@ from litellm.types.llms.bedrock import (
ToolInputSchemaBlock as BedrockToolInputSchemaBlock,
)
from litellm.types.llms.bedrock import ToolJsonSchemaBlock as BedrockToolJsonSchemaBlock
from litellm.types.llms.bedrock import SearchResultBlock
from litellm.types.llms.bedrock import ToolResultBlock as BedrockToolResultBlock
from litellm.types.llms.bedrock import (
ToolResultContentBlock as BedrockToolResultContentBlock,
@ -4115,14 +4116,19 @@ def _convert_to_bedrock_tool_call_result(
# If `search_results` is present, we intentionally prefer it over `content`
# to avoid generating mixed text + searchResult blocks.
search_results = message.get("search_results")
used_search_results = False
if isinstance(search_results, list):
for result in search_results:
if not isinstance(result, dict):
continue
tool_result_content_blocks.append(
BedrockToolResultContentBlock(searchResult=result) # type: ignore[arg-type]
BedrockToolResultContentBlock(
searchResult=cast(SearchResultBlock, result)
)
)
else:
used_search_results = len(tool_result_content_blocks) > 0
if not used_search_results:
if isinstance(message["content"], str):
tool_result_content_blocks.append(
BedrockToolResultContentBlock(text=message["content"])
@ -4209,7 +4215,7 @@ def _convert_to_bedrock_tool_call_result(
tool_result = BedrockToolResultBlock(
content=tool_result_content_blocks, toolUseId=id
)
if isinstance(search_results, list):
if used_search_results:
tool_result["status"] = cast(Literal["success"], "success")
content_block = BedrockContentBlock(toolResult=tool_result)

View file

@ -486,6 +486,14 @@ def _count_messages(
use_default_image_token_count,
default_token_count,
)
elif key == "search_results" and isinstance(value, list):
from litellm.litellm_core_utils.prompt_templates.common_utils import (
extract_search_results_text,
)
search_results_text = extract_search_results_text(value)
if search_results_text:
num_tokens += params.count_function(search_results_text)
else:
# Skip unsupported keys instead of raising an error
continue

View file

@ -1911,25 +1911,33 @@ class AmazonConverseConfig(BaseConfig):
for citations_block in citations_content_blocks:
block_text = ""
for content_part in citations_block.get("content", []):
if isinstance(content_part, dict):
_text = content_part.get("text")
if isinstance(_text, str):
block_text += _text
raw_content = citations_block.get("content")
if isinstance(raw_content, list):
for content_part in raw_content:
if isinstance(content_part, dict):
_text = content_part.get("text")
if isinstance(_text, str):
block_text += _text
if block_text:
citations_text_parts.append(block_text)
for citation in citations_block.get("citations", []):
raw_citations = citations_block.get("citations")
if not isinstance(raw_citations, list):
continue
for citation in raw_citations:
if not isinstance(citation, dict):
continue
location = citation.get("location", {})
search_location = (
location.get("searchResultLocation", {})
if isinstance(location, dict)
else {}
)
location = citation.get("location")
if not isinstance(location, dict):
continue
search_location = location.get("searchResultLocation")
if not isinstance(search_location, dict):
continue
start = search_location.get("start")
end = search_location.get("end")
if not isinstance(start, int) or not isinstance(end, int):
@ -2137,14 +2145,17 @@ class AmazonConverseConfig(BaseConfig):
citations_text, annotations = self._transform_citations_to_annotations(
citationsContentBlocks
)
citations_included_in_content = False
if citations_text:
if not content_str:
content_str = citations_text
citations_included_in_content = True
elif content_str.strip() == ".":
# Bedrock may emit the cited sentence in citationsContent and only
# punctuation in text blocks; stitch them for user-facing content.
content_str = citations_text + content_str
if annotations:
citations_included_in_content = True
if annotations and citations_included_in_content:
chat_completion_message["annotations"] = annotations
if reasoningContentBlocks is not None:

View file

@ -1714,6 +1714,160 @@ async def test_tool_message_search_results_maps_to_bedrock_search_result_block()
)
@pytest.mark.asyncio
async def test_tool_message_empty_search_results_falls_back_to_content():
"""Empty search_results must not skip normal tool content processing."""
from litellm.litellm_core_utils.prompt_templates.factory import (
BedrockConverseMessagesProcessor,
_bedrock_converse_messages_pt,
)
messages = [
{"role": "user", "content": "hello"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "tooluse_empty_search",
"type": "function",
"function": {"name": "lookup", "arguments": "{}"},
}
],
},
{
"role": "tool",
"tool_call_id": "tooluse_empty_search",
"content": "fallback tool text",
"search_results": [],
},
]
result = _bedrock_converse_messages_pt(
messages=messages,
model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
llm_provider="bedrock_converse",
)
async_result = (
await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async(
messages=messages,
model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
llm_provider="bedrock_converse",
)
)
assert result == async_result
tool_result = result[2]["content"][0]["toolResult"]
assert tool_result["toolUseId"] == "tooluse_empty_search"
assert "status" not in tool_result
assert len(tool_result["content"]) == 1
assert tool_result["content"][0]["text"] == "fallback tool text"
def test_transform_response_omits_annotations_when_citations_not_stitched():
from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig
from litellm.types.utils import ModelResponse
response_json = {
"metrics": {"latencyMs": 100},
"output": {
"message": {
"role": "assistant",
"content": [
{
"citationsContent": {
"content": [{"text": "cited sentence only in citations"}],
"citations": [
{
"location": {
"searchResultLocation": {
"start": 0,
"end": 5,
}
},
"source": "https://example.com",
"title": "Example",
}
],
}
},
{"text": "separate assistant answer"},
],
}
},
"stopReason": "end_turn",
"usage": {
"inputTokens": 10,
"outputTokens": 5,
"totalTokens": 15,
"cacheReadInputTokenCount": 0,
"cacheReadInputTokens": 0,
"cacheWriteInputTokenCount": 0,
"cacheWriteInputTokens": 0,
},
}
class MockResponse:
def json(self):
return response_json
@property
def text(self):
return json.dumps(response_json)
config = AmazonConverseConfig()
result = config._transform_response(
model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
response=MockResponse(),
model_response=ModelResponse(),
stream=False,
logging_obj=None,
optional_params={},
api_key=None,
data=None,
messages=[],
encoding=None,
)
message = result.choices[0].message
assert message.content == "separate assistant answer"
assert message.model_dump().get("annotations") is None
def test_extract_search_results_text_counts_hidden_tool_payload():
from litellm.litellm_core_utils.prompt_templates.common_utils import (
convert_content_list_to_str,
extract_search_results_text,
)
from litellm.litellm_core_utils.token_counter import token_counter
hidden = "x" * 500
message = {
"role": "tool",
"content": "small",
"search_results": [
{
"source": "s",
"title": "t",
"content": [{"text": hidden}],
}
],
}
assert extract_search_results_text(message["search_results"]) == hidden
assert len(convert_content_list_to_str(message)) > len("small")
tokens_with_search = token_counter(
model="gpt-3.5-turbo",
messages=[message],
)
tokens_without_search = token_counter(
model="gpt-3.5-turbo",
messages=[{"role": "tool", "content": "small"}],
)
assert tokens_with_search > tokens_without_search
@pytest.mark.asyncio
async def test_assistant_tool_calls_cache_control():
"""Test that assistant tool_calls with cache_control generate cachePoint blocks."""