Merge pull request #38808 from BerriAI/litellm_headroom_ccr_streaming_responses

fix(headroom): resolve CCR retrieval on streaming /v1/responses
This commit is contained in:
Mateo Wang 2026-09-03 13:13:18 -07:00 • committed by GitHub
commit 80250807db
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
13 changed files with 1156 additions and 103 deletions

View file

@ -932,7 +932,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
return OpenAiResponsesToChatCompletionStreamIterator(streaming_response, sync_stream, json_mode)
def _convert_content_str_to_input_text(self, content: str, role: str) -> dict[str, object]:
if role == "user" or role == "system" or role == "tool":
if role in ("user", "system", "developer", "tool"):
return {"type": "input_text", "text": content}
else:
return {"type": "output_text", "text": content}

View file

@ -2337,6 +2337,9 @@ class CustomStreamWrapper:
else:
self.sent_last_chunk = True
processed_chunk: Final = self.finish_reason_handler()
if self.stream_options is None:
usage: Final = calculate_total_usage(chunks=self.chunks)
processed_chunk._hidden_params["usage"] = usage # pyright: ignore[reportPrivateUsage] # sync parity
# see sync __next__'s sibling branch: deliberately do NOT restore
# here - this chunk is still this call's own data, and restoring
# before returning it would corrupt the caller's own log

View file

