This commit is contained in:
devin-ai-integration[bot] 2026-08-27 23:19:21 +00:00 • committed by GitHub
commit 9f5866333b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 258 additions and 15 deletions

View file

@ -1,6 +1,7 @@
# this is a patch to allow for agentic loops covering llm_http_handler.py and openai sdk based calling flows for the .completion() api
import json
from collections.abc import Mapping
from typing import Final, cast
from litellm._logging import verbose_logger
@ -9,8 +10,11 @@ from litellm.litellm_core_utils.agentic_loop_settings import (
DEFAULT_MAX_AGENTIC_LOOPS,
validated_max_agentic_loops,
)
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObject
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
from litellm.types.integrations.custom_logger import (
CHAT_COMPLETION_AGENTIC_SURFACE,
HEADROOM_CONVERTED_STREAM_KEY,
NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES,
AgenticLoopPlan,
AgenticLoopRequestPatch,
@ -50,6 +54,12 @@ def _post_hook_overridden(callback: CustomLogger) -> bool:
return getattr(func, "__func__", func) is not getattr(base, "__func__", base)
def _converted_stream_requested(kwargs: Mapping[str, object]) -> bool:
return bool(
kwargs.get("_code_interpreter_interception_converted_stream") or kwargs.get(HEADROOM_CONVERTED_STREAM_KEY)
)
def _coerce_int(value: object, default: int) -> int:
return int(value) if isinstance(value, (int, str)) else default
@ -87,16 +97,24 @@ def _check_agentic_loop_safety(
return fingerprint
def _wrap_response_as_fake_stream(response: object) -> object:
if getattr(response, "object", None) == "chat.completion.chunk":
def _wrap_response_as_fake_stream(
response: object,
*,
model: str,
custom_llm_provider: str,
logging_obj: object,
) -> object:
if isinstance(response, CustomStreamWrapper):
return response
if not hasattr(response, "choices"):
if not isinstance(response, ModelResponse) or not isinstance(logging_obj, LiteLLMLoggingObject):
return response
from litellm.llms.base_llm.base_model_iterator import (
convert_model_response_to_streaming,
)
return convert_model_response_to_streaming(cast(ModelResponse, response))
return CustomStreamWrapper(
completion_stream=MockResponseIterator(model_response=response),
model=model,
custom_llm_provider=custom_llm_provider,
logging_obj=logging_obj,
)
def _add_agentic_loop_metadata(kwargs_for_followup: dict[str, object]) -> None:
@ -177,8 +195,13 @@ async def _execute_chat_completion_agentic_plan(
model,
str(e),
)
if kwargs.get("_code_interpreter_interception_converted_stream") and not depth:
return _wrap_response_as_fake_stream(response_followup)
if _converted_stream_requested(kwargs) and not depth:
return _wrap_response_as_fake_stream(
response_followup,
model=model,
custom_llm_provider=custom_llm_provider,
logging_obj=logging_obj,
)
return response_followup
finally:
try:
@ -302,9 +325,14 @@ async def maybe_run_chat_completion_agentic_loop(
str(e),
)
if kwargs.get("_code_interpreter_interception_converted_stream") and not depth and hasattr(response, "choices"):
if _converted_stream_requested(kwargs) and not depth:
return cast(
"ModelResponse | CustomStreamWrapper",
_wrap_response_as_fake_stream(response),
_wrap_response_as_fake_stream(
response,
model=model,
custom_llm_provider=custom_llm_provider,
logging_obj=logging_obj,
),
)
return None

View file

@ -10238,6 +10238,18 @@
"description": "AWS Bedrock runtime endpoint URL",
"title": "Aws Bedrock Runtime Endpoint"
},
"aws_external_id": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"description": "External ID required by the target role's trust policy on sts:AssumeRole",
"title": "Aws External Id"
},
"aws_profile_name": {
"anyOf": [
{
@ -25237,6 +25249,9 @@
},
{
"$ref": "#/components/schemas/ChatCompletionImageObject"
},
{
"$ref": "#/components/schemas/ChatCompletionToolReferenceObject"
}
]
},
@ -25324,6 +25339,26 @@
"title": "ChatCompletionToolParamFunctionChunk",
"type": "object"
},
"ChatCompletionToolReferenceObject": {
"description": "Anthropic tool-search result block, carried through untouched so it survives a round trip.",
"properties": {
"tool_name": {
"title": "Tool Name",
"type": "string"
},
"type": {
"const": "tool_reference",
"title": "Type",
"type": "string"
}
},
"required": [
"type",
"tool_name"
],
"title": "ChatCompletionToolReferenceObject",
"type": "object"
},
"ChatCompletionUserMessage": {
"properties": {
"cache_control": {

View file

@ -37,8 +37,12 @@ from litellm.proxy.guardrails.guardrail_hooks.content_text import (
from litellm.proxy.spend_tracking.compression_savings import HEADROOM_GUARDRAIL_PROVIDER
from litellm.secret_managers.main import get_secret_str
from litellm.types.guardrails import GuardrailEventHooks, Mode
from litellm.types.integrations.custom_logger import AgenticLoopPlan, AgenticLoopRequestPatch
from litellm.types.utils import GenericGuardrailAPIInputs
from litellm.types.integrations.custom_logger import (
HEADROOM_CONVERTED_STREAM_KEY,
AgenticLoopPlan,
AgenticLoopRequestPatch,
)
from litellm.types.utils import CallTypes, GenericGuardrailAPIInputs
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
@ -713,6 +717,23 @@ class HeadroomGuardrail(CustomGuardrail):
return {**inputs, "structured_messages": compressed, "tools": merged_tools} # pyright: ignore[reportReturnType]
async def async_pre_call_deployment_hook(
self,
kwargs: Mapping[str, Any],
call_type: CallTypes | None,
) -> dict[str, Any] | None: # mutable-ok: overrides CustomLogger hook whose contract is a plain dict
if call_type not in (CallTypes.completion, CallTypes.acompletion):
return None
if not kwargs.get("stream"):
return None
if not has_headroom_retrieve_tool(kwargs.get("tools")):
return None
return { # mutable-ok: the hook contract is a plain dict the router merges into the request kwargs
**kwargs,
"stream": False,
HEADROOM_CONVERTED_STREAM_KEY: True,
}
async def async_should_run_agentic_loop(
self,
response: Any,

View file

@ -251,6 +251,7 @@ _UNTRUSTED_ROOT_CONTROL_FIELDS: Final = (
"_code_interpreter_interception_converted_stream",
"_code_interpreter_interception_sandbox_key",
"_code_interpreter_interception_session_scoped",
"_headroom_interception_converted_stream",
"max_agentic_loops",
# Recomputed below from the actual caller-controlled timeout sources (headers and
# body fields); a client-forged value here would let a request either dodge cooldown

View file

@ -5,8 +5,14 @@ from pydantic import BaseModel, Field
CHAT_COMPLETION_AGENTIC_SURFACE: Final = "chat_completions"
RESPONSES_AGENTIC_SURFACE: Final = "responses"
CODE_INTERPRETER_INTERCEPTION_PREFIX: Final = "_code_interpreter_interception"
HEADROOM_INTERCEPTION_PREFIX: Final = "_headroom_interception"
HEADROOM_CONVERTED_STREAM_KEY: Final = f"{HEADROOM_INTERCEPTION_PREFIX}_converted_stream"
NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES: Final = frozenset(
("_websearch_interception", "_compression_interception")
(
"_websearch_interception",
"_compression_interception",
HEADROOM_INTERCEPTION_PREFIX,
)
)
INTERCEPTION_INTERNAL_PREFIXES: Final = frozenset(
(

View file

@ -3495,6 +3495,7 @@ agentic_loop_internal_litellm_params: Final = [
"_code_interpreter_interception_converted_stream",
"_websearch_interception_emit_native_blocks",
"_websearch_interception_converted_stream",
"_headroom_interception_converted_stream",
]
# Proxy-owned callback credentials, stamped from admin-configured team/key callback

View file

@ -17,14 +17,18 @@ Tests cover:
- CCR: headroom_retrieve tool injected when compressed messages contain hashes
- CCR: async_should_run_agentic_loop returns True when response has headroom_retrieve tool calls
- CCR: async_build_agentic_loop_plan calls retrieve endpoint and builds follow-up messages
- CCR: streaming /chat/completions is converted to a non-streaming call so the agentic
loop resolves the retrieve tool call, then fake-streamed back to the client
"""
import json
import time
from typing import Optional
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
import respx
from fastapi import HTTPException
import litellm
@ -38,7 +42,11 @@ from litellm.proxy.guardrails.guardrail_hooks.headroom.headroom import (
from litellm.proxy.spend_tracking.compression_savings import (
extract_compression_saved_tokens,
)
from litellm.types.utils import GenericGuardrailAPIInputs
from litellm.types.integrations.custom_logger import HEADROOM_CONVERTED_STREAM_KEY
from litellm.types.utils import (
CallTypes,
GenericGuardrailAPIInputs,
)
FAKE_API_BASE = "https://headroom.example.com"
FAKE_API_KEY = "test-key"
@ -1893,6 +1901,144 @@ async def test_fail_open_returns_original_parts_shapes():
assert [m["content"] for m in messages] == [m["content"] for m in PARTS_MESSAGES]
CCR_HASH = "b573993006976af767214fac"
def _retrieve_tool_definition() -> dict:
return {
"type": "function",
"function": {
"name": HEADROOM_RETRIEVE_TOOL_NAME,
"description": "retrieve compressed content",
"parameters": {"type": "object", "properties": {"hash": {"type": "string"}}},
},
}
def _openai_completion_payload(message: dict, finish_reason: str) -> dict:
return {
"id": "chatcmpl-ccr",
"object": "chat.completion",
"created": 1700000000,
"model": "gpt-4o",
"choices": [{"index": 0, "message": message, "finish_reason": finish_reason}],
"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
}
def _openai_tool_call_payload() -> dict:
return _openai_completion_payload(
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_ccr",
"type": "function",
"function": {
"name": HEADROOM_RETRIEVE_TOOL_NAME,
"arguments": json.dumps({"hash": CCR_HASH}),
},
}
],
},
"tool_calls",
)
def _openai_text_payload(content: str) -> dict:
return _openai_completion_payload({"role": "assistant", "content": content}, "stop")
@pytest.mark.parametrize(
"call_type, stream, tools, expect_conversion",
[
(CallTypes.acompletion, True, [_retrieve_tool_definition()], True),
(CallTypes.completion, True, [_retrieve_tool_definition()], True),
(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.anthropic_messages, True, [_retrieve_tool_definition()], False),
],
)
@pytest.mark.asyncio
async def test_pre_call_deployment_hook_converts_stream_only_for_ccr_chat_completions(
guardrail: HeadroomGuardrail,
call_type: CallTypes,
stream: bool,
tools: Optional[list],
expect_conversion: bool,
):
kwargs = {"model": "gpt-4o", "stream": stream, "tools": tools}
result = await guardrail.async_pre_call_deployment_hook(kwargs=kwargs, call_type=call_type)
if not expect_conversion:
assert result is None
assert kwargs["stream"] is stream
return
assert result is not None
assert result["stream"] is False
assert result[HEADROOM_CONVERTED_STREAM_KEY] is True
assert kwargs["stream"] is True
@pytest.mark.asyncio
async def test_streaming_chat_completion_resolves_ccr_retrieval_end_to_end(
guardrail: HeadroomGuardrail,
respx_mock: respx.MockRouter,
monkeypatch: pytest.MonkeyPatch,
):
"""Regression test for streaming /chat/completions: the retrieve tool call the
model emits must be resolved by the agentic loop instead of being streamed back
to a client that never declared the tool."""
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/chat/completions").mock(
side_effect=[
httpx.Response(200, json=_openai_tool_call_payload()),
httpx.Response(200, json=_openai_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.acompletion(
model="openai/gpt-4o",
messages=[{"role": "user", "content": f"summarize hash={CCR_HASH}"}],
tools=[_retrieve_tool_definition()],
stream=True,
litellm_call_id="ccr-call-id",
)
chunks = [chunk async for chunk in response]
streamed_text = "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices)
assert streamed_text == final_answer
assert not any(chunk.choices and chunk.choices[0].delta.tool_calls for chunk in chunks)
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["messages"])
assert not any(key.startswith("_headroom_interception") for key in followup_body)
# ---------------------------------------------------------------------------
# LIT-5018: the turn the model is being asked to act on is never compressed.
#

View file

@ -29466,6 +29466,11 @@ export interface components {
* @description AWS Bedrock runtime endpoint URL
*/
aws_bedrock_runtime_endpoint?: string | null;
/**
* Aws External Id
* @description External ID required by the target role's trust policy on sts:AssumeRole
*/
aws_external_id?: string | null;
/**
* Aws Profile Name
* @description AWS profile name for credential retrieval