diff --git a/examples/agent_patterns/deterministic.py b/examples/agent_patterns/deterministic.py index 0c163afe..76bffb3e 100644 --- a/examples/agent_patterns/deterministic.py +++ b/examples/agent_patterns/deterministic.py @@ -2,7 +2,7 @@ import asyncio 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. diff --git a/examples/cai/simple_one_tool_test.py b/examples/cai/simple_one_tool_test.py index 056ad636..d7578f9c 100644 --- a/examples/cai/simple_one_tool_test.py +++ b/examples/cai/simple_one_tool_test.py @@ -18,12 +18,13 @@ from cai.util import fix_litellm_transcription_annotations, color # Load environment variables load_dotenv() -# Initialize OpenAI client -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) +# NOTE: This is needed when using LiteLLM Proxy Server +# +# 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) async def main(): # Apply litellm patch to fix the __annotations__ error diff --git a/examples/cai/simple_one_tool_test_streamed.py b/examples/cai/simple_one_tool_test_streamed.py index 4a8475f8..9708b3d4 100644 --- a/examples/cai/simple_one_tool_test_streamed.py +++ b/examples/cai/simple_one_tool_test_streamed.py @@ -14,17 +14,18 @@ from openai.types.responses import ResponseTextDeltaEvent from cai.sdk.agents import Runner, set_default_openai_client from cai.agents import get_agent_by_name from cai.util import fix_litellm_transcription_annotations, color - +from cai.sdk.agents import Agent, OpenAIChatCompletionsModel # Load environment variables load_dotenv() -# Initialize OpenAI client -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) +# NOTE: This is needed when using LiteLLM Proxy Server +# +# 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) async def main(): # Apply litellm patch to fix the __annotations__ error @@ -34,7 +35,7 @@ async def main(): # Get the one_tool agent agent = get_agent_by_name("one_tool_agent") - + print("Testing one_tool agent with a simple hello message (streaming mode)...") print(f"Using model: {os.getenv('CAI_MODEL', 'default')}") diff --git a/src/cai/agents/one_tool.py b/src/cai/agents/one_tool.py index 6d08c576..6dd09aad 100644 --- a/src/cai/agents/one_tool.py +++ b/src/cai/agents/one_tool.py @@ -9,11 +9,13 @@ from openai import AsyncOpenAI # Get model from environment or use default model_name = os.getenv('CAI_MODEL', "qwen2.5:14b") -# Create OpenAI client for the agent -openai_client = AsyncOpenAI( - base_url = os.getenv('LITELLM_BASE_URL', 'http://localhost:4000'), - api_key=os.getenv('LITELLM_API_KEY', 'key') -) +# NOTE: This is needed when using LiteLLM Proxy Server +# +# # Create OpenAI client for the agent +# 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 # is_qwen = "qwen" in model_name.lower() @@ -58,7 +60,7 @@ one_tool_agent = Agent( ], model=OpenAIChatCompletionsModel( model=model_name, - openai_client=openai_client, + openai_client=AsyncOpenAI(), ) ) diff --git a/src/cai/cli.py b/src/cai/cli.py index 4efc90f8..e19e4b89 100644 --- a/src/cai/cli.py +++ b/src/cai/cli.py @@ -120,11 +120,14 @@ from cai.agents import get_agent_by_name # Load environment variables from .env file load_dotenv() -external_client = AsyncOpenAI( - base_url = os.getenv('LITELLM_BASE_URL', 'http://localhost:4000'), - api_key=os.getenv('LITELLM_API_KEY', 'key')) +# NOTE: This is needed when using LiteLLM Proxy Server +# +# 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) # llm_model=os.getenv('LLM_MODEL', 'gpt-4o-mini') @@ -140,8 +143,8 @@ agent = Agent( instructions=instructions, model=OpenAIChatCompletionsModel( model=llm_model, - # openai_client=AsyncOpenAI() # original OpenAI servers - openai_client = external_client + openai_client=AsyncOpenAI() # original OpenAI servers + # openai_client = external_client # LiteLLM Proxy Server ) ) diff --git a/src/cai/repl/ui/banner.py b/src/cai/repl/ui/banner.py index 453b316c..006f83cd 100644 --- a/src/cai/repl/ui/banner.py +++ b/src/cai/repl/ui/banner.py @@ -5,6 +5,7 @@ Module for displaying the CAI banner and welcome message. import os import glob import logging +import sys from configparser import ConfigParser # 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.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(): - """Get the CAI version from setup.cfg.""" + """Get the CAI version from pyproject.toml.""" version = "unknown" try: - config = ConfigParser() - config.read('setup.cfg') - version = config.get('metadata', 'version') - except Exception: # pylint: disable=broad-except - logging.warning("Could not read version from setup.cfg") + # Determine which TOML parser to use + if sys.version_info >= (3, 11): + toml_parser = tomllib + else: + 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 diff --git a/src/cai/sdk/agents/models/openai_chatcompletions.py b/src/cai/sdk/agents/models/openai_chatcompletions.py index c0c1ba18..de5e80ad 100644 --- a/src/cai/sdk/agents/models/openai_chatcompletions.py +++ b/src/cai/sdk/agents/models/openai_chatcompletions.py @@ -9,7 +9,8 @@ import litellm from collections.abc import AsyncIterator, Iterable from dataclasses import dataclass, field 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.types import ChatModel @@ -87,6 +88,9 @@ if TYPE_CHECKING: from ..model_settings import ModelSettings +# Suppress debug info from litellm +litellm.suppress_debug_info = True + _USER_AGENT = f"Agents/Python {__version__}" _HEADERS = {"User-Agent": _USER_AGENT} @@ -238,9 +242,11 @@ class OpenAIChatCompletionsModel(Model): continue # Get the delta content - delta = choices[0].get('delta', None) - if not delta and hasattr(choices[0], 'delta'): + delta = None + if hasattr(choices[0], 'delta'): delta = choices[0].delta + elif isinstance(choices[0], dict) and 'delta' in choices[0]: + delta = choices[0]['delta'] if not delta: continue @@ -579,6 +585,10 @@ class OpenAIChatCompletionsModel(Model): tracing: ModelTracing, stream: bool = False, ) -> 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) if system_instructions: @@ -602,9 +612,6 @@ class OpenAIChatCompletionsModel(Model): for handoff in handoffs: converted_tools.append(ToolConverter.convert_handoff_tool(handoff)) - # if self.is_ollama: - # converted_tools = [] - if _debug.DONT_LOG_MODEL_DATA: logger.debug("Calling LLM") else: @@ -619,7 +626,7 @@ class OpenAIChatCompletionsModel(Model): # 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 - + # Prepare kwargs for the API call kwargs = { "model": self.model, @@ -639,87 +646,239 @@ class OpenAIChatCompletionsModel(Model): "extra_headers": _HEADERS, } - # Previously articulated as - # ret = await self._get_client().chat.completions.create(**kwargs) - # but now we're using litellm + # Model adjustments + if any(x in self.model for x in ["claude"]): + litellm.drop_params = True - if self.is_ollama: - # 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"] - } + # Error encountered: Error code: 400 - {'error': {'code': 'invalid_request_error', + # 'message': "'tool_choice' is only allowed when 'tools' are specified", + # 'type': 'invalid_request_error', 'param': None}} + # + # if has no tools, remove tool_choice + if not converted_tools: + kwargs.pop("tool_choice", None) + + # BadRequestError encountered: litellm.BadRequestError: AnthropicException - + # 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 - 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}" + # Filter out NotGiven values to avoid JSON serialization issues + filtered_kwargs = {} + for key, value in kwargs.items(): + if value is not NOT_GIVEN: + filtered_kwargs[key] = value + kwargs = filtered_kwargs - # 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 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 + try: + if self.is_ollama: + return await self._fetch_response_litellm_ollama(kwargs, model_settings, tool_choice, stream, parallel_tool_calls) else: - # Non-streaming mode - ollama_api_base = get_ollama_api_base().rstrip('/v1') # Remove /v1 if present - ret = litellm.completion( - **ollama_kwargs, - api_base=ollama_api_base, - custom_llm_provider="ollama" + return await self._fetch_response_litellm_openai(kwargs, model_settings, tool_choice, stream, parallel_tool_calls) + + # ret = await self._get_client().chat.completions.create(**kwargs) + + # if isinstance(ret, ChatCompletion): + # 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 - else: - # Standard LiteLLM handling + if retry_info and 'retryDelay' in retry_info: + retry_delay = int(retry_info['retryDelay'].rstrip('s')) + 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) + 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 def _get_client(self) -> AsyncOpenAI: @@ -734,7 +893,7 @@ class _Converter: cls, tool_choice: Literal["auto", "required", "none"] | str | None ) -> ChatCompletionToolChoiceOptionParam | NotGiven: if tool_choice is None: - return None + return "auto" elif tool_choice == "auto": return "auto" elif tool_choice == "required": diff --git a/src/cai/util.py b/src/cai/util.py index d986b6dc..a24217e4 100644 --- a/src/cai/util.py +++ b/src/cai/util.py @@ -153,3 +153,113 @@ def fix_litellm_transcription_annotations(): except (ImportError, AttributeError): # If the import fails or the attribute doesn't exist, the patch couldn't be applied 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