@ -99,9 +99,12 @@ from litellm.types.containers.main import (
)
from litellm.types.files import StreamingMediaUploadConfig, TwoStepFileUploadConfig
from litellm.types.integrations.custom_logger import (
NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES,
AgenticLoopPlan,
AgenticLoopRequestPatch,
AgenticLoopSafetyError,
converted_stream_requested,
is_interception_internal_key,
)
from litellm.types.llms.anthropic_messages.anthropic_response import (
AnthropicMessagesResponse,
@ -2760,6 +2763,7 @@ class BaseLLMHTTPHandler:
)
if self._has_agentic_completion_hook(logging_obj):
agentic_kwargs: Final = dict(litellm_params) # mutable-ok: agentic hooks mutate kwargs in place
final_response: Final = run_async_function(
self._call_agentic_completion_hooks,
response=initial_response,
@ -2770,10 +2774,19 @@ class BaseLLMHTTPHandler:
logging_obj=logging_obj,
stream=False,
custom_llm_provider=custom_llm_provider,
kwargs=dict(litellm_params),
kwargs=agentic_kwargs,
api_surface="responses",
)
return final_response if final_response is not None else initial_response
result: Final = final_response if final_response is not None else initial_response
if converted_stream_requested(agentic_kwargs) and not agentic_kwargs.get("_agentic_loop_depth"):
return self._wrap_responses_response_as_fake_stream(
result=result,
model=model,
responses_api_provider_config=responses_api_provider_config,
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,
)
return result
return initial_response
@ -2939,6 +2952,7 @@ class BaseLLMHTTPHandler:
logging_obj=logging_obj,
)
agentic_kwargs: Final = dict(litellm_params) # mutable-ok: agentic hooks mutate kwargs in place
final_response: Final = await self._call_agentic_completion_hooks(
response=initial_response,
model=model,
@ -2948,15 +2962,12 @@ class BaseLLMHTTPHandler:
logging_obj=logging_obj,
stream=False,
custom_llm_provider=custom_llm_provider,
kwargs=dict(litellm_params),
kwargs=agentic_kwargs,
api_surface="responses",
)
result: Final = final_response if final_response is not None else initial_response
interception_converted_stream: Final = litellm_params.get(
"_code_interpreter_interception_converted_stream"
) or litellm_params.get("_websearch_interception_converted_stream")
if interception_converted_stream and not litellm_params.get("_agentic_loop_depth"):
if converted_stream_requested(agentic_kwargs) and not agentic_kwargs.get("_agentic_loop_depth"):
return self._wrap_responses_response_as_fake_stream(
result=result,
model=model,
@ -5420,8 +5431,7 @@ class BaseLLMHTTPHandler:
kwargs_for_followup: Final = {
k: v
for k, v in kwargs.items()
if not k.startswith("_websearch_interception")
and not k.startswith("_compression_interception")
if not is_interception_internal_key(k, prefixes=NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES)
and k != "_code_interpreter_interception_converted_stream"
and k not in internal_keys
and k not in optional_params

View file

@ -33,8 +33,9 @@ import time
import uuid
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from itertools import accumulate
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Union, cast
from typing import TYPE_CHECKING, Any, Final, NamedTuple, Union, cast
from openai.types.responses.response_function_tool_call import ResponseFunctionToolCall
from pydantic import BaseModel, TypeAdapter
@ -42,6 +43,7 @@ from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_proxy_logger
from litellm.completion_extras.litellm_responses_transformation.transformation import (
LiteLLMResponsesTransformationHandler,
OpenAiResponsesToChatCompletionStreamIterator,
)
from litellm.llms.base_llm.guardrail_translation.base_translation import (
@ -74,6 +76,7 @@ from litellm.types.llms.openai import (
OutputTextDoneEvent,
ResponseAPIUsage,
ResponseCompletedEvent,
ResponsesAPIOptionalRequestParams,
ResponsesAPIResponse,
ResponsesAPIStreamEvents,
ResponsesAPIStreamingResponse,
@ -115,6 +118,199 @@ class ResponsesStreamChunk(TypedDict, total=False):
content_index: ReadOnly[int]
_PATCHABLE_ITEM_FIELDS: Final[Mapping[str, str]] = MappingProxyType(
{"function_call_output": "output", "message": "content"}
)
_EMPTY_RESPONSES_REQUEST: Final[ResponsesAPIOptionalRequestParams] = {}
def _item_rewrite_field(item: Mapping[str, object]) -> str | None:
item_type: Final = item.get("type")
if item_type is None:
return "content" if "content" in item else None
if not isinstance(item_type, str):
return None
return _PATCHABLE_ITEM_FIELDS.get(item_type)
def _rewritten_input_item(item: Mapping[str, object], rewritten: object) -> Mapping[str, object] | None:
field: Final = _item_rewrite_field(item)
if field is None or not isinstance(rewritten, Mapping):
return None
rewritten_content: Final = rewritten.get("content")
if isinstance(item.get(field), str) and isinstance(rewritten_content, str):
return {**item, field: rewritten_content} # mutable-ok: request input items must stay JSON-plain dicts
rewritten_row: Final = cast("AllMessageValues", rewritten) # cast-ok: guardrails hand back chat-shaped rows
converted_items, _ = LiteLLMResponsesTransformationHandler().convert_chat_completion_messages_to_responses_api(
[rewritten_row] # mutable-ok: converter signature takes a list
)
if len(converted_items) != 1 or not isinstance(converted_items[0], Mapping):
return None
first_converted: Final = cast("Mapping[str, object]", converted_items[0]) # cast-ok: isinstance-checked above
converted_value: Final = first_converted.get(field)
if converted_value is None:
return None
return {**item, field: converted_value} # mutable-ok: request input items must stay JSON-plain dicts
def _is_function_call_item(item: object) -> bool:
return isinstance(item, Mapping) and item.get("type") in ("function_call", "custom_tool_call")
def _last_message_role(messages: Sequence[object]) -> str | None:
if not messages:
return None
last: Final = messages[-1]
role: Final = last.get("role") if isinstance(last, Mapping) else getattr(last, "role", None)
return role if isinstance(role, str) else None
def _provenance_unit_bounds(
raw_input: Sequence[object],
solo_conversions: Sequence[Sequence[object]],
) -> tuple[tuple[int, int], ...]:
trailing_roles: Final = tuple(
accumulate(
(_last_message_role(messages) for messages in solo_conversions),
lambda previous, current: current if current is not None else previous,
)
)
start_indexes: Final = tuple(
index
for index in range(len(raw_input))
if index == 0 or not (_is_function_call_item(raw_input[index]) and trailing_roles[index - 1] == "assistant")
)
return tuple(zip(start_indexes, (*start_indexes[1:], len(raw_input))))
def _input_item_provenance(
raw_input: Sequence[object],
expected_messages: Sequence[object],
) -> tuple[Mapping[int, int], frozenset[int]] | None:
if not all(isinstance(item, Mapping) for item in raw_input):
return None
solo_conversions: Final = tuple(
LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
input=cast("ResponseInputParam", [item]), # cast-ok: items checked as Mappings above
responses_api_request=_EMPTY_RESPONSES_REQUEST,
)
for item in raw_input
)
full_conversion: Final = tuple(
LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
input=cast("ResponseInputParam", list(raw_input)), # cast-ok: items checked as Mappings above
responses_api_request=_EMPTY_RESPONSES_REQUEST,
)
)
if full_conversion != tuple(expected_messages):
return None
units: Final = _provenance_unit_bounds(raw_input, solo_conversions)
unit_messages: Final = tuple(
tuple(solo_conversions[start])
if end - start == 1
else tuple(
LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
input=cast("ResponseInputParam", list(raw_input[start:end])), # cast-ok: checked as Mappings above
responses_api_request=_EMPTY_RESPONSES_REQUEST,
)
)
for start, end in units
)
if tuple(message for messages in unit_messages for message in messages) != full_conversion:
return None
boundaries: Final = tuple(accumulate((len(messages) for messages in unit_messages), initial=0))
item_for_message: Final = MappingProxyType(
{
message_index: start
for unit_index, (start, end) in enumerate(units)
if end - start == 1
for message_index in range(boundaries[unit_index], boundaries[unit_index + 1])
}
)
tainted: Final = frozenset(
message_index
for unit_index, (start, end) in enumerate(units)
if end - start > 1
for message_index in range(boundaries[unit_index], boundaries[unit_index + 1])
)
return item_for_message, tainted
class _RequestFields(NamedTuple):
input: tuple[object, ...]
instructions: str | None
class _ExtractedInputs(NamedTuple):
inputs: GenericGuardrailAPIInputs
task_mappings: tuple[tuple[int, int | None], ...]
def _patched_request_fields(
raw_input: object,
instructions: object,
original_messages: Sequence[object],
structured_messages: Sequence[object],
) -> _RequestFields | None:
if not isinstance(raw_input, list) or len(original_messages) != len(structured_messages):
return None
offset: Final = 1 if instructions else 0
provenance: Final = _input_item_provenance(raw_input, tuple(original_messages)[offset:])
if provenance is None:
return None
item_for_message, tainted = provenance
changed: Final = tuple(
(index, rewritten)
for index, (original, rewritten) in enumerate(zip(original_messages, structured_messages))
if original != rewritten
)
instruction_rewrites: Final = tuple(rewritten for index, rewritten in changed if index < offset)
rewritten_instructions: Final = (
instruction_rewrites[0].get("content")
if instruction_rewrites and isinstance(instruction_rewrites[0], Mapping)
else instructions
)
instructions_value: Final = rewritten_instructions if isinstance(rewritten_instructions, str) else None
if rewritten_instructions is not None and instructions_value is None:
return None
body_changes: Final = tuple((index - offset, rewritten) for index, rewritten in changed if index >= offset)
if any(message_index in tainted or message_index not in item_for_message for message_index, _ in body_changes):
return None
replacements: Final = MappingProxyType(
{
item_for_message[message_index]: _rewritten_input_item(
cast("Mapping[str, object]", raw_input[item_for_message[message_index]]), # cast-ok: checked Mappings
rewritten,
)
for message_index, rewritten in body_changes
}
)
if len(replacements) != len(body_changes) or any(item is None for item in replacements.values()):
return None
return _RequestFields(
input=tuple(replacements.get(index, item) for index, item in enumerate(raw_input)),
instructions=instructions_value,
)
def _patch_or_convert_request_fields(
raw_input: object,
instructions: object,
original_messages: Sequence[object],
structured_messages: Sequence[AllMessageValues],
) -> _RequestFields | None:
if not isinstance(structured_messages, list):
return None
patched: Final = _patched_request_fields(raw_input, instructions, original_messages, structured_messages)
if patched is not None:
return patched
input_items, converted_instructions = (
LiteLLMResponsesTransformationHandler().convert_chat_completion_messages_to_responses_api(structured_messages)
)
return _RequestFields(input=tuple(input_items), instructions=converted_instructions)
def _next_stream_sequence_number(responses_so_far: Sequence[Any] | None) -> int:
sequence_numbers: Final = (
item.get("sequence_number") if isinstance(item, dict) else getattr(item, "sequence_number", None)
@ -162,9 +358,8 @@ class OpenAIResponsesHandler(BaseTranslation):
Handles both string input and list of message objects.
"""
input_data: Final[str | ResponseInputParam | None] = data.get("input")
if input_data is None:
if not isinstance(input_data, (str, list)):
return data
structured_messages: Final = self.get_structured_messages(data)
raw_tools: Final = data.get("tools")
original_tools: Final[tuple[Mapping[str, object], ...]] = (
@ -173,94 +368,93 @@ class OpenAIResponsesHandler(BaseTranslation):
flattened_tool_groups: Final = tuple(
form.chat_tools for form in LiteLLMCompletionResponsesConfig.responses_tools_to_chat_forms(original_tools)
)
flattened_tools: Final = tuple(
cast(ChatCompletionToolParam, tool) # cast-ok: mcp tools ride along in the guardrail's tool list
for group in flattened_tool_groups
for tool in group
)
tools_to_check: Final[list[ChatCompletionToolParam]] = list( # mutable-ok: guardrail inputs want a list
copy.deepcopy(flattened_tools)
)
# Handle simple string input
if isinstance(input_data, str):
inputs = GenericGuardrailAPIInputs(texts=[input_data])
if tools_to_check:
inputs["tools"] = tools_to_check
if structured_messages:
inputs["structured_messages"] = structured_messages
# Include model information if available
model = data.get("model")
if model:
inputs["model"] = model
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs=inputs,
request_data=data,
input_type="request",
logging_obj=litellm_logging_obj,
)
guardrailed_texts = guardrailed_inputs.get("texts", [])
data["input"] = guardrailed_texts[0] if guardrailed_texts else input_data
self._apply_guardrailed_tools_to_data(
data, original_tools, flattened_tool_groups, guardrailed_inputs.get("tools")
)
verbose_proxy_logger.debug("OpenAI Responses API: Processed string input")
return data
# Handle list input (ResponseInputParam)
if not isinstance(input_data, list):
extracted: Final = self._extract_guardrail_inputs(data, input_data, flattened_tool_groups)
if not extracted.inputs.get("texts"):
return data
if structured_messages:
extracted.inputs["structured_messages"] = structured_messages
guardrailed_inputs: Final = await guardrail_to_apply.apply_guardrail(
inputs=extracted.inputs,
request_data=data,
input_type="request",
logging_obj=litellm_logging_obj,
)
self._apply_guardrailed_tools_to_data(
data, original_tools, flattened_tool_groups, guardrailed_inputs.get("tools")
)
written_back: Final = self._written_back_request_fields(data, structured_messages, guardrailed_inputs)
if written_back is not None:
data["input"] = list(written_back.input) # mutable-ok: JSON body
if written_back.instructions is None:
data.pop("instructions", None)
else:
data["instructions"] = written_back.instructions # rebind-ok: data is an out-param
elif isinstance(input_data, str):
guardrailed_texts: Final = guardrailed_inputs.get("texts") or ()
data["input"] = guardrailed_texts[0] if guardrailed_texts else input_data # rebind-ok: data is an out-param
else:
await self._apply_guardrail_responses_to_input(
messages=input_data,
responses=guardrailed_inputs.get("texts") or (),
task_mappings=extracted.task_mappings,
)
verbose_proxy_logger.debug("OpenAI Responses API: Processed input messages: %s", data.get("input"))
return data
def _extract_guardrail_inputs(
self,
data: Mapping[str, object],
input_data: "str | ResponseInputParam",
flattened_tool_groups: Sequence[Sequence[Mapping[str, object]]],
) -> _ExtractedInputs:
texts_to_check: Final[list[str]] = []
images_to_check: Final[list[str]] = []
task_mappings: Final[list[tuple[int, int | None]]] = []
# Step 1: Extract all text content, images, and tools
for msg_idx, message in enumerate(input_data):
self._extract_input_text_and_images(
message=message,
msg_idx=msg_idx,
texts_to_check=texts_to_check,
images_to_check=images_to_check,
task_mappings=task_mappings,
tools_to_check: Final[list[ChatCompletionToolParam]] = list( # mutable-ok: guardrail inputs want a list
copy.deepcopy(
tuple(
cast(ChatCompletionToolParam, tool) # cast-ok: mcp tools ride along in the guardrail's tool list
for group in flattened_tool_groups
for tool in group
)
)
)
if isinstance(input_data, str):
texts_to_check.append(input_data)
else:
for msg_idx, message in enumerate(input_data):
self._extract_input_text_and_images(
message=message,
msg_idx=msg_idx,
texts_to_check=texts_to_check,
images_to_check=images_to_check,
task_mappings=task_mappings,
)
inputs: Final = GenericGuardrailAPIInputs(texts=texts_to_check)
if images_to_check:
inputs["images"] = images_to_check
if tools_to_check:
inputs["tools"] = tools_to_check
model: Final = data.get("model")
if isinstance(model, str):
inputs["model"] = model
return _ExtractedInputs(inputs=inputs, task_mappings=tuple(task_mappings))
# Step 2: Apply guardrail to all texts in batch
if texts_to_check:
inputs = GenericGuardrailAPIInputs(texts=texts_to_check)
if images_to_check:
inputs["images"] = images_to_check
if tools_to_check:
inputs["tools"] = tools_to_check
if structured_messages:
inputs["structured_messages"] = structured_messages
# Include model information if available
model = data.get("model")
if model:
inputs["model"] = model
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs=inputs,
request_data=data,
input_type="request",
logging_obj=litellm_logging_obj,
)
guardrailed_texts = guardrailed_inputs.get("texts", [])
self._apply_guardrailed_tools_to_data(
data, original_tools, flattened_tool_groups, guardrailed_inputs.get("tools")
)
# Step 3: Map guardrail responses back to original input structure
await self._apply_guardrail_responses_to_input(
messages=input_data,
responses=guardrailed_texts,
task_mappings=task_mappings,
)
verbose_proxy_logger.debug("OpenAI Responses API: Processed input messages: %s", input_data)
return data
@staticmethod
def _written_back_request_fields(
data: Mapping[str, object],
structured_messages: Sequence[AllMessageValues] | None,
guardrailed_inputs: GenericGuardrailAPIInputs,
) -> _RequestFields | None:
guardrailed: Final = guardrailed_inputs.get("structured_messages")
if guardrailed is None or guardrailed is structured_messages:
return None
return _patch_or_convert_request_fields(
data.get("input"),
data.get("instructions"),
structured_messages or (),
guardrailed,
)
def extract_request_tool_names(self, data: dict) -> list[str]:
"""Extract tool names from Responses API request (tools[].name for function
@ -331,8 +525,8 @@ class OpenAIResponsesHandler(BaseTranslation):
async def _apply_guardrail_responses_to_input(
self,
messages: Any, # Can be List[Dict[str, Any]] or ResponseInputParam
responses: list[str],
task_mappings: list[tuple[int, int | None]],
responses: Sequence[str],
task_mappings: Sequence[tuple[int, int | None]],
) -> None:
"""
Apply guardrail responses back to input messages.

View file

@ -916,9 +916,10 @@ class CompresrGuardrail(CustomGuardrail):
def _mirror_texts_channel(input_texts: object, applied: _CompressionResult) -> list[object] | None:
"""Compressed content mirrored into the Responses `texts` channel.
The chat/Anthropic handlers round-trip ``structured_messages``; the
Responses translation cannot rebuild its input from chat messages and
instead writes back through ``texts``. This matches by value, so a
The chat/Anthropic/Responses handlers round-trip
``structured_messages``; translations without that round-trip write
back through ``texts``, so the compressed content is mirrored there
too. This matches by value, so a
replacement is applied only when it is unambiguous: one compression per
text, and every occurrence in ``texts`` accounted for by a compressed
target. Anything else is left uncompressed rather than risk a wrong or

View file

@ -50,6 +50,9 @@ if TYPE_CHECKING:
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
BYPASS_HEADER: Final = "x-headroom-bypass"
_STREAM_CONVERTIBLE_CALL_TYPES: Final = frozenset(
(CallTypes.completion, CallTypes.acompletion, CallTypes.responses, CallTypes.aresponses)
)
HEADROOM_RETRIEVE_TOOL_NAME: Final = "headroom_retrieve"
_HASH_PATTERN: Final = re.compile(r"hash=([a-f0-9]{24})")
_HASH_CACHE_TTL_SECONDS: Final = 15 * 60
@ -725,6 +728,10 @@ class HeadroomGuardrail(CustomGuardrail):
verbose_proxy_logger.debug("Headroom: %s header set; skipping compression", BYPASS_HEADER)
return inputs
if request_data.get("background"):
verbose_proxy_logger.debug("Headroom: background request; skipping compression")
return inputs
structured_messages: Final = inputs.get("structured_messages")
if not _is_object_list(structured_messages) or not structured_messages:
return inputs
@ -826,9 +833,9 @@ class HeadroomGuardrail(CustomGuardrail):
) -> dict[str, Any] | None: # mutable-ok: overrides CustomLogger hook whose contract is a plain dict
base_result: Final = await super().async_pre_call_deployment_hook(kwargs, call_type)
effective: Final = base_result if base_result is not None else kwargs
if call_type not in (CallTypes.completion, CallTypes.acompletion):
if call_type not in _STREAM_CONVERTIBLE_CALL_TYPES:
return base_result
if not effective.get("stream"):
if not effective.get("stream") or effective.get("background"):
return base_result
if not has_headroom_retrieve_tool(effective.get("tools")):
return base_result

View file

@ -562,6 +562,12 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
hidden_params: Final = getattr(chunk, "_hidden_params", None)
if hidden_params is not None:
chunk_dict["_hidden_params"] = dict(hidden_params) if isinstance(hidden_params, dict) else hidden_params
if (
chunk_dict.get("usage") is None
and isinstance(hidden_params, dict)
and hidden_params.get("usage") is not None
):
chunk_dict["usage"] = hidden_params["usage"]
return chunk_dict
def create_reasoning_summary_text_done_event(

View file

@ -1,3 +1,4 @@
from collections.abc import Mapping
from typing import Any, Final
from pydantic import BaseModel, Field
@ -29,6 +30,13 @@ def is_interception_internal_key(
return any(key.startswith(prefix) for prefix in prefixes)
CONVERTED_STREAM_KEYS: Final = frozenset(f"{prefix}_converted_stream" for prefix in INTERCEPTION_INTERNAL_PREFIXES)
def converted_stream_requested(params: Mapping[str, object]) -> bool:
return any(bool(params.get(key)) for key in CONVERTED_STREAM_KEYS)
class AgenticLoopSafetyError(ValueError):
"""
Raised when an agentic-loop safety rail refuses a rerun.

View file

@ -1295,6 +1295,19 @@ def test_text_plus_tool_calls_sequence():
# =============================================================================
def test_developer_message_content_uses_input_text():
handler = LiteLLMResponsesTransformationHandler()
input_items, instructions = handler.convert_chat_completion_messages_to_responses_api(
[{"role": "developer", "content": "Always answer in French."}]
)
assert instructions is None
assert input_items == [
{"type": "message", "role": "developer", "content": [{"type": "input_text", "text": "Always answer in French."}]}
]
def test_tool_message_output_uses_input_text_not_output_text():
"""
Test that tool message content uses input_text type, not output_text.

View file

@ -4761,6 +4761,43 @@ async def test_async_stream_assembled_response_keeps_vertex_traffic_type(logging
assert assembled._hidden_params["provider_specific_fields"]["traffic_type"] == "ON_DEMAND_FLEX"
@pytest.mark.asyncio
async def test_async_fake_stream_final_chunk_carries_hidden_usage(logging_obj: Logging):
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
from litellm.types.utils import ModelResponse
model_response = ModelResponse(
id="chatcmpl-fake-stream",
model="my-random-model",
choices=[
{
"index": 0,
"message": {"role": "assistant", "content": "hello world"},
"finish_reason": "stop",
}
],
)
model_response.usage = Usage(prompt_tokens=1234, completion_tokens=7, total_tokens=1241)
wrapper = CustomStreamWrapper(
completion_stream=MockResponseIterator(model_response=model_response),
model="my-random-model",
custom_llm_provider="anthropic",
logging_obj=logging_obj,
)
final_chunk = None
async for chunk in wrapper:
final_chunk = chunk
assert final_chunk is not None
hidden_usage = final_chunk._hidden_params.get("usage")
assert hidden_usage is not None
assert hidden_usage.prompt_tokens == 1234
assert hidden_usage.completion_tokens == 7
assert hidden_usage.total_tokens == 1241
class TestStableStreamingResponseId:
"""
All chunks of one streamed response must share the same top-level id

View file

@ -1329,6 +1329,523 @@ class TestOpenAIResponsesHandlerToolInjection:
assert "injected_tool" in names
COMPRESSED_MARKER = "[compressed document; retrieve the full text with hash=b573993006976af767214fac]"
class StructuredRewriteGuardrail(CustomGuardrail):
"""Guardrail that rewrites whole messages via structured_messages and leaves
texts untouched, the way message-compressing guardrails do."""
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional[Any] = None,
) -> GenericGuardrailAPIInputs:
messages = list(inputs.get("structured_messages") or [])
first_user = next(i for i, m in enumerate(messages) if m.get("role") == "user")
rewritten = [
{**m, "content": COMPRESSED_MARKER} if i == first_user else m for i, m in enumerate(messages)
]
return {**inputs, "structured_messages": rewritten}
class ToolOutputRewriteGuardrail(CustomGuardrail):
"""Guardrail that compresses the first tool-result row, the way Headroom does."""
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional[Any] = None,
) -> GenericGuardrailAPIInputs:
messages = list(inputs.get("structured_messages") or [])
first_tool = next(i for i, m in enumerate(messages) if isinstance(m, dict) and m.get("role") == "tool")
rewritten = [
{**m, "content": COMPRESSED_MARKER} if i == first_tool else m for i, m in enumerate(messages)
]
return {**inputs, "structured_messages": rewritten}
class DroppingRewriteGuardrail(CustomGuardrail):
"""Guardrail that rewrites the first user row and drops the last row, so the
rewrite can only land through the full-conversion fallback."""
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional[Any] = None,
) -> GenericGuardrailAPIInputs:
messages = list(inputs.get("structured_messages") or [])
first_user = next(i for i, m in enumerate(messages) if isinstance(m, dict) and m.get("role") == "user")
rewritten = [
{**m, "content": COMPRESSED_MARKER} if i == first_user else m for i, m in enumerate(messages)
]
return {**inputs, "structured_messages": rewritten[:-1]}
def _texts(item: dict) -> list[str]:
content = item.get("content")
if isinstance(content, str):
return [content]
return [part["text"] for part in content]
class TestStructuredMessagesWriteBack:
"""A guardrail's structured_messages rewrite must land in the Responses request,
not only the per-text mapping the chat handler shares with it."""
@pytest.mark.asyncio
async def test_list_input_gets_rewritten_messages_and_keeps_instructions(self):
handler = OpenAIResponsesHandler()
data = {
"model": "gpt-5.6",
"instructions": "Answer from the memo only.",
"input": [
{"role": "user", "content": "memo " * 400},
{"role": "assistant", "content": "Understood."},
{"role": "user", "content": "What is the codename?"},
],
}
result = await handler.process_input_messages(data, StructuredRewriteGuardrail())
assert result["instructions"] == "Answer from the memo only."
user_items = [item for item in result["input"] if item.get("role") == "user"]
assert [_texts(item) for item in user_items] == [[COMPRESSED_MARKER], ["What is the codename?"]]
assert not any(item.get("role") == "system" for item in result["input"])
assert _texts(next(item for item in result["input"] if item.get("role") == "assistant")) == ["Understood."]
@pytest.mark.asyncio
async def test_string_input_becomes_rewritten_message_list(self):
handler = OpenAIResponsesHandler()
data = {"model": "gpt-5.6", "input": "memo " * 400}
result = await handler.process_input_messages(data, StructuredRewriteGuardrail())
assert [_texts(item) for item in result["input"]] == [[COMPRESSED_MARKER]]
assert "instructions" not in result
@pytest.mark.asyncio
async def test_developer_item_preserved_verbatim_by_row_patch(self):
handler = OpenAIResponsesHandler()
developer_item = {"role": "developer", "content": "Always answer in French."}
data = {
"model": "gpt-5.6",
"input": [
developer_item,
{"role": "user", "content": "memo " * 400},
{"role": "user", "content": "What is the codename?"},
],
}
result = await handler.process_input_messages(data, StructuredRewriteGuardrail())
assert result["input"][0] is developer_item
assert developer_item["content"] == "Always answer in French."
assert _texts(result["input"][1]) == [COMPRESSED_MARKER]
assert _texts(result["input"][2]) == ["What is the codename?"]
@pytest.mark.asyncio
async def test_reasoning_and_function_call_items_survive_tool_output_compression(self):
handler = OpenAIResponsesHandler()
reasoning_item = {
"id": "rs_123",
"type": "reasoning",
"summary": [],
"encrypted_content": "gAAAAA-signed-reasoning",
}
function_call_item = {
"id": "fc_123",
"type": "function_call",
"call_id": "call_abc",
"name": "read_document",
"arguments": '{"path": "memo.txt"}',
"status": "completed",
}
data = {
"model": "gpt-5.6",
"instructions": "Answer from the memo only.",
"input": [
reasoning_item,
function_call_item,
{"type": "function_call_output", "call_id": "call_abc", "output": "memo " * 400},
{"role": "user", "content": "What is the codename?"},
],
}
result = await handler.process_input_messages(data, ToolOutputRewriteGuardrail())
assert result["instructions"] == "Answer from the memo only."
assert result["input"][0] is reasoning_item
assert reasoning_item["encrypted_content"] == "gAAAAA-signed-reasoning"
assert result["input"][1] is function_call_item
assert function_call_item["id"] == "fc_123"
assert result["input"][2] == {
"type": "function_call_output",
"call_id": "call_abc",
"output": COMPRESSED_MARKER,
}
assert result["input"][3] == {"role": "user", "content": "What is the codename?"}
@pytest.mark.asyncio
async def test_web_search_call_item_preserved_verbatim(self):
handler = OpenAIResponsesHandler()
web_search_item = {
"id": "ws_123",
"type": "web_search_call",
"status": "completed",
"action": {"type": "search", "query": "codename memo"},
}
data = {
"model": "gpt-5.6",
"input": [
web_search_item,
{"role": "user", "content": "memo " * 400},
{"role": "user", "content": "What is the codename?"},
],
}
result = await handler.process_input_messages(data, StructuredRewriteGuardrail())
assert result["input"][0] is web_search_item
assert _texts(result["input"][1]) == [COMPRESSED_MARKER]
assert _texts(result["input"][2]) == ["What is the codename?"]
@pytest.mark.asyncio
async def test_row_count_change_falls_back_to_full_conversion(self):
handler = OpenAIResponsesHandler()
data = {
"model": "gpt-5.6",
"input": [
{"role": "developer", "content": "Always answer in French."},
{"role": "user", "content": "memo " * 400},
{"role": "user", "content": "What is the codename?"},
],
}
result = await handler.process_input_messages(data, DroppingRewriteGuardrail())
assert len(result["input"]) == 2
developer = next(item for item in result["input"] if item.get("role") == "developer")
assert developer["content"] == [{"type": "input_text", "text": "Always answer in French."}]
assert _texts(next(item for item in result["input"] if item.get("role") == "user")) == [COMPRESSED_MARKER]
@pytest.mark.asyncio
async def test_same_inputs_object_back_keeps_the_text_mapping(self):
handler = OpenAIResponsesHandler()
original_input = [
{"role": "user", "content": "Hello"},
{"role": "user", "content": [{"type": "input_text", "text": "Again"}]},
]
data = {"model": "gpt-5.6", "input": original_input}
result = await handler.process_input_messages(data, MockGuardrail())
assert result["input"] is original_input
assert [_texts(item) for item in result["input"]] == [["Hello [GUARDRAILED]"], ["Again [GUARDRAILED]"]]
class AllToolOutputsRewriteGuardrail(CustomGuardrail):
"""Guardrail that compresses every tool-result row."""
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional[Any] = None,
) -> GenericGuardrailAPIInputs:
messages = list(inputs.get("structured_messages") or [])
rewritten = [
{**m, "content": COMPRESSED_MARKER} if isinstance(m, dict) and m.get("role") == "tool" else m
for m in messages
]
return {**inputs, "structured_messages": rewritten}
class AssistantRewriteGuardrail(CustomGuardrail):
"""Guardrail that rewrites the first assistant row's content."""
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional[Any] = None,
) -> GenericGuardrailAPIInputs:
messages = list(inputs.get("structured_messages") or [])
first = next(i for i, m in enumerate(messages) if isinstance(m, dict) and m.get("role") == "assistant")
rewritten = [{**m, "content": COMPRESSED_MARKER} if i == first else m for i, m in enumerate(messages)]
return {**inputs, "structured_messages": rewritten}
class DictStructuredMessagesGuardrail(CustomGuardrail):
"""Guardrail that hands back a raw evaluation dict instead of a message list,
the way HiddenLayer v2 does."""
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional[Any] = None,
) -> GenericGuardrailAPIInputs:
return {**inputs, "structured_messages": {"evaluation": "allowed", "messages": []}}
def _parallel_tool_call_input() -> list:
return [
{"id": "fc_1", "type": "function_call", "call_id": "call_1", "name": "read_a", "arguments": "{}"},
{"id": "fc_2", "type": "function_call", "call_id": "call_2", "name": "read_b", "arguments": "{}"},
{"type": "function_call_output", "call_id": "call_1", "output": "memo " * 400},
{"type": "function_call_output", "call_id": "call_2", "output": "note " * 400},
{"role": "user", "content": "What is the codename?"},
]
class TestProvenancePatching:
"""The O(n) provenance pass must keep patching rewritten rows in place for the
shapes real agent loops produce, and fall back safely everywhere else."""
@pytest.mark.asyncio
async def test_parallel_tool_call_outputs_both_patched(self):
handler = OpenAIResponsesHandler()
raw_input = _parallel_tool_call_input()
fc_1, fc_2 = raw_input[0], raw_input[1]
data = {"model": "gpt-5.6", "input": raw_input}
result = await handler.process_input_messages(data, AllToolOutputsRewriteGuardrail())
assert result["input"][0] is fc_1
assert result["input"][1] is fc_2
assert result["input"][2] == {"type": "function_call_output", "call_id": "call_1", "output": COMPRESSED_MARKER}
assert result["input"][3] == {"type": "function_call_output", "call_id": "call_2", "output": COMPRESSED_MARKER}
assert result["input"][4] == {"role": "user", "content": "What is the codename?"}
@pytest.mark.asyncio
async def test_assistant_turn_with_tool_call_keeps_items_verbatim(self):
handler = OpenAIResponsesHandler()
assistant_item = {"role": "assistant", "content": "Let me read the memo."}
function_call_item = {
"id": "fc_9",
"type": "function_call",
"call_id": "call_9",
"name": "read_document",
"arguments": '{"path": "memo.txt"}',
}
data = {
"model": "gpt-5.6",
"input": [
assistant_item,
function_call_item,
{"type": "function_call_output", "call_id": "call_9", "output": "memo " * 400},
{"role": "user", "content": "What is the codename?"},
],
}
result = await handler.process_input_messages(data, ToolOutputRewriteGuardrail())
assert result["input"][0] is assistant_item
assert result["input"][1] is function_call_item
assert result["input"][2] == {"type": "function_call_output", "call_id": "call_9", "output": COMPRESSED_MARKER}
@pytest.mark.asyncio
async def test_rewrite_of_merged_tool_call_message_falls_back(self):
handler = OpenAIResponsesHandler()
raw_input = _parallel_tool_call_input()
data = {"model": "gpt-5.6", "input": raw_input}
result = await handler.process_input_messages(data, AssistantRewriteGuardrail())
assert not any(item is original for item in result["input"] for original in raw_input)
assistant_items = [item for item in result["input"] if item.get("role") == "assistant"]
assert [_texts(item) for item in assistant_items] == [[COMPRESSED_MARKER]]
@pytest.mark.asyncio
async def test_rewrite_of_lone_function_call_message_falls_back(self):
handler = OpenAIResponsesHandler()
data = {
"model": "gpt-5.6",
"input": [
{"id": "fc_1", "type": "function_call", "call_id": "call_1", "name": "read_a", "arguments": "{}"},
{"type": "function_call_output", "call_id": "call_1", "output": "memo memo"},
{"role": "user", "content": "What is the codename?"},
],
}
raw_input = data["input"]
result = await handler.process_input_messages(data, AssistantRewriteGuardrail())
assert not any(item is original for item in result["input"] for original in raw_input)
assistant_items = [item for item in result["input"] if item.get("role") == "assistant"]
assert [_texts(item) for item in assistant_items] == [[COMPRESSED_MARKER]]
def test_provenance_bails_on_non_mapping_item(self):
from litellm.llms.openai.responses.guardrail_translation.handler import _input_item_provenance
assert _input_item_provenance(["not a mapping"], []) is None
def test_provenance_bails_when_expected_messages_disagree(self):
from litellm.llms.openai.responses.guardrail_translation.handler import _input_item_provenance
assert _input_item_provenance([{"role": "user", "content": "hi"}], [{"role": "user", "content": "bye"}]) is None
def test_provenance_bails_on_unpredicted_merge(self):
from litellm.llms.openai.responses.guardrail_translation.handler import _input_item_provenance
from litellm.responses.litellm_completion_transformation.transformation import (
LiteLLMCompletionResponsesConfig,
)
raw_input = [
{"id": "fc_1", "type": "function_call", "call_id": "call_1", "name": "read_a", "arguments": "{}"},
{"role": "assistant", "content": "Reading the memo now."},
]
expected = LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
input=raw_input, responses_api_request={}
)
assert len(expected) == 1
assert _input_item_provenance(raw_input, expected) is None
def test_provenance_maps_and_taints_parallel_tool_calls(self):
from litellm.llms.openai.responses.guardrail_translation.handler import _input_item_provenance
from litellm.responses.litellm_completion_transformation.transformation import (
LiteLLMCompletionResponsesConfig,
)
raw_input = _parallel_tool_call_input()
expected = LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
input=raw_input, responses_api_request={}
)
provenance = _input_item_provenance(raw_input, expected)
assert provenance is not None
item_for_message, tainted = provenance
assert tainted == {0}
assert dict(item_for_message) == {1: 2, 2: 3, 3: 4}
class TestDictStructuredMessagesGuard:
"""A guardrail handing back a non-list structured_messages payload must not
blow up the request; the write-back is skipped instead."""
@pytest.mark.asyncio
async def test_list_input_survives_dict_structured_messages(self):
handler = OpenAIResponsesHandler()
original_input = [{"role": "user", "content": "Hello"}]
data = {"model": "gpt-5.6", "input": original_input}
result = await handler.process_input_messages(data, DictStructuredMessagesGuardrail())
assert result["input"] is original_input
assert result["input"] == [{"role": "user", "content": "Hello"}]
@pytest.mark.asyncio
async def test_string_input_survives_dict_structured_messages(self):
handler = OpenAIResponsesHandler()
data = {"model": "gpt-5.6", "input": "Hello there"}
result = await handler.process_input_messages(data, DictStructuredMessagesGuardrail())
assert result["input"] == "Hello there"
class SystemRewriteGuardrail(CustomGuardrail):
"""Guardrail that rewrites the system row, the way prompt-hardening guardrails do."""
def __init__(self, rewritten_content: Any = COMPRESSED_MARKER):
super().__init__()
self.rewritten_content = rewritten_content
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional[Any] = None,
) -> GenericGuardrailAPIInputs:
messages = list(inputs.get("structured_messages") or [])
first = next(i for i, m in enumerate(messages) if isinstance(m, dict) and m.get("role") == "system")
rewritten = [
{**m, "content": self.rewritten_content} if i == first else m for i, m in enumerate(messages)
]
return {**inputs, "structured_messages": rewritten}
class TestPatchEdgeBranches:
@pytest.mark.asyncio
async def test_multimodal_user_item_rewritten_through_conversion(self):
handler = OpenAIResponsesHandler()
data = {
"model": "gpt-5.6",
"input": [
{"role": "user", "content": [{"type": "input_text", "text": "memo " * 400}]},
{"role": "user", "content": "What is the codename?"},
],
}
result = await handler.process_input_messages(data, StructuredRewriteGuardrail())
assert _texts(result["input"][0]) == [COMPRESSED_MARKER]
assert result["input"][1] == {"role": "user", "content": "What is the codename?"}
@pytest.mark.asyncio
async def test_instructions_rewrite_lands_in_instructions_field(self):
handler = OpenAIResponsesHandler()
user_item = {"role": "user", "content": "What is the codename?"}
data = {
"model": "gpt-5.6",
"instructions": "Answer from the memo only.",
"input": [user_item],
}
result = await handler.process_input_messages(data, SystemRewriteGuardrail())
assert result["instructions"] == COMPRESSED_MARKER
assert result["input"][0] is user_item
@pytest.mark.asyncio
async def test_non_string_instructions_rewrite_falls_back(self):
handler = OpenAIResponsesHandler()
user_item = {"role": "user", "content": "What is the codename?"}
data = {
"model": "gpt-5.6",
"instructions": "Answer from the memo only.",
"input": [user_item],
}
result = await handler.process_input_messages(
data, SystemRewriteGuardrail(rewritten_content=[{"type": "text", "text": COMPRESSED_MARKER}])
)
assert result["input"][0] is not user_item
@pytest.mark.asyncio
async def test_unpredicted_merge_falls_back_through_patch(self):
handler = OpenAIResponsesHandler()
raw_input = [
{"id": "fc_1", "type": "function_call", "call_id": "call_1", "name": "read_a", "arguments": "{}"},
{"role": "assistant", "content": "Reading the memo now."},
{"type": "function_call_output", "call_id": "call_1", "output": "memo memo"},
{"role": "user", "content": "memo " * 400},
]
data = {"model": "gpt-5.6", "input": raw_input}
result = await handler.process_input_messages(data, StructuredRewriteGuardrail())
assert not any(item is original for item in result["input"] for original in raw_input)
user_items = [item for item in result["input"] if item.get("role") == "user"]
assert _texts(user_items[0]) == [COMPRESSED_MARKER]
def test_item_rewrite_field_ignores_non_string_type(self):
from litellm.llms.openai.responses.guardrail_translation.handler import _item_rewrite_field
assert _item_rewrite_field({"type": 123, "content": "hello"}) is None
class ToolEditingGuardrail(CustomGuardrail):
"""Guardrail that rewrites the flattened chat tools it was handed through ``edit``"""

