v0 implementation of context_management (#29090)

This commit is contained in:
JT 2026-05-27 16:55:35 -07:00 • committed by GitHub
parent e529e3856e
commit 7a0bee5107
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 215 additions and 122 deletions

View file

@ -1393,10 +1393,10 @@ def convert_to_gemini_tool_call_invoke(
if tool_calls is not None:
for idx, tool in enumerate(tool_calls):
if "function" in tool:
gemini_function_call: Optional[
VertexFunctionCall
] = _gemini_tool_call_invoke_helper(
function_call_params=tool["function"]
gemini_function_call: Optional[VertexFunctionCall] = (
_gemini_tool_call_invoke_helper(
function_call_params=tool["function"]
)
)
if gemini_function_call is not None:
part_dict: VertexPartType = {
@ -1574,9 +1574,7 @@ def convert_to_gemini_tool_call_result( # noqa: PLR0915
file_data = (
file_content.get("file_data", "")
if isinstance(file_content, dict)
else file_content
if isinstance(file_content, str)
else ""
else file_content if isinstance(file_content, str) else ""
)
if file_data:
@ -2081,9 +2079,9 @@ def _sanitize_empty_text_content(
if isinstance(content, str):
if not content or not content.strip():
message = cast(AllMessageValues, dict(message)) # Make a copy
message[
"content"
] = "[System: Empty message content sanitised to satisfy protocol]"
message["content"] = (
"[System: Empty message content sanitised to satisfy protocol]"
)
verbose_logger.debug(
f"_sanitize_empty_text_content: Replaced empty text content in {message.get('role')} message"
)
@ -2423,9 +2421,9 @@ def anthropic_messages_pt( # noqa: PLR0915
# Convert ChatCompletionImageUrlObject to dict if needed
image_url_value = m["image_url"]
if isinstance(image_url_value, str):
image_url_input: Union[
str, dict[str, Any]
] = image_url_value
image_url_input: Union[str, dict[str, Any]] = (
image_url_value
)
else:
# ChatCompletionImageUrlObject or dict case - convert to dict
image_url_input = {
@ -2452,9 +2450,9 @@ def anthropic_messages_pt( # noqa: PLR0915
)
if "cache_control" in _content_element:
_anthropic_content_element[
"cache_control"
] = _content_element["cache_control"]
_anthropic_content_element["cache_control"] = (
_content_element["cache_control"]
)
user_content.append(_anthropic_content_element)
elif m.get("type", "") == "text":
m = cast(ChatCompletionTextObject, m)
@ -2514,9 +2512,9 @@ def anthropic_messages_pt( # noqa: PLR0915
)
if "cache_control" in _content_element:
_anthropic_content_text_element[
"cache_control"
] = _content_element["cache_control"]
_anthropic_content_text_element["cache_control"] = (
_content_element["cache_control"]
)
user_content.append(_anthropic_content_text_element)
@ -2649,9 +2647,9 @@ def anthropic_messages_pt( # noqa: PLR0915
original_content_element=dict(assistant_content_block),
)
if "cache_control" in _content_element:
_anthropic_text_content_element[
"cache_control"
] = _content_element["cache_control"]
_anthropic_text_content_element["cache_control"] = (
_content_element["cache_control"]
)
text_element = _anthropic_text_content_element
# Interleave: each thinking block precedes its server tool group.
@ -2811,9 +2809,9 @@ def anthropic_messages_pt( # noqa: PLR0915
)
if "cache_control" in _content_element:
_anthropic_text_content_element[
"cache_control"
] = _content_element["cache_control"]
_anthropic_text_content_element["cache_control"] = (
_content_element["cache_control"]
)
assistant_content.append(_anthropic_text_content_element)
@ -5255,9 +5253,7 @@ def default_response_schema_prompt(response_schema: dict) -> str:
prompt_str = """Use this JSON schema:
```json
{}
```""".format(
response_schema
)
```""".format(response_schema)
return prompt_str

View file

@ -4,35 +4,10 @@ model_list:
model: openai/gpt-4o
api_key: os.environ/OPENAI_API_KEY
- model_name: text-embedding-3-small
- model_name: claude-haiku
litellm_params:
model: openai/text-embedding-3-small
api_key: os.environ/OPENAI_API_KEY
- model_name: bedrock-claude-sonnet-3.5
litellm_params:
model: "bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0"
aws_region_name: "us-east-1"
- model_name: bedrock-claude-sonnet-4
litellm_params:
model: "bedrock/us.anthropic.claude-sonnet-4-20250514-v1:0"
aws_region_name: "us-east-1"
- model_name: bedrock-claude-sonnet-4.5
litellm_params:
model: "bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0"
aws_region_name: "us-east-1"
- model_name: bedrock-claude-opus-4.5
litellm_params:
model: "bedrock/converse/us.anthropic.claude-opus-4-5-20251101-v1:0"
aws_region_name: "us-east-1"
- model_name: bedrock-nova-premier
litellm_params:
model: "bedrock/us.amazon.nova-premier-v1:0"
aws_region_name: "us-east-1"
model: claude-haiku-4-5-20251001
api_key: os.environ/ANTHROPIC_API_KEY
# MCP Server Configuration
mcp_servers:

View file

@ -2,8 +2,12 @@
Handler for transforming responses api requests to litellm.completion requests
"""
from typing import Any, Coroutine, Dict, Optional, Union
from typing import Any, Coroutine, Dict, List, Optional, Union
import base64
import random
import string
import json
import litellm
from litellm.responses.litellm_completion_transformation.streaming_iterator import (
LiteLLMCompletionStreamingIterator,
@ -29,6 +33,7 @@ class LiteLLMCompletionTransformationHandler:
custom_llm_provider: Optional[str] = None,
_is_async: bool = False,
stream: Optional[bool] = None,
context_management: Optional[List[Dict[str, Any]]] = None,
extra_headers: Optional[Dict[str, Any]] = None,
**kwargs,
) -> Union[
@ -38,14 +43,16 @@ class LiteLLMCompletionTransformationHandler:
Any, Any, Union[ResponsesAPIResponse, BaseResponsesAPIStreamingIterator]
],
]:
litellm_completion_request: dict = LiteLLMCompletionResponsesConfig.transform_responses_api_request_to_chat_completion_request(
model=model,
input=input,
responses_api_request=responses_api_request,
custom_llm_provider=custom_llm_provider,
stream=stream,
extra_headers=extra_headers,
**kwargs,
litellm_completion_request: dict = (
LiteLLMCompletionResponsesConfig.transform_responses_api_request_to_chat_completion_request(
model=model,
input=input,
responses_api_request=responses_api_request,
custom_llm_provider=custom_llm_provider,
stream=stream,
extra_headers=extra_headers,
**kwargs,
)
)
if _is_async:
@ -53,6 +60,7 @@ class LiteLLMCompletionTransformationHandler:
litellm_completion_request=litellm_completion_request,
request_input=input,
responses_api_request=responses_api_request,
context_management=context_management,
**kwargs,
)
@ -68,10 +76,12 @@ class LiteLLMCompletionTransformationHandler:
)
if isinstance(litellm_completion_response, ModelResponse):
responses_api_response: ResponsesAPIResponse = LiteLLMCompletionResponsesConfig.transform_chat_completion_response_to_responses_api_response(
chat_completion_response=litellm_completion_response,
request_input=input,
responses_api_request=responses_api_request,
responses_api_response: ResponsesAPIResponse = (
LiteLLMCompletionResponsesConfig.transform_chat_completion_response_to_responses_api_response(
chat_completion_response=litellm_completion_response,
request_input=input,
responses_api_request=responses_api_request,
)
)
return responses_api_response
@ -94,6 +104,7 @@ class LiteLLMCompletionTransformationHandler:
litellm_completion_request: dict,
request_input: Union[str, ResponseInputParam],
responses_api_request: ResponsesAPIOptionalRequestParams,
context_management: Optional[List[Dict[str, Any]]] = None,
**kwargs,
) -> Union[ResponsesAPIResponse, BaseResponsesAPIStreamingIterator]:
previous_response_id: Optional[str] = responses_api_request.get(
@ -105,10 +116,27 @@ class LiteLLMCompletionTransformationHandler:
litellm_completion_request=litellm_completion_request,
)
# breakpoint()
compacted = False
if context_management:
litellm_completion_request["messages"], compacted = (
await LiteLLMCompletionResponsesConfig._transform_context_management(
model=litellm_completion_request.get("model", ""),
input=litellm_completion_request.get("messages", []),
context_management=context_management,
**kwargs,
)
)
if compacted:
compact_input = litellm_completion_request["messages"][0]
acompletion_args = {}
acompletion_args.update(kwargs)
acompletion_args.update(litellm_completion_request)
if "context_management" in acompletion_args:
del acompletion_args["context_management"]
litellm_completion_response: Union[
ModelResponse, litellm.CustomStreamWrapper
] = await litellm.acompletion(
@ -116,11 +144,26 @@ class LiteLLMCompletionTransformationHandler:
)
if isinstance(litellm_completion_response, ModelResponse):
responses_api_response: ResponsesAPIResponse = LiteLLMCompletionResponsesConfig.transform_chat_completion_response_to_responses_api_response(
chat_completion_response=litellm_completion_response,
request_input=request_input,
responses_api_request=responses_api_request,
responses_api_response: ResponsesAPIResponse = (
LiteLLMCompletionResponsesConfig.transform_chat_completion_response_to_responses_api_response(
chat_completion_response=litellm_completion_response,
request_input=request_input,
responses_api_request=responses_api_request,
)
)
if compacted:
responses_api_response.output.append(
{
"id": "cmp_"
+ "".join(
random.choices(string.ascii_letters + string.digits, k=20)
),
"type": "compaction",
"encrypted_content": base64.b64encode(
json.dumps(compact_input).encode("utf-8")
).decode("utf-8"),
}
)
return responses_api_response

View file

@ -3,6 +3,8 @@ Handles transforming from Responses API -> LiteLLM completion (Chat Completion
"""
from collections.abc import Sequence
import json
import base64
from typing import Any, Dict, List, Literal, Optional, Set, Tuple, Union, cast
from openai.types.responses import ResponseFunctionToolCall
@ -99,6 +101,7 @@ class LiteLLMCompletionResponsesConfig:
"tools",
"top_p",
"user",
"context_management",
]
@staticmethod
@ -152,41 +155,78 @@ class LiteLLMCompletionResponsesConfig:
# Return as-is for unknown formats
return tool_choice
type Messages = List[
Union[
AllMessageValues,
GenericChatCompletionMessage,
ChatCompletionMessageToolCall,
ChatCompletionResponseMessage,
Message,
]
]
@staticmethod
async def _compact_input(
model: str,
input: Union[str, ResponseInputParam],
) -> Union[str, ResponseInputParam]:
input: Messages,
**kwargs,
) -> Messages:
"""
Make a 2nd LLM call to compact/summarize the conversation history.
Returns the compacted input as a single user message list.
"""
import litellm
summary_args = {}
# summary_args.update(kwargs)
if "context_management" in summary_args:
del summary_args["context_management"]
summary_args["model"] = model
conversation_history = json.dumps(input)
summary_args["messages"] = [
{
"role": "user",
"content": f"Provide a summary of the conversation below in your response directly with no special formatting. Focus on the key points and important details. Be concise but include relevant context that would help answer the next user message. The summary should be in a compact format that captures the essence of the conversation history.\n\n<conversation>\n{conversation_history}\n</conversation>",
}
]
# breakpoint()
summary_response = litellm.completion(**summary_args)
if not isinstance(summary_response, ModelResponse):
raise ValueError("Expected a ModelResponse object")
summary = (
summary_response.choices[0].message.content
if summary_response.choices
else ""
)
return [
{
"type": "message",
"role": "user",
"content": "",
"role": "developer",
"content": f"The following is a compacted summary of the prior conversation. Treat it as historical context, not as a new user request. Use it only to answer the latest user message.\n\n<conversation_summary>\nEarlier conversation: {summary}\n</conversation_summary>",
}
]
@staticmethod
def _cheap_token_counter(input: Union[str, ResponseInputParam]) -> int:
def _cheap_token_counter(input: Messages) -> int:
"""
Cheaply estimate the token count of the input.
~4 chars per token for strings; for message lists, stringify first.
Note: 1. not using an actual tokenizer
2. transformation from ResponseInputParam to str is ugly and not precise
"""
pass
json_str = json.dumps(input)
return len(json_str) // 4
@staticmethod
async def _transform_context_management(
model: str,
input: Union[str, ResponseInputParam],
input: Messages,
context_management: Optional[List[Dict[str, Any]]],
) -> Union[str, ResponseInputParam]:
**kwargs,
) -> Tuple[Messages, bool]:
"""
Handle context_management compaction for the Responses API -> Chat Completion path.
@ -195,17 +235,42 @@ class LiteLLMCompletionResponsesConfig:
Returns the (possibly compacted) input.
"""
pass
if not context_management:
return input, False
ctx_mgmt_type = context_management[0].get("type", "")
ctx_mgmt_threshold = context_management[0].get("compact_threshold", 0)
if not ctx_mgmt_type or not ctx_mgmt_threshold:
return input, False
@staticmethod
def should_execute_compaction(
input_token_size: int,
context_management: Optional[List[Dict[str, Any]]],
) -> bool:
"""
Check if compaction should be executed
"""
pass
# Note: for now skip compaction if input is a string or only has 1 message
if len(input) <= 1:
return input, False
# Note: exclude the last user message from token count since that's the new input we want to preserve in full
input_token_size = LiteLLMCompletionResponsesConfig._cheap_token_counter(
input[:-1]
)
if input_token_size < ctx_mgmt_threshold:
return input, False
compacted_input = await LiteLLMCompletionResponsesConfig._compact_input(
model=model,
input=input[:-1],
**kwargs,
)
# Append the latest user message back to the compacted history
compacted_input.append(input[-1])
return compacted_input, True
# @staticmethod
# def should_execute_compaction(
# input_token_size: int,
# context_management: Optional[List[Dict[str, Any]]],
# ) -> bool:
# """
# Check if compaction should be executed
# """
# pass
@staticmethod
def transform_responses_api_request_to_chat_completion_request(
@ -244,7 +309,7 @@ class LiteLLMCompletionResponsesConfig:
elif isinstance(reasoning_param, str):
# reasoning could be a string directly
reasoning_effort = reasoning_param
litellm_completion_request: dict = {
"messages": LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
input=input,
@ -1028,6 +1093,14 @@ class LiteLLMCompletionResponsesConfig:
return LiteLLMCompletionResponsesConfig._transform_responses_api_function_call_to_chat_completion_message(
function_call=input_item
)
elif input_item.get("type") == "compaction":
return [
json.loads(
base64.b64decode(input_item.get("encrypted_content")).decode(
"utf-8"
)
)
]
else:
content = input_item.get("content")
# Handle None content: Responses API allows None content, but GenericChatCompletionMessage requires content
@ -2177,9 +2250,9 @@ class LiteLLMCompletionResponsesConfig:
hasattr(completion_details, "reasoning_tokens")
and completion_details.reasoning_tokens is not None
):
output_details_dict[
"reasoning_tokens"
] = completion_details.reasoning_tokens
output_details_dict["reasoning_tokens"] = (
completion_details.reasoning_tokens
)
else:
output_details_dict["reasoning_tokens"] = 0

View file

@ -437,6 +437,7 @@ async def aresponses(
user: Optional[str] = None,
service_tier: Optional[str] = None,
safety_identifier: Optional[str] = None,
context_management: Optional[Iterable[Dict[str, Any]]] = None,
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
# The extra values given here take precedence over values defined on the client or passed to this method.
extra_headers: Optional[Dict[str, Any]] = None,
@ -450,6 +451,7 @@ async def aresponses(
"""
Async: Handles responses API requests by reusing the synchronous function
"""
# breakpoint()
local_vars = locals()
try:
loop = asyncio.get_event_loop()
@ -540,6 +542,7 @@ async def aresponses(
top_p=top_p,
truncation=truncation,
user=user,
context_management=context_management,
extra_headers=extra_headers,
extra_query=extra_query,
extra_body=extra_body,
@ -730,6 +733,7 @@ def responses(
user: Optional[str] = None,
service_tier: Optional[str] = None,
safety_identifier: Optional[str] = None,
context_management: Optional[List[Dict[str, Any]]] = None,
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
# The extra values given here take precedence over values defined on the client or passed to this method.
extra_headers: Optional[Dict[str, Any]] = None,
@ -919,6 +923,7 @@ def responses(
**emulated_kwargs,
)
# breakpoint()
if responses_api_provider_config is None:
return litellm_completion_transformation_handler.response_api_handler(
model=model,
@ -927,6 +932,7 @@ def responses(
custom_llm_provider=custom_llm_provider,
_is_async=_is_async,
stream=stream,
context_management=context_management,
extra_headers=extra_headers,
extra_body=extra_body,
**kwargs,
@ -1115,11 +1121,11 @@ def delete_responses(
raise ValueError("custom_llm_provider is required but passed as None")
# get provider config
responses_api_provider_config: Optional[
BaseResponsesAPIConfig
] = ProviderConfigManager.get_provider_responses_api_config(
model=None,
provider=custom_llm_provider,
responses_api_provider_config: Optional[BaseResponsesAPIConfig] = (
ProviderConfigManager.get_provider_responses_api_config(
model=None,
provider=custom_llm_provider,
)
)
if responses_api_provider_config is None:
@ -1296,11 +1302,11 @@ def get_responses(
raise ValueError("custom_llm_provider is required but passed as None")
# get provider config
responses_api_provider_config: Optional[
BaseResponsesAPIConfig
] = ProviderConfigManager.get_provider_responses_api_config(
model=None,
provider=custom_llm_provider,
responses_api_provider_config: Optional[BaseResponsesAPIConfig] = (
ProviderConfigManager.get_provider_responses_api_config(
model=None,
provider=custom_llm_provider,
)
)
if responses_api_provider_config is None:
@ -1454,11 +1460,11 @@ def list_input_items(
if custom_llm_provider is None:
raise ValueError("custom_llm_provider is required but passed as None")
responses_api_provider_config: Optional[
BaseResponsesAPIConfig
] = ProviderConfigManager.get_provider_responses_api_config(
model=None,
provider=custom_llm_provider,
responses_api_provider_config: Optional[BaseResponsesAPIConfig] = (
ProviderConfigManager.get_provider_responses_api_config(
model=None,
provider=custom_llm_provider,
)
)
if responses_api_provider_config is None:
@ -1613,11 +1619,11 @@ def cancel_responses(
raise ValueError("custom_llm_provider is required but passed as None")
# get provider config
responses_api_provider_config: Optional[
BaseResponsesAPIConfig
] = ProviderConfigManager.get_provider_responses_api_config(
model=None,
provider=custom_llm_provider,
responses_api_provider_config: Optional[BaseResponsesAPIConfig] = (
ProviderConfigManager.get_provider_responses_api_config(
model=None,
provider=custom_llm_provider,
)
)
if responses_api_provider_config is None:
@ -1801,11 +1807,11 @@ def compact_responses(
raise ValueError("custom_llm_provider is required but passed as None")
# get provider config
responses_api_provider_config: Optional[
BaseResponsesAPIConfig
] = ProviderConfigManager.get_provider_responses_api_config(
model=model,
provider=custom_llm_provider,
responses_api_provider_config: Optional[BaseResponsesAPIConfig] = (
ProviderConfigManager.get_provider_responses_api_config(
model=model,
provider=custom_llm_provider,
)
)
if responses_api_provider_config is None: