Address issues with litellm and agents SDK

Squeezed in some additional features such as:
- updated get_version()
- various examples for streaming and non-streaming testing

Signed-off-by: Víctor Mayoral Vilches <v.mayoralv@gmail.com>
This commit is contained in:
Víctor Mayoral Vilches 2025-04-02 14:03:00 +02:00
parent 7cda5445a4
commit 0877dfda81
8 changed files with 421 additions and 116 deletions

View File

@ -2,7 +2,7 @@ import asyncio
from pydantic import BaseModel from pydantic import BaseModel
from agents import Agent, Runner, trace from cai.sdk.agents import Agent, Runner, trace
""" """
This example demonstrates a deterministic flow, where each step is performed by an agent. This example demonstrates a deterministic flow, where each step is performed by an agent.

View File

@ -18,12 +18,13 @@ from cai.util import fix_litellm_transcription_annotations, color
# Load environment variables # Load environment variables
load_dotenv() load_dotenv()
# Initialize OpenAI client # NOTE: This is needed when using LiteLLM Proxy Server
external_client = AsyncOpenAI( #
base_url=os.getenv('LITELLM_BASE_URL', 'http://localhost:4000'), # external_client = AsyncOpenAI(
api_key=os.getenv('LITELLM_API_KEY', 'key') # base_url=os.getenv('LITELLM_BASE_URL', 'http://localhost:4000'),
) # api_key=os.getenv('LITELLM_API_KEY', 'key')
set_default_openai_client(external_client) # )
# set_default_openai_client(external_client)
async def main(): async def main():
# Apply litellm patch to fix the __annotations__ error # Apply litellm patch to fix the __annotations__ error

View File

@ -14,17 +14,18 @@ from openai.types.responses import ResponseTextDeltaEvent
from cai.sdk.agents import Runner, set_default_openai_client from cai.sdk.agents import Runner, set_default_openai_client
from cai.agents import get_agent_by_name from cai.agents import get_agent_by_name
from cai.util import fix_litellm_transcription_annotations, color from cai.util import fix_litellm_transcription_annotations, color
from cai.sdk.agents import Agent, OpenAIChatCompletionsModel
# Load environment variables # Load environment variables
load_dotenv() load_dotenv()
# Initialize OpenAI client # NOTE: This is needed when using LiteLLM Proxy Server
external_client = AsyncOpenAI( #
base_url=os.getenv('LITELLM_BASE_URL', 'http://localhost:4000'), # external_client = AsyncOpenAI(
api_key=os.getenv('LITELLM_API_KEY', 'key') # base_url=os.getenv('LITELLM_BASE_URL', 'http://localhost:4000'),
) # api_key=os.getenv('LITELLM_API_KEY', 'key')
set_default_openai_client(external_client) # )
# set_default_openai_client(external_client)
async def main(): async def main():
# Apply litellm patch to fix the __annotations__ error # Apply litellm patch to fix the __annotations__ error
@ -34,7 +35,7 @@ async def main():
# Get the one_tool agent # Get the one_tool agent
agent = get_agent_by_name("one_tool_agent") agent = get_agent_by_name("one_tool_agent")
print("Testing one_tool agent with a simple hello message (streaming mode)...") print("Testing one_tool agent with a simple hello message (streaming mode)...")
print(f"Using model: {os.getenv('CAI_MODEL', 'default')}") print(f"Using model: {os.getenv('CAI_MODEL', 'default')}")

View File

@ -9,11 +9,13 @@ from openai import AsyncOpenAI
# Get model from environment or use default # Get model from environment or use default
model_name = os.getenv('CAI_MODEL', "qwen2.5:14b") model_name = os.getenv('CAI_MODEL', "qwen2.5:14b")
# Create OpenAI client for the agent # NOTE: This is needed when using LiteLLM Proxy Server
openai_client = AsyncOpenAI( #
base_url = os.getenv('LITELLM_BASE_URL', 'http://localhost:4000'), # # Create OpenAI client for the agent
api_key=os.getenv('LITELLM_API_KEY', 'key') # openai_client = AsyncOpenAI(
) # base_url = os.getenv('LITELLM_BASE_URL', 'http://localhost:4000'),
# api_key=os.getenv('LITELLM_API_KEY', 'key')
# )
# # Check if we're using a Qwen model # # Check if we're using a Qwen model
# is_qwen = "qwen" in model_name.lower() # is_qwen = "qwen" in model_name.lower()
@ -58,7 +60,7 @@ one_tool_agent = Agent(
], ],
model=OpenAIChatCompletionsModel( model=OpenAIChatCompletionsModel(
model=model_name, model=model_name,
openai_client=openai_client, openai_client=AsyncOpenAI(),
) )
) )

View File

@ -120,11 +120,14 @@ from cai.agents import get_agent_by_name
# Load environment variables from .env file # Load environment variables from .env file
load_dotenv() load_dotenv()
external_client = AsyncOpenAI( # NOTE: This is needed when using LiteLLM Proxy Server
base_url = os.getenv('LITELLM_BASE_URL', 'http://localhost:4000'), #
api_key=os.getenv('LITELLM_API_KEY', 'key')) # external_client = AsyncOpenAI(
# base_url = os.getenv('LITELLM_BASE_URL', 'http://localhost:4000'),
# api_key=os.getenv('LITELLM_API_KEY', 'key'))
#
# set_default_openai_client(external_client)
set_default_openai_client(external_client)
set_tracing_disabled(True) set_tracing_disabled(True)
# llm_model=os.getenv('LLM_MODEL', 'gpt-4o-mini') # llm_model=os.getenv('LLM_MODEL', 'gpt-4o-mini')
@ -140,8 +143,8 @@ agent = Agent(
instructions=instructions, instructions=instructions,
model=OpenAIChatCompletionsModel( model=OpenAIChatCompletionsModel(
model=llm_model, model=llm_model,
# openai_client=AsyncOpenAI() # original OpenAI servers openai_client=AsyncOpenAI() # original OpenAI servers
openai_client = external_client # openai_client = external_client # LiteLLM Proxy Server
) )
) )

View File

@ -5,6 +5,7 @@ Module for displaying the CAI banner and welcome message.
import os import os
import glob import glob
import logging import logging
import sys
from configparser import ConfigParser from configparser import ConfigParser
# Third-party imports # Third-party imports
@ -13,16 +14,44 @@ from rich.console import Console # pylint: disable=import-error
from rich.panel import Panel # pylint: disable=import-error from rich.panel import Panel # pylint: disable=import-error
from rich.table import Table # pylint: disable=import-error from rich.table import Table # pylint: disable=import-error
# For reading TOML files
if sys.version_info >= (3, 11):
import tomllib
else:
try:
import tomli as tomllib
except ImportError:
# If tomli is not available, we'll handle it in the get_version function
pass
def get_version(): def get_version():
"""Get the CAI version from setup.cfg.""" """Get the CAI version from pyproject.toml."""
version = "unknown" version = "unknown"
try: try:
config = ConfigParser() # Determine which TOML parser to use
config.read('setup.cfg') if sys.version_info >= (3, 11):
version = config.get('metadata', 'version') toml_parser = tomllib
except Exception: # pylint: disable=broad-except else:
logging.warning("Could not read version from setup.cfg") try:
import tomli as toml_parser
except ImportError:
logging.warning("Could not import tomli. Falling back to manual parsing.")
# Simple manual parsing for version only
with open('pyproject.toml', 'r', encoding='utf-8') as f:
for line in f:
if line.strip().startswith('version = '):
# Extract version from line like 'version = "0.4.0"'
version = line.split('=')[1].strip().strip('"\'')
return version
return version
# Use proper TOML parser if available
with open('pyproject.toml', 'rb') as f:
config = toml_parser.load(f)
version = config.get('project', {}).get('version', 'unknown')
except Exception as e: # pylint: disable=broad-except
logging.warning("Could not read version from pyproject.toml: %s", e)
return version return version

View File

@ -9,7 +9,8 @@ import litellm
from collections.abc import AsyncIterator, Iterable from collections.abc import AsyncIterator, Iterable
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any, Literal, cast, overload from typing import TYPE_CHECKING, Any, Literal, cast, overload
from cai.util import get_ollama_api_base from cai.util import get_ollama_api_base, fix_message_list
from wasabi import color
from openai import NOT_GIVEN, AsyncOpenAI, AsyncStream, NotGiven from openai import NOT_GIVEN, AsyncOpenAI, AsyncStream, NotGiven
from openai.types import ChatModel from openai.types import ChatModel
@ -87,6 +88,9 @@ if TYPE_CHECKING:
from ..model_settings import ModelSettings from ..model_settings import ModelSettings
# Suppress debug info from litellm
litellm.suppress_debug_info = True
_USER_AGENT = f"Agents/Python {__version__}" _USER_AGENT = f"Agents/Python {__version__}"
_HEADERS = {"User-Agent": _USER_AGENT} _HEADERS = {"User-Agent": _USER_AGENT}
@ -238,9 +242,11 @@ class OpenAIChatCompletionsModel(Model):
continue continue
# Get the delta content # Get the delta content
delta = choices[0].get('delta', None) delta = None
if not delta and hasattr(choices[0], 'delta'): if hasattr(choices[0], 'delta'):
delta = choices[0].delta delta = choices[0].delta
elif isinstance(choices[0], dict) and 'delta' in choices[0]:
delta = choices[0]['delta']
if not delta: if not delta:
continue continue
@ -579,6 +585,10 @@ class OpenAIChatCompletionsModel(Model):
tracing: ModelTracing, tracing: ModelTracing,
stream: bool = False, stream: bool = False,
) -> ChatCompletion | tuple[Response, AsyncStream[ChatCompletionChunk]]: ) -> ChatCompletion | tuple[Response, AsyncStream[ChatCompletionChunk]]:
# start by re-fetching self.is_ollama
self.is_ollama = os.getenv('OLLAMA') is not None and os.getenv('OLLAMA').lower() == 'true'
converted_messages = _Converter.items_to_messages(input) converted_messages = _Converter.items_to_messages(input)
if system_instructions: if system_instructions:
@ -602,9 +612,6 @@ class OpenAIChatCompletionsModel(Model):
for handoff in handoffs: for handoff in handoffs:
converted_tools.append(ToolConverter.convert_handoff_tool(handoff)) converted_tools.append(ToolConverter.convert_handoff_tool(handoff))
# if self.is_ollama:
# converted_tools = []
if _debug.DONT_LOG_MODEL_DATA: if _debug.DONT_LOG_MODEL_DATA:
logger.debug("Calling LLM") logger.debug("Calling LLM")
else: else:
@ -619,7 +626,7 @@ class OpenAIChatCompletionsModel(Model):
# Match the behavior of Responses where store is True when not given # Match the behavior of Responses where store is True when not given
store = model_settings.store if model_settings.store is not None else True store = model_settings.store if model_settings.store is not None else True
# Prepare kwargs for the API call # Prepare kwargs for the API call
kwargs = { kwargs = {
"model": self.model, "model": self.model,
@ -639,87 +646,239 @@ class OpenAIChatCompletionsModel(Model):
"extra_headers": _HEADERS, "extra_headers": _HEADERS,
} }
# Previously articulated as # Model adjustments
# ret = await self._get_client().chat.completions.create(**kwargs) if any(x in self.model for x in ["claude"]):
# but now we're using litellm litellm.drop_params = True
if self.is_ollama: # Error encountered: Error code: 400 - {'error': {'code': 'invalid_request_error',
# Filter out parameters not supported by Ollama # 'message': "'tool_choice' is only allowed when 'tools' are specified",
ollama_supported_params = { # 'type': 'invalid_request_error', 'param': None}}
"model": kwargs["model"], #
"messages": kwargs["messages"], # if has no tools, remove tool_choice
"temperature": kwargs["temperature"] if kwargs["temperature"] is not NOT_GIVEN else None, if not converted_tools:
"top_p": kwargs["top_p"] if kwargs["top_p"] is not NOT_GIVEN else None, kwargs.pop("tool_choice", None)
"max_tokens": kwargs["max_tokens"] if kwargs["max_tokens"] is not NOT_GIVEN else None,
"stream": kwargs["stream"], # BadRequestError encountered: litellm.BadRequestError: AnthropicException -
"extra_headers": kwargs["extra_headers"] # b'{"type":"error","error":
} # {"type":"invalid_request_error","message":"store: Extra inputs are not permitted"}}'
#
kwargs.pop("store", None)
# Modify the messages to remove system message for Ollama # Filter out NotGiven values to avoid JSON serialization issues
if ollama_supported_params["messages"] and ollama_supported_params["messages"][0].get("role") == "system": filtered_kwargs = {}
# Extract the system message for key, value in kwargs.items():
system_content = ollama_supported_params["messages"][0].get("content", "") if value is not NOT_GIVEN:
# Remove it from the messages filtered_kwargs[key] = value
ollama_supported_params["messages"] = ollama_supported_params["messages"][1:] kwargs = filtered_kwargs
# If there are user messages, prepend system to first user
if ollama_supported_params["messages"] and ollama_supported_params["messages"][0].get("role") == "user":
# Prepend the system instruction to the first user message, with a separator
user_content = ollama_supported_params["messages"][0].get("content", "")
if isinstance(user_content, str):
ollama_supported_params["messages"][0]["content"] = f"System: {system_content}\n\nUser: {user_content}"
# Remove None values try:
ollama_kwargs = {k: v for k, v in ollama_supported_params.items() if v is not None} if self.is_ollama:
return await self._fetch_response_litellm_ollama(kwargs, model_settings, tool_choice, stream, parallel_tool_calls)
if stream:
# For streaming with Ollama, we need to create a Response object first
response = Response(
id=FAKE_RESPONSES_ID,
created_at=time.time(),
model=self.model,
object="response",
output=[],
tool_choice="auto" if tool_choice is None else cast(Literal["auto", "required", "none"], tool_choice)
if tool_choice != NOT_GIVEN
else "auto",
top_p=model_settings.top_p,
temperature=model_settings.temperature,
tools=[],
parallel_tool_calls=parallel_tool_calls or False,
usage={
"completion_tokens": 0,
"prompt_tokens": 0,
"total_tokens": 0,
"input_tokens": 0,
"input_tokens_details": {
"cached_tokens": 0
},
"output_tokens": 0,
"output_tokens_details": {
"reasoning_tokens": 0
}
},
)
# Get the streaming object
ollama_api_base = get_ollama_api_base().rstrip('/v1') # Remove /v1 if present
stream_obj = await litellm.acompletion(
**ollama_kwargs,
api_base=ollama_api_base,
custom_llm_provider="ollama"
)
return response, stream_obj
else: else:
# Non-streaming mode return await self._fetch_response_litellm_openai(kwargs, model_settings, tool_choice, stream, parallel_tool_calls)
ollama_api_base = get_ollama_api_base().rstrip('/v1') # Remove /v1 if present
ret = litellm.completion( # ret = await self._get_client().chat.completions.create(**kwargs)
**ollama_kwargs,
api_base=ollama_api_base, # if isinstance(ret, ChatCompletion):
custom_llm_provider="ollama" # return ret
# response = Response(
# id=FAKE_RESPONSES_ID,
# created_at=time.time(),
# model=self.model,
# object="response",
# output=[],
# tool_choice=cast(Literal["auto", "required", "none"], tool_choice)
# if tool_choice != NOT_GIVEN
# else "auto",
# top_p=model_settings.top_p,
# temperature=model_settings.temperature,
# tools=[],
# parallel_tool_calls=parallel_tool_calls or False,
# )
# return response, ret
except litellm.exceptions.BadRequestError as e:
print(color("BadRequestError encountered: " + str(e), fg="yellow"))
if "LLM Provider NOT provided" in str(e):
# Create a copy of params to avoid overwriting the original
# ones
try:
return await self._fetch_response_litellm_ollama(kwargs, model_settings, tool_choice, stream, parallel_tool_calls)
except litellm.exceptions.BadRequestError as e: # pylint: disable=W0621,C0301 # noqa: E501
#
# CTRL-C handler for ollama models
#
if "invalid message content type" in str(e):
kwargs["messages"] = fix_message_list(
kwargs["messages"])
return await self._fetch_response_litellm_ollama(kwargs, model_settings, tool_choice, stream, parallel_tool_calls)
else:
raise e
elif ("An assistant message with 'tool_calls'" in str(e) or
"`tool_use` blocks must be followed by a user message with `tool_result`" in str(e)): # noqa: E501 # pylint: disable=C0301
print(f"Error: {str(e)}")
# NOTE: EDGE CASE: Report Agent CTRL C error
#
# This fix CTRL-C error when message list is incomplete
# When a tool is not finished but the LLM generates a tool call
kwargs["messages"] = fix_message_list(kwargs["messages"])
return await self._fetch_response_litellm_openai(kwargs, model_settings, tool_choice, stream, parallel_tool_calls)
# this captures an error related to the fact
# that the messages list contains an empty
# content position
elif "expected a string, got null" in str(e):
print(f"Error: {str(e)}")
# Fix for null content in messages
kwargs["messages"] = [
msg if msg.get("content") is not None else
{**msg, "content": ""} for msg in kwargs["messages"]
]
return await self._fetch_response_litellm_openai(kwargs, model_settings, tool_choice, stream, parallel_tool_calls)
# Handle Anthropic error for empty text content blocks
elif ("text content blocks must be non-empty" in str(e) or
"cache_control cannot be set for empty text blocks" in str(e)): # noqa
print(f"Error: {str(e)}")
# Fix for empty content in messages for Anthropic models
kwargs["messages"] = [
msg if msg.get("content") not in [None, ""] else
{
**msg,
"content": "Empty content block"
} for msg in kwargs["messages"]
]
return await self._fetch_response_litellm_openai(kwargs, model_settings, tool_choice, stream, parallel_tool_calls)
else:
raise e
except litellm.exceptions.RateLimitError as e:
print("Rate Limit Error:" + str(e))
# Try to extract retry delay from error response or use default
retry_delay = 60 # Default delay in seconds
try:
# Extract the JSON part from the error message
json_str = str(e.message).split('VertexAIException - ')[-1]
error_details = json.loads(json_str)
retry_info = next(
(detail for detail in error_details.get('error', {}).get('details', [])
if detail.get('@type') == 'type.googleapis.com/google.rpc.RetryInfo'),
None
) )
return ret if retry_info and 'retryDelay' in retry_info:
else: retry_delay = int(retry_info['retryDelay'].rstrip('s'))
# Standard LiteLLM handling except Exception as parse_error:
print(f"Could not parse retry delay, using default: {parse_error}")
print(f"Waiting {retry_delay} seconds before retrying...")
time.sleep(retry_delay)
# fall back to ollama if openai API fails
except Exception as e: # pylint: disable=W0718
print(color("Error encountered: " + str(e), fg="yellow"))
try:
return await self._fetch_response_litellm_ollama(kwargs, model_settings, tool_choice, stream, parallel_tool_calls)
except Exception as execp: # pylint: disable=W0718
print("Error: " + str(execp))
return None
async def _fetch_response_litellm_openai(
self,
kwargs: dict,
model_settings: ModelSettings,
tool_choice: ChatCompletionToolChoiceOptionParam | NotGiven,
stream: bool,
parallel_tool_calls: bool
) -> ChatCompletion | tuple[Response, AsyncStream[ChatCompletionChunk]]:
"""Handle standard LiteLLM API calls for OpenAI and compatible models."""
if stream:
# Standard LiteLLM handling for streaming
ret = litellm.completion(**kwargs) ret = litellm.completion(**kwargs)
stream_obj = await litellm.acompletion(**kwargs)
response = Response(
id=FAKE_RESPONSES_ID,
created_at=time.time(),
model=self.model,
object="response",
output=[],
tool_choice="auto" if tool_choice is None or tool_choice == NOT_GIVEN else cast(Literal["auto", "required", "none"], tool_choice),
top_p=model_settings.top_p,
temperature=model_settings.temperature,
tools=[],
parallel_tool_calls=parallel_tool_calls or False,
)
return response, stream_obj
else:
# Standard OpenAI handling for non-streaming
ret = litellm.completion(**kwargs)
return ret
async def _fetch_response_litellm_ollama(
self,
kwargs: dict,
model_settings: ModelSettings,
tool_choice: ChatCompletionToolChoiceOptionParam | NotGiven,
stream: bool,
parallel_tool_calls: bool
) -> ChatCompletion | tuple[Response, AsyncStream[ChatCompletionChunk]]:
# Filter out parameters not supported by Ollama
ollama_supported_params = {
"model": kwargs["model"],
"messages": kwargs["messages"],
"temperature": kwargs["temperature"] if kwargs["temperature"] is not NOT_GIVEN else None,
"top_p": kwargs["top_p"] if kwargs["top_p"] is not NOT_GIVEN else None,
"max_tokens": kwargs["max_tokens"] if kwargs["max_tokens"] is not NOT_GIVEN else None,
"stream": kwargs["stream"],
"extra_headers": kwargs["extra_headers"]
}
# Modify the messages to remove system message for Ollama
if ollama_supported_params["messages"] and ollama_supported_params["messages"][0].get("role") == "system":
# Extract the system message
system_content = ollama_supported_params["messages"][0].get("content", "")
# Remove it from the messages
ollama_supported_params["messages"] = ollama_supported_params["messages"][1:]
# If there are user messages, prepend system to first user
if ollama_supported_params["messages"] and ollama_supported_params["messages"][0].get("role") == "user":
# Prepend the system instruction to the first user message, with a separator
user_content = ollama_supported_params["messages"][0].get("content", "")
if isinstance(user_content, str):
ollama_supported_params["messages"][0]["content"] = f"System: {system_content}\n\nUser: {user_content}"
# Remove None values
ollama_kwargs = {k: v for k, v in ollama_supported_params.items() if v is not None}
if stream:
# For streaming with Ollama, we need to create a Response object first
response = Response(
id=FAKE_RESPONSES_ID,
created_at=time.time(),
model=self.model,
object="response",
output=[],
tool_choice="auto" if tool_choice is None or tool_choice == NOT_GIVEN else cast(Literal["auto", "required", "none"], tool_choice),
top_p=model_settings.top_p,
temperature=model_settings.temperature,
tools=[],
parallel_tool_calls=parallel_tool_calls or False,
)
# Get the streaming object
stream_obj = await litellm.acompletion(
**ollama_kwargs,
api_base=get_ollama_api_base().rstrip('/v1'),
custom_llm_provider="ollama"
)
return response, stream_obj
else:
# Non-streaming mode
ret = litellm.completion(
**ollama_kwargs,
api_base=get_ollama_api_base().rstrip('/v1'),
custom_llm_provider="ollama"
)
return ret return ret
def _get_client(self) -> AsyncOpenAI: def _get_client(self) -> AsyncOpenAI:
@ -734,7 +893,7 @@ class _Converter:
cls, tool_choice: Literal["auto", "required", "none"] | str | None cls, tool_choice: Literal["auto", "required", "none"] | str | None
) -> ChatCompletionToolChoiceOptionParam | NotGiven: ) -> ChatCompletionToolChoiceOptionParam | NotGiven:
if tool_choice is None: if tool_choice is None:
return None return "auto"
elif tool_choice == "auto": elif tool_choice == "auto":
return "auto" return "auto"
elif tool_choice == "required": elif tool_choice == "required":

View File

@ -153,3 +153,113 @@ def fix_litellm_transcription_annotations():
except (ImportError, AttributeError): except (ImportError, AttributeError):
# If the import fails or the attribute doesn't exist, the patch couldn't be applied # If the import fails or the attribute doesn't exist, the patch couldn't be applied
return False return False
def fix_message_list(messages): # pylint: disable=R0914,R0915,R0912
"""
Sanitizes the message list passed as a parameter to align with the
OpenAI API message format.
Adjusts the message list to comply with the following rules:
1. A tool call id appears no more than twice.
2. Each tool call id appears as a pair, and both messages
must have content.
3. If a tool call id appears alone (without a pair), it is removed.
4. There cannot be empty messages.
Args:
messages (List[dict]): List of message dictionaries containing
role, content, and optionally tool_calls or
tool_call_id fields.
Returns:
List[dict]: Sanitized list of messages with invalid tool calls
and empty messages removed.
"""
# Step 1: Filter and discard empty messages (considered empty if 'content'
# is None or only whitespace)
cleaned_messages = []
for msg in messages:
content = msg.get("content")
if content is not None and content.strip():
cleaned_messages.append(msg)
messages = cleaned_messages
# Step 2: Collect tool call id occurrences.
# In assistant messages, iterate through 'tool_calls' list.
# In 'tool' type messages, use the 'tool_call_id' key.
tool_calls_occurrences = {}
for i, msg in enumerate(messages):
if msg.get("role") == "assistant" and isinstance(
msg.get("tool_calls"), list):
for j, tool_call in enumerate(msg["tool_calls"]):
tc_id = tool_call.get("id")
if tc_id:
tool_calls_occurrences.setdefault(
tc_id, []).append((i, "assistant", j))
elif msg.get("role") == "tool" and msg.get("tool_call_id"):
tc_id = msg["tool_call_id"]
tool_calls_occurrences.setdefault(
tc_id, []).append(
(i, "tool", None))
# Step 3: Mark invalid or extra occurrences for removal
removal_messages = set() # Indices of messages (tool type) to remove
# Maps message index (assistant) to set of indices (in tool_calls) to
# remove
removal_assistant_entries = {}
for tc_id, occurrences in tool_calls_occurrences.items():
# Only 2 occurrences allowed. Mark extras for removal.
valid_occurrences = occurrences[:2]
extra_occurrences = occurrences[2:]
for occ in extra_occurrences:
msg_idx, typ, j = occ
if typ == "assistant":
removal_assistant_entries.setdefault(msg_idx, set()).add(j)
elif typ == "tool":
removal_messages.add(msg_idx)
# If valid occurrences aren't exactly 2 (i.e., a lonely tool call),
# mark for removal
if len(valid_occurrences) != 2:
for occ in valid_occurrences:
msg_idx, typ, j = occ
if typ == "assistant":
removal_assistant_entries.setdefault(
msg_idx, set()).add(j)
elif typ == "tool":
removal_messages.add(msg_idx)
else:
# If exactly 2 occurrences, ensure both have content
remove_pair = False
for occ in valid_occurrences:
msg_idx, typ, _ = occ
msg_content = messages[msg_idx].get("content")
if msg_content is None or not msg_content.strip():
remove_pair = True
break
if remove_pair:
for occ in valid_occurrences:
msg_idx, typ, j = occ
if typ == "assistant":
removal_assistant_entries.setdefault(
msg_idx, set()).add(j)
elif typ == "tool":
removal_messages.add(msg_idx)
# Step 4: Build new message list applying removals
new_messages = []
for i, msg in enumerate(messages):
# Skip if message (tool type) is marked for removal
if i in removal_messages:
continue
# For assistant messages, remove marked tool_calls
if msg.get("role") == "assistant" and "tool_calls" in msg:
new_tool_calls = []
for j, tc in enumerate(msg["tool_calls"]):
if j not in removal_assistant_entries.get(i, set()):
new_tool_calls.append(tc)
msg["tool_calls"] = new_tool_calls
# If after modification message has no content and no tool_calls,
# discard it
msg_content = msg.get("content")
if ((msg_content is None or not msg_content.strip()) and
not msg.get("tool_calls")):
continue
new_messages.append(msg)
return new_messages