View file

@ -193,6 +193,33 @@ async def test_apply_guardrail_compresses_and_returns_structured_messages(
assert "headroom" in _applied_guardrails(request_data)
@pytest.mark.asyncio
async def test_apply_guardrail_leaves_background_requests_uncompressed(
guardrail: HeadroomGuardrail,
):
inputs = GenericGuardrailAPIInputs(
texts=["A" * 5000],
structured_messages=ORIGINAL_MESSAGES,
)
request_data = {"model": "gpt-4o", "background": True}
with patch.object(
guardrail.async_handler,
"post",
new_callable=AsyncMock,
return_value=_make_compress_response(COMPRESSED_MESSAGES),
) as post:
result = await guardrail.apply_guardrail(
inputs=inputs,
request_data=request_data,
input_type="request",
)
assert result is inputs
post.assert_not_awaited()
assert _recorded_guardrail_entries(request_data) == []
def _recorded_guardrail_response(request_data: dict) -> dict:
entries = request_data["metadata"]["standard_logging_guardrail_information"]
assert len(entries) == 1
@ -954,6 +981,40 @@ async def test_passthrough_handler_does_not_log_headroom_as_run(
assert "headroom" not in _applied_guardrails(data)
@pytest.mark.asyncio
async def test_responses_request_sends_compressed_input_and_retrieve_tool_upstream(
guardrail: HeadroomGuardrail,
):
"""Regression for LIT-6494: on /v1/responses the compressed messages must be
written back into `input`, not only the retrieve tool into `tools`, or the
model keeps reading the full document and never calls headroom_retrieve."""
from litellm.llms.openai.responses.guardrail_translation.handler import OpenAIResponsesHandler
data = {
"model": "gpt-5.6",
"instructions": ORIGINAL_MESSAGES[0]["content"],
"input": [{"role": m["role"], "content": m["content"]} for m in ORIGINAL_MESSAGES[1:]],
"tools": [{"type": "function", "name": "get_weather", "parameters": {"type": "object", "properties": {}}}],
}
with patch.object(
guardrail.async_handler,
"post",
new_callable=AsyncMock,
return_value=_make_compress_response(COMPRESSED_MESSAGES_WITH_HASH),
):
result = await OpenAIResponsesHandler().process_input_messages(data=data, guardrail_to_apply=guardrail)
assert result["instructions"] == ORIGINAL_MESSAGES[0]["content"]
assert [item["content"] for item in result["input"]] == [
COMPRESSED_MESSAGES_WITH_HASH[0]["content"],
ORIGINAL_MESSAGES[2]["content"],
ORIGINAL_MESSAGES[3]["content"],
]
assert "A" * 5000 not in json.dumps(result["input"])
assert [tool["name"] for tool in result["tools"]] == ["get_weather", HEADROOM_RETRIEVE_TOOL_NAME]
@pytest.mark.asyncio
async def test_apply_guardrail_http_error_raises():
guardrail = _make_guardrail()
@ -1950,6 +2011,58 @@ def _openai_text_payload(content: str) -> dict:
return _openai_completion_payload({"role": "assistant", "content": content}, "stop")
def _responses_retrieve_tool_definition() -> dict:
return {"type": "function", **_retrieve_tool_definition()["function"]}
def _openai_responses_payload(output_item: dict) -> dict:
return {
"id": "resp_ccr",
"object": "response",
"created_at": 1700000000,
"status": "completed",
"model": "gpt-4o",
"output": [output_item],
"usage": {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15},
"parallel_tool_calls": True,
"tool_choice": "auto",
"tools": [],
"error": None,
"incomplete_details": None,
"instructions": None,
"metadata": {},
"temperature": 1.0,
"top_p": 1.0,
"text": {"format": {"type": "text"}},
"truncation": "disabled",
}
def _openai_responses_retrieve_call_payload() -> dict:
return _openai_responses_payload(
{
"type": "function_call",
"id": "fc_ccr",
"call_id": "call_ccr",
"name": HEADROOM_RETRIEVE_TOOL_NAME,
"arguments": json.dumps({"hash": CCR_HASH}),
"status": "completed",
}
)
def _openai_responses_text_payload(text: str) -> dict:
return _openai_responses_payload(
{
"type": "message",
"id": "msg_ccr",
"role": "assistant",
"status": "completed",
"content": [{"type": "output_text", "text": text, "annotations": []}],
}
)
@pytest.mark.parametrize(
"call_type, stream, tools, expect_conversion",
[
@ -1958,12 +2071,14 @@ def _openai_text_payload(content: str) -> dict:
(CallTypes.acompletion, False, [_retrieve_tool_definition()], False),
(CallTypes.acompletion, True, [{"type": "function", "function": {"name": "get_weather"}}], False),
(CallTypes.acompletion, True, None, False),
(CallTypes.aresponses, True, [_retrieve_tool_definition()], False),
(CallTypes.aresponses, True, [_retrieve_tool_definition()], True),
(CallTypes.responses, True, [_responses_retrieve_tool_definition()], True),
(CallTypes.aresponses, False, [_retrieve_tool_definition()], False),
(CallTypes.anthropic_messages, True, [_retrieve_tool_definition()], False),
],
)
@pytest.mark.asyncio
async def test_pre_call_deployment_hook_converts_stream_only_for_ccr_chat_completions(
async def test_pre_call_deployment_hook_converts_stream_only_for_ccr_chat_completions_and_responses(
guardrail: HeadroomGuardrail,
call_type: CallTypes,
stream: bool,
@ -1986,6 +2101,22 @@ async def test_pre_call_deployment_hook_converts_stream_only_for_ccr_chat_comple
assert kwargs["stream"] is True
@pytest.mark.asyncio
async def test_pre_call_deployment_hook_leaves_background_streams_alone(guardrail: HeadroomGuardrail):
kwargs = {
"model": "gpt-4o",
"stream": True,
"background": True,
"tools": [_responses_retrieve_tool_definition()],
}
result = await guardrail.async_pre_call_deployment_hook(kwargs=kwargs, call_type=CallTypes.aresponses)
assert result is kwargs
assert HEADROOM_CONVERTED_STREAM_KEY not in kwargs
assert kwargs["stream"] is True
@pytest.mark.asyncio
async def test_pre_call_deployment_hook_still_compresses_for_deployment_level_configs(
guardrail: HeadroomGuardrail,
@ -2094,6 +2225,117 @@ async def test_streaming_chat_completion_resolves_ccr_retrieval_end_to_end(
assert not any(key.startswith("_headroom_interception") for key in followup_body)
@pytest.mark.asyncio
async def test_streaming_responses_resolves_ccr_retrieval_end_to_end(
guardrail: HeadroomGuardrail,
respx_mock: respx.MockRouter,
monkeypatch: pytest.MonkeyPatch,
):
"""Regression test for LIT-6481: streaming /v1/responses must resolve the
retrieve tool call server-side exactly like streaming /chat/completions does,
instead of streaming a headroom_retrieve function_call to the client."""
original_content = "the full uncompressed document"
final_answer = "the document says hello"
guardrail._issued_hashes_by_call_id["ccr-call-id"] = (
frozenset({CCR_HASH}),
time.monotonic() + 999,
)
monkeypatch.setenv("OPENAI_API_KEY", "sk-test")
monkeypatch.setattr(litellm, "callbacks", [guardrail])
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
upstream = respx_mock.post("https://api.openai.com/v1/responses").mock(
side_effect=[
httpx.Response(200, json=_openai_responses_retrieve_call_payload()),
httpx.Response(200, json=_openai_responses_text_payload(final_answer)),
]
)
with patch.object(
guardrail.async_handler,
"get",
new_callable=AsyncMock,
return_value=_make_retrieve_response(original_content),
) as mock_get:
response = await litellm.aresponses(
model="openai/gpt-4o",
input=[{"role": "user", "content": f"summarize hash={CCR_HASH}"}],
tools=[_responses_retrieve_tool_definition()],
stream=True,
litellm_call_id="ccr-call-id",
)
events = [event async for event in response]
streamed_text = "".join(
getattr(event, "delta", "") for event in events if getattr(event, "type", None) == "response.output_text.delta"
)
assert streamed_text == final_answer
assert not any("function_call" in str(getattr(event, "type", "")) for event in events)
assert not any(
getattr(getattr(event, "item", None), "type", None) == "function_call" for event in events
)
mock_get.assert_called_once()
assert CCR_HASH in (mock_get.call_args.kwargs.get("url") or mock_get.call_args.args[0])
assert len(upstream.calls) == 2
followup_body = json.loads(upstream.calls[1].request.content)
assert not followup_body.get("stream")
assert original_content in json.dumps(followup_body["input"])
assert not any(key.startswith("_headroom_interception") for key in followup_body)
def test_sync_streaming_responses_resolves_ccr_retrieval_end_to_end(
guardrail: HeadroomGuardrail,
respx_mock: respx.MockRouter,
monkeypatch: pytest.MonkeyPatch,
):
"""The synchronous responses() path converts the stream the same way, so it
must hand back a stream iterator with the resolved answer rather than the
completed response object."""
original_content = "the full uncompressed document"
final_answer = "the document says hello"
guardrail._issued_hashes_by_call_id["ccr-call-id"] = (
frozenset({CCR_HASH}),
time.monotonic() + 999,
)
monkeypatch.setenv("OPENAI_API_KEY", "sk-test")
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
monkeypatch.setattr(litellm, "callbacks", [guardrail])
upstream = respx_mock.post("https://api.openai.com/v1/responses").mock(
side_effect=[
httpx.Response(200, json=_openai_responses_retrieve_call_payload()),
httpx.Response(200, json=_openai_responses_text_payload(final_answer)),
]
)
with patch.object(
guardrail.async_handler,
"get",
new_callable=AsyncMock,
return_value=_make_retrieve_response(original_content),
) as mock_get:
response = litellm.responses(
model="openai/gpt-4o",
input=[{"role": "user", "content": f"summarize hash={CCR_HASH}"}],
tools=[_responses_retrieve_tool_definition()],
stream=True,
litellm_call_id="ccr-call-id",
)
events = list(response)
streamed_text = "".join(
getattr(event, "delta", "") for event in events if getattr(event, "type", None) == "response.output_text.delta"
)
assert streamed_text == final_answer
assert not any(
getattr(getattr(event, "item", None), "type", None) == "function_call" for event in events
)
mock_get.assert_called_once()
assert len(upstream.calls) == 2
assert not json.loads(upstream.calls[1].request.content).get("stream")
# ---------------------------------------------------------------------------
# LIT-5018: the turn the model is being asked to act on is never compressed.
#

View file

@ -25,6 +25,7 @@ from litellm.types.utils import (
ModelResponse,
ModelResponseStream,
StreamingChoices,
Usage,
)
CHAT_COMPLETION_ID = "chatcmpl-77d33d09-effa-4cd2-9c0d-c742d4358256"
@ -527,6 +528,20 @@ async def test_streaming_response_id_falls_back_when_upstream_yields_nothing():
assert response_ids[0].startswith("resp_")
def test_completed_event_restores_usage_hidden_by_stream_options_none():
final_chunk = _chunk("", finish_reason="stop")
final_chunk._hidden_params = {"usage": Usage(prompt_tokens=117, completion_tokens=5, total_tokens=122)}
iterator = _build_iterator([_chunk("the document says hello"), final_chunk])
events = list(iterator)
completed = next(
event for event in events if getattr(event, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED
)
assert completed.response.usage.input_tokens == 117
assert completed.response.usage.output_tokens == 5
def test_object_tool_call_arguments_stream_as_valid_json():
"""A provider that sends decoded object arguments must still stream valid JSON.