mirror of https://github.com/aliasrobotics/cai.git
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:
parent
7cda5445a4
commit
0877dfda81
|
|
@ -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.
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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')}")
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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(),
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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":
|
||||||
|
|
|
||||||
110
src/cai/util.py
110
src/cai/util.py
|
|
@ -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
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue