mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
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
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:
parent
65f8e21ebd
commit
0eb3658e9b
5 changed files with 220 additions and 16 deletions
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue