mirror of https://github.com/aliasrobotics/cai.git
Streaming now working for OLLAMA models, but not the rest. Figure it out
Signed-off-by: Víctor Mayoral Vilches <v.mayoralv@gmail.com>
This commit is contained in:
parent
cff20480b7
commit
aef3f5e73c
|
|
@ -0,0 +1,61 @@
|
|||
"""
|
||||
A simple example to test the one_tool agent with streaming.
|
||||
|
||||
This script demonstrates how to initialize and run the one_tool agent
|
||||
using Runner.run_streamed() with a simple hello message to verify everything
|
||||
is working correctly, with streaming output.
|
||||
"""
|
||||
|
||||
import os
|
||||
import asyncio
|
||||
from dotenv import load_dotenv
|
||||
from openai import AsyncOpenAI
|
||||
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
|
||||
|
||||
|
||||
# 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)
|
||||
|
||||
async def main():
|
||||
# Apply litellm patch to fix the __annotations__ error
|
||||
patch_applied = fix_litellm_transcription_annotations()
|
||||
if not patch_applied:
|
||||
print(color("Something went wrong patching LiteLLM fix_litellm_transcription_annotations", color="red"))
|
||||
|
||||
# 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')}")
|
||||
|
||||
# Stream indicator
|
||||
print("\nAgent response (streaming):")
|
||||
print("-" * 40)
|
||||
print("Agent: ", end="", flush=True)
|
||||
|
||||
# Run the agent with a simple test message in streaming mode
|
||||
result = Runner.run_streamed(agent, "Explain what you can do in a few sentences. Don't use any tools.")
|
||||
|
||||
# Process the streaming response
|
||||
async for event in result.stream_events():
|
||||
if event.type == "raw_response_event" and isinstance(event.data, ResponseTextDeltaEvent):
|
||||
# Print the delta with a visible marker for each token
|
||||
# print(f"{event.data.delta}|", end="", flush=True)
|
||||
|
||||
print(f"{event.data.delta}", end="", flush=True)
|
||||
|
||||
print("\n" + "-" * 40)
|
||||
print("\nTest completed successfully!")
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
|
|
@ -1,30 +0,0 @@
|
|||
from cai.sdk.agents import Agent, Runner, AsyncOpenAI, OpenAIChatCompletionsModel
|
||||
import asyncio
|
||||
|
||||
spanish_agent = Agent(
|
||||
name="Spanish agent",
|
||||
instructions="You only speak Spanish.",
|
||||
model="o3-mini",
|
||||
)
|
||||
|
||||
english_agent = Agent(
|
||||
name="English agent",
|
||||
instructions="You only speak English",
|
||||
model=OpenAIChatCompletionsModel(
|
||||
model="gpt-4o",
|
||||
openai_client=AsyncOpenAI()
|
||||
),
|
||||
)
|
||||
|
||||
triage_agent = Agent(
|
||||
name="Triage agent",
|
||||
instructions="Handoff to the appropriate agent based on the language of the request.",
|
||||
handoffs=[spanish_agent, english_agent],
|
||||
model="gpt-3.5-turbo",
|
||||
)
|
||||
|
||||
async def main():
|
||||
result = await Runner.run(english_agent, input="Hey there, what's up?")
|
||||
print(result.final_output)
|
||||
|
||||
asyncio.run(main())
|
||||
|
|
@ -216,15 +216,43 @@ class OpenAIChatCompletionsModel(Model):
|
|||
)
|
||||
|
||||
# The usage is only available in the last chunk
|
||||
usage = chunk.usage
|
||||
if hasattr(chunk, 'usage'):
|
||||
usage = chunk.usage
|
||||
# For Ollama/LiteLLM streams that don't have usage attribute
|
||||
else:
|
||||
usage = None
|
||||
|
||||
if not chunk.choices or not chunk.choices[0].delta:
|
||||
# Handle different stream chunk formats
|
||||
if hasattr(chunk, 'choices') and chunk.choices:
|
||||
choices = chunk.choices
|
||||
elif hasattr(chunk, 'delta') and chunk.delta:
|
||||
# Some providers might return delta directly
|
||||
choices = [{"delta": chunk.delta}]
|
||||
elif isinstance(chunk, dict) and 'choices' in chunk:
|
||||
choices = chunk['choices']
|
||||
else:
|
||||
# Skip chunks that don't contain choice data
|
||||
continue
|
||||
|
||||
if not choices or len(choices) == 0:
|
||||
continue
|
||||
|
||||
# Get the delta content
|
||||
delta = choices[0].get('delta', None)
|
||||
if not delta and hasattr(choices[0], 'delta'):
|
||||
delta = choices[0].delta
|
||||
|
||||
if not delta:
|
||||
continue
|
||||
|
||||
delta = chunk.choices[0].delta
|
||||
|
||||
# Handle text
|
||||
if delta.content:
|
||||
content = None
|
||||
if hasattr(delta, 'content') and delta.content is not None:
|
||||
content = delta.content
|
||||
elif isinstance(delta, dict) and 'content' in delta and delta['content'] is not None:
|
||||
content = delta['content']
|
||||
|
||||
if content:
|
||||
if not state.text_content_index_and_output:
|
||||
# Initialize a content tracker for streaming text
|
||||
state.text_content_index_and_output = (
|
||||
|
|
@ -263,16 +291,22 @@ class OpenAIChatCompletionsModel(Model):
|
|||
# Emit the delta for this segment of content
|
||||
yield ResponseTextDeltaEvent(
|
||||
content_index=state.text_content_index_and_output[0],
|
||||
delta=delta.content,
|
||||
delta=content,
|
||||
item_id=FAKE_RESPONSES_ID,
|
||||
output_index=0,
|
||||
type="response.output_text.delta",
|
||||
)
|
||||
# Accumulate the text into the response part
|
||||
state.text_content_index_and_output[1].text += delta.content
|
||||
state.text_content_index_and_output[1].text += content
|
||||
|
||||
# Handle refusals (model declines to answer)
|
||||
refusal_content = None
|
||||
if hasattr(delta, 'refusal') and delta.refusal:
|
||||
refusal_content = delta.refusal
|
||||
elif isinstance(delta, dict) and 'refusal' in delta and delta['refusal']:
|
||||
refusal_content = delta['refusal']
|
||||
|
||||
if refusal_content:
|
||||
if not state.refusal_content_index_and_output:
|
||||
# Initialize a content tracker for streaming refusal text
|
||||
state.refusal_content_index_and_output = (
|
||||
|
|
@ -307,36 +341,66 @@ class OpenAIChatCompletionsModel(Model):
|
|||
# Emit the delta for this segment of refusal
|
||||
yield ResponseRefusalDeltaEvent(
|
||||
content_index=state.refusal_content_index_and_output[0],
|
||||
delta=delta.refusal,
|
||||
delta=refusal_content,
|
||||
item_id=FAKE_RESPONSES_ID,
|
||||
output_index=0,
|
||||
type="response.refusal.delta",
|
||||
)
|
||||
# Accumulate the refusal string in the output part
|
||||
state.refusal_content_index_and_output[1].refusal += delta.refusal
|
||||
state.refusal_content_index_and_output[1].refusal += refusal_content
|
||||
|
||||
# Handle tool calls
|
||||
# Because we don't know the name of the function until the end of the stream, we'll
|
||||
# save everything and yield events at the end
|
||||
tool_calls = None
|
||||
if hasattr(delta, 'tool_calls') and delta.tool_calls:
|
||||
for tc_delta in delta.tool_calls:
|
||||
if tc_delta.index not in state.function_calls:
|
||||
state.function_calls[tc_delta.index] = ResponseFunctionToolCall(
|
||||
tool_calls = delta.tool_calls
|
||||
elif isinstance(delta, dict) and 'tool_calls' in delta and delta['tool_calls']:
|
||||
tool_calls = delta['tool_calls']
|
||||
|
||||
if tool_calls:
|
||||
for tc_delta in tool_calls:
|
||||
tc_index = tc_delta.index if hasattr(tc_delta, 'index') else tc_delta.get('index', 0)
|
||||
if tc_index not in state.function_calls:
|
||||
state.function_calls[tc_index] = ResponseFunctionToolCall(
|
||||
id=FAKE_RESPONSES_ID,
|
||||
arguments="",
|
||||
name="",
|
||||
type="function_call",
|
||||
call_id="",
|
||||
)
|
||||
tc_function = tc_delta.function
|
||||
|
||||
state.function_calls[tc_delta.index].arguments += (
|
||||
tc_function.arguments if tc_function else ""
|
||||
) or ""
|
||||
state.function_calls[tc_delta.index].name += (
|
||||
tc_function.name if tc_function else ""
|
||||
) or ""
|
||||
state.function_calls[tc_delta.index].call_id += tc_delta.id or ""
|
||||
|
||||
tc_function = None
|
||||
if hasattr(tc_delta, 'function'):
|
||||
tc_function = tc_delta.function
|
||||
elif isinstance(tc_delta, dict) and 'function' in tc_delta:
|
||||
tc_function = tc_delta['function']
|
||||
|
||||
if tc_function:
|
||||
# Handle both object and dict formats
|
||||
args = ""
|
||||
if hasattr(tc_function, 'arguments'):
|
||||
args = tc_function.arguments or ""
|
||||
elif isinstance(tc_function, dict) and 'arguments' in tc_function:
|
||||
args = tc_function.get('arguments', "") or ""
|
||||
|
||||
name = ""
|
||||
if hasattr(tc_function, 'name'):
|
||||
name = tc_function.name or ""
|
||||
elif isinstance(tc_function, dict) and 'name' in tc_function:
|
||||
name = tc_function.get('name', "") or ""
|
||||
|
||||
state.function_calls[tc_index].arguments += args
|
||||
state.function_calls[tc_index].name += name
|
||||
|
||||
# Handle call_id in both formats
|
||||
call_id = ""
|
||||
if hasattr(tc_delta, 'id'):
|
||||
call_id = tc_delta.id or ""
|
||||
elif isinstance(tc_delta, dict) and 'id' in tc_delta:
|
||||
call_id = tc_delta.get('id', "") or ""
|
||||
|
||||
state.function_calls[tc_index].call_id += call_id
|
||||
|
||||
function_call_starting_index = 0
|
||||
if state.text_content_index_and_output:
|
||||
|
|
@ -423,31 +487,38 @@ class OpenAIChatCompletionsModel(Model):
|
|||
|
||||
final_response = response.model_copy()
|
||||
final_response.output = outputs
|
||||
|
||||
final_response.usage = (
|
||||
ResponseUsage(
|
||||
input_tokens=usage.prompt_tokens,
|
||||
output_tokens=usage.completion_tokens,
|
||||
total_tokens=usage.total_tokens,
|
||||
input_tokens=usage.prompt_tokens if usage and hasattr(usage, 'prompt_tokens') else 0,
|
||||
output_tokens=usage.completion_tokens if usage and hasattr(usage, 'completion_tokens') else 0,
|
||||
total_tokens=usage.total_tokens if usage and hasattr(usage, 'total_tokens') else 0,
|
||||
output_tokens_details=OutputTokensDetails(
|
||||
reasoning_tokens=usage.completion_tokens_details.reasoning_tokens
|
||||
if hasattr(usage, 'completion_tokens_details')
|
||||
if usage and hasattr(usage, 'completion_tokens_details')
|
||||
and usage.completion_tokens_details
|
||||
and hasattr(usage.completion_tokens_details, 'reasoning_tokens')
|
||||
and usage.completion_tokens_details.reasoning_tokens
|
||||
else 0
|
||||
),
|
||||
input_tokens_details={
|
||||
"prompt_tokens": usage.prompt_tokens if usage.prompt_tokens else 0,
|
||||
"prompt_tokens": usage.prompt_tokens if usage and hasattr(usage, 'prompt_tokens') else 0,
|
||||
"cached_tokens": usage.prompt_tokens_details.cached_tokens
|
||||
if hasattr(usage, 'prompt_tokens_details')
|
||||
if usage and hasattr(usage, 'prompt_tokens_details')
|
||||
and usage.prompt_tokens_details
|
||||
and hasattr(usage.prompt_tokens_details, 'cached_tokens')
|
||||
and usage.prompt_tokens_details.cached_tokens
|
||||
else 0
|
||||
},
|
||||
)
|
||||
if usage
|
||||
else None
|
||||
if usage is not None
|
||||
else ResponseUsage(
|
||||
input_tokens=0,
|
||||
output_tokens=0,
|
||||
total_tokens=0,
|
||||
output_tokens_details=OutputTokensDetails(reasoning_tokens=0),
|
||||
input_tokens_details={"prompt_tokens": 0, "cached_tokens": 0}
|
||||
)
|
||||
)
|
||||
|
||||
yield ResponseCompletedEvent(
|
||||
|
|
@ -462,6 +533,11 @@ class OpenAIChatCompletionsModel(Model):
|
|||
"input_tokens": usage.prompt_tokens,
|
||||
"output_tokens": usage.completion_tokens,
|
||||
}
|
||||
else:
|
||||
span_generation.span_data.usage = {
|
||||
"input_tokens": 0,
|
||||
"output_tokens": 0,
|
||||
}
|
||||
|
||||
@overload
|
||||
async def _fetch_response(
|
||||
|
|
@ -559,7 +635,38 @@ class OpenAIChatCompletionsModel(Model):
|
|||
"extra_headers": _HEADERS,
|
||||
}
|
||||
|
||||
# Previously articulated as
|
||||
# ret = await self._get_client().chat.completions.create(**kwargs)
|
||||
# but now we're using litellm
|
||||
|
||||
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"]
|
||||
}
|
||||
|
||||
# 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(
|
||||
|
|
@ -590,18 +697,20 @@ class OpenAIChatCompletionsModel(Model):
|
|||
},
|
||||
)
|
||||
# Get the streaming object
|
||||
ollama_api_base = get_ollama_api_base().rstrip('/v1') # Remove /v1 if present
|
||||
stream_obj = await litellm.acompletion(
|
||||
**kwargs,
|
||||
api_base=get_ollama_api_base(),
|
||||
custom_llm_provider="openai"
|
||||
**ollama_kwargs,
|
||||
api_base=ollama_api_base,
|
||||
custom_llm_provider="ollama"
|
||||
)
|
||||
return response, stream_obj
|
||||
else:
|
||||
# Non-streaming mode
|
||||
ollama_api_base = get_ollama_api_base().rstrip('/v1') # Remove /v1 if present
|
||||
ret = litellm.completion(
|
||||
**kwargs,
|
||||
api_base=get_ollama_api_base(),
|
||||
custom_llm_provider="openai"
|
||||
**ollama_kwargs,
|
||||
api_base=ollama_api_base,
|
||||
custom_llm_provider="ollama"
|
||||
)
|
||||
return ret
|
||||
else:
|
||||
|
|
|
|||
Loading…
Reference in New Issue