mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge a14daf18d6 into 44d84360fb
This commit is contained in:
commit
9f5866333b
8 changed files with 258 additions and 15 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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": {
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
#
|
||||
|
|
|
|||
5
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
5
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue