mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
commit
80250807db
13 changed files with 1156 additions and 103 deletions
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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``"""
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
#
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue