mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix: address remaining review feedback
- Empty stdout/stderr now produces outputs=None (matching OpenAI parity)
instead of outputs=[{logs:""}], in both streaming and non-streaming paths
- Fix test fixture to use real Anthropic type "bash_code_execution_tool_result"
instead of "code_execution_tool_result"
- Add test for empty-output → outputs=None behavior
- Add unit tests for _extract_tool_result_output_items: Pydantic objects,
plain dicts (post-model_dump), empty/missing provider_specific_fields,
and in-place substitution preserving output ordering
This commit is contained in:
parent
4be1d76fd7
commit
5b3e84f383
4 changed files with 243 additions and 6 deletions
|
|
@ -716,6 +716,9 @@ class ModelResponseIterator:
|
|||
logs = str(content)
|
||||
tool_input = self._server_tool_inputs.get(call_id, {})
|
||||
code = tool_input.get("command", "") if isinstance(tool_input, dict) else ""
|
||||
log_outputs = (
|
||||
[OutputCodeInterpreterCallLog(type="logs", logs=logs)] if logs else None
|
||||
)
|
||||
results.append(
|
||||
OutputCodeInterpreterCall(
|
||||
type="code_interpreter_call",
|
||||
|
|
@ -723,7 +726,7 @@ class ModelResponseIterator:
|
|||
code=code,
|
||||
container_id=None,
|
||||
status="completed",
|
||||
outputs=[OutputCodeInterpreterCallLog(type="logs", logs=logs)],
|
||||
outputs=log_outputs,
|
||||
)
|
||||
)
|
||||
return results
|
||||
|
|
|
|||
|
|
@ -1778,6 +1778,11 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
logs = "".join(parts)
|
||||
else:
|
||||
logs = str(content)
|
||||
log_outputs = (
|
||||
[OutputCodeInterpreterCallLog(type="logs", logs=logs)]
|
||||
if logs
|
||||
else None
|
||||
)
|
||||
code_interpreter_results.append(
|
||||
OutputCodeInterpreterCall(
|
||||
type="code_interpreter_call",
|
||||
|
|
@ -1785,9 +1790,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
code=code_by_id.get(call_id, ""),
|
||||
container_id=container_id,
|
||||
status="completed",
|
||||
outputs=[
|
||||
OutputCodeInterpreterCallLog(type="logs", logs=logs)
|
||||
],
|
||||
outputs=log_outputs,
|
||||
)
|
||||
)
|
||||
provider_specific_fields["code_interpreter_results"] = (
|
||||
|
|
|
|||
|
|
@ -1410,10 +1410,10 @@ def test_streaming_code_execution_input_assembled_from_deltas():
|
|||
"type": "content_block_start",
|
||||
"index": 1,
|
||||
"content_block": {
|
||||
"type": "code_execution_tool_result",
|
||||
"type": "bash_code_execution_tool_result",
|
||||
"tool_use_id": "srvtoolu_01AAA",
|
||||
"content": {
|
||||
"type": "code_execution_result",
|
||||
"type": "bash_code_execution_result",
|
||||
"stdout": "hello\n",
|
||||
"stderr": "",
|
||||
"return_code": 0,
|
||||
|
|
@ -1445,3 +1445,71 @@ def test_streaming_code_execution_input_assembled_from_deltas():
|
|||
assert code_results[0].id == "srvtoolu_01AAA"
|
||||
assert code_results[0].code == "echo hello"
|
||||
assert code_results[0].outputs[0].logs == "hello\n"
|
||||
|
||||
|
||||
def test_empty_output_produces_null_outputs():
|
||||
"""
|
||||
When both stdout and stderr are empty, outputs should be None
|
||||
(matching OpenAI's native behavior) rather than [{logs: ""}].
|
||||
"""
|
||||
chunks = [
|
||||
{
|
||||
"type": "message_start",
|
||||
"message": {
|
||||
"id": "msg_01XYZ",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [],
|
||||
"usage": {"input_tokens": 100, "output_tokens": 1},
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": 0,
|
||||
"content_block": {
|
||||
"type": "server_tool_use",
|
||||
"id": "srvtoolu_01AAA",
|
||||
"name": "bash_code_execution",
|
||||
"input": {"command": "true"},
|
||||
},
|
||||
},
|
||||
{"type": "content_block_stop", "index": 0},
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": 1,
|
||||
"content_block": {
|
||||
"type": "bash_code_execution_tool_result",
|
||||
"tool_use_id": "srvtoolu_01AAA",
|
||||
"content": {
|
||||
"type": "bash_code_execution_result",
|
||||
"stdout": "",
|
||||
"stderr": "",
|
||||
"return_code": 0,
|
||||
},
|
||||
},
|
||||
},
|
||||
{"type": "content_block_stop", "index": 1},
|
||||
{
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": "end_turn"},
|
||||
"usage": {"output_tokens": 50},
|
||||
},
|
||||
]
|
||||
|
||||
iterator = ModelResponseIterator(None, sync_stream=True)
|
||||
|
||||
code_results = None
|
||||
for chunk in chunks:
|
||||
parsed = iterator.chunk_parser(chunk)
|
||||
psf = None
|
||||
if parsed.choices and parsed.choices[0].delta:
|
||||
psf = getattr(parsed.choices[0].delta, "provider_specific_fields", None)
|
||||
if psf and "code_interpreter_results" in psf:
|
||||
code_results = psf["code_interpreter_results"]
|
||||
|
||||
assert code_results is not None, "No code_interpreter_results emitted"
|
||||
assert len(code_results) == 1
|
||||
assert code_results[0].id == "srvtoolu_01AAA"
|
||||
assert (
|
||||
code_results[0].outputs is None
|
||||
), f"Expected outputs=None for empty execution, got {code_results[0].outputs}"
|
||||
|
|
|
|||
|
|
@ -0,0 +1,163 @@
|
|||
"""
|
||||
Tests for the Responses API _extract_tool_result_output_items path
|
||||
and the non-streaming _hidden_params propagation of code_interpreter_results.
|
||||
"""
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.responses.litellm_completion_transformation.transformation import (
|
||||
LiteLLMCompletionResponsesConfig,
|
||||
)
|
||||
from litellm.types.responses.main import (
|
||||
OutputCodeInterpreterCall,
|
||||
OutputCodeInterpreterCallLog,
|
||||
)
|
||||
from litellm.types.utils import Choices, Message, ModelResponse
|
||||
|
||||
|
||||
def _make_model_response(code_interpreter_results=None, provider_specific_fields=None):
|
||||
"""Helper to build a ModelResponse with provider_specific_fields on the message."""
|
||||
psf = provider_specific_fields or {}
|
||||
if code_interpreter_results is not None:
|
||||
psf["code_interpreter_results"] = code_interpreter_results
|
||||
msg = Message(content="test", provider_specific_fields=psf if psf else None)
|
||||
choice = Choices(index=0, message=msg, finish_reason="stop")
|
||||
resp = ModelResponse()
|
||||
resp.choices = [choice]
|
||||
return resp
|
||||
|
||||
|
||||
def test_extract_tool_result_output_items_from_pydantic_objects():
|
||||
"""Non-streaming path: code_interpreter_results are Pydantic OutputCodeInterpreterCall objects."""
|
||||
items = [
|
||||
OutputCodeInterpreterCall(
|
||||
type="code_interpreter_call",
|
||||
id="srvtoolu_01AAA",
|
||||
code="echo hello",
|
||||
container_id=None,
|
||||
status="completed",
|
||||
outputs=[OutputCodeInterpreterCallLog(type="logs", logs="hello\n")],
|
||||
),
|
||||
OutputCodeInterpreterCall(
|
||||
type="code_interpreter_call",
|
||||
id="srvtoolu_01BBB",
|
||||
code="echo world",
|
||||
container_id=None,
|
||||
status="completed",
|
||||
outputs=[OutputCodeInterpreterCallLog(type="logs", logs="world\n")],
|
||||
),
|
||||
]
|
||||
resp = _make_model_response(code_interpreter_results=items)
|
||||
result = LiteLLMCompletionResponsesConfig._extract_tool_result_output_items(resp)
|
||||
assert len(result) == 2
|
||||
assert result[0].id == "srvtoolu_01AAA"
|
||||
assert result[1].id == "srvtoolu_01BBB"
|
||||
|
||||
|
||||
def test_extract_tool_result_output_items_from_dicts():
|
||||
"""Streaming path: after model_dump(), code_interpreter_results are plain dicts."""
|
||||
items = [
|
||||
{
|
||||
"type": "code_interpreter_call",
|
||||
"id": "srvtoolu_01AAA",
|
||||
"code": "echo hello",
|
||||
"container_id": None,
|
||||
"status": "completed",
|
||||
"outputs": [{"type": "logs", "logs": "hello\n"}],
|
||||
},
|
||||
]
|
||||
resp = _make_model_response(code_interpreter_results=items)
|
||||
result = LiteLLMCompletionResponsesConfig._extract_tool_result_output_items(resp)
|
||||
assert len(result) == 1
|
||||
assert result[0]["id"] == "srvtoolu_01AAA"
|
||||
|
||||
|
||||
def test_extract_tool_result_output_items_empty():
|
||||
"""No code_interpreter_results → empty list."""
|
||||
resp = _make_model_response()
|
||||
result = LiteLLMCompletionResponsesConfig._extract_tool_result_output_items(resp)
|
||||
assert result == []
|
||||
|
||||
|
||||
def test_extract_tool_result_output_items_no_provider_specific_fields():
|
||||
"""Message with no provider_specific_fields → empty list."""
|
||||
msg = Message(content="test")
|
||||
choice = Choices(index=0, message=msg, finish_reason="stop")
|
||||
resp = ModelResponse()
|
||||
resp.choices = [choice]
|
||||
result = LiteLLMCompletionResponsesConfig._extract_tool_result_output_items(resp)
|
||||
assert result == []
|
||||
|
||||
|
||||
def test_in_place_substitution_preserves_ordering():
|
||||
"""
|
||||
function_call items matching code_interpreter_results should be replaced
|
||||
in-place, preserving the original output ordering.
|
||||
|
||||
Simulates: [message, function_call(exec1), function_call(regular), function_call(exec2)]
|
||||
Expected: [message, code_interpreter_call(exec1), function_call(regular), code_interpreter_call(exec2)]
|
||||
"""
|
||||
code_results = [
|
||||
OutputCodeInterpreterCall(
|
||||
type="code_interpreter_call",
|
||||
id="srvtoolu_01AAA",
|
||||
code="echo first",
|
||||
container_id=None,
|
||||
status="completed",
|
||||
outputs=[OutputCodeInterpreterCallLog(type="logs", logs="first\n")],
|
||||
),
|
||||
OutputCodeInterpreterCall(
|
||||
type="code_interpreter_call",
|
||||
id="srvtoolu_01CCC",
|
||||
code="echo third",
|
||||
container_id=None,
|
||||
status="completed",
|
||||
outputs=[OutputCodeInterpreterCallLog(type="logs", logs="third\n")],
|
||||
),
|
||||
]
|
||||
resp = _make_model_response(code_interpreter_results=code_results)
|
||||
|
||||
# Build a mock responses_output list with interleaved items
|
||||
class MockItem:
|
||||
def __init__(self, type, call_id=None):
|
||||
self.type = type
|
||||
self.call_id = call_id
|
||||
|
||||
msg_item = MockItem(type="message")
|
||||
fc_exec1 = MockItem(type="function_call", call_id="srvtoolu_01AAA")
|
||||
fc_regular = MockItem(type="function_call", call_id="srvtoolu_01BBB")
|
||||
fc_exec2 = MockItem(type="function_call", call_id="srvtoolu_01CCC")
|
||||
|
||||
responses_output = [msg_item, fc_exec1, fc_regular, fc_exec2]
|
||||
|
||||
# Apply the same logic as _transform_chat_completion_choices_to_responses_output
|
||||
tool_result_items = (
|
||||
LiteLLMCompletionResponsesConfig._extract_tool_result_output_items(resp)
|
||||
)
|
||||
if tool_result_items:
|
||||
result_by_id = {
|
||||
(item.get("id") if isinstance(item, dict) else item.id): item
|
||||
for item in tool_result_items
|
||||
}
|
||||
replaced_ids = set(result_by_id.keys())
|
||||
responses_output = [
|
||||
(
|
||||
result_by_id[getattr(item, "call_id", None)]
|
||||
if (
|
||||
getattr(item, "type", None) == "function_call"
|
||||
and getattr(item, "call_id", None) in replaced_ids
|
||||
)
|
||||
else item
|
||||
)
|
||||
for item in responses_output
|
||||
]
|
||||
|
||||
# Verify ordering: message, code_interpreter(AAA), function_call(BBB), code_interpreter(CCC)
|
||||
assert len(responses_output) == 4
|
||||
assert responses_output[0].type == "message"
|
||||
assert responses_output[1].type == "code_interpreter_call"
|
||||
assert responses_output[1].id == "srvtoolu_01AAA"
|
||||
assert responses_output[2].type == "function_call"
|
||||
assert responses_output[2].call_id == "srvtoolu_01BBB"
|
||||
assert responses_output[3].type == "code_interpreter_call"
|
||||
assert responses_output[3].id == "srvtoolu_01CCC"
|
||||
Loading…
Add table
Reference in a new issue