merge with reset-v0.4.0

This commit is contained in:
Mery-Sanz 2025-05-07 10:36:30 +02:00
commit 7fe22d7312
17 changed files with 2240 additions and 491 deletions

View File

View File

@ -175,7 +175,8 @@ def run_cai_cli(starting_agent, context_variables=None, stream=False, max_turns=
ACTIVE_TIME = 0 ACTIVE_TIME = 0
idle_time = 0 idle_time = 0
console = Console() console = Console()
last_model = os.getenv('CAI_MODEL', 'qwen2.5:14b')
last_agent_type = os.getenv('CAI_AGENT_TYPE', 'one_tool_agent')
# Initialize command completer and key bindings # Initialize command completer and key bindings
command_completer = FuzzyCommandCompleter() command_completer = FuzzyCommandCompleter()
current_text = [''] current_text = ['']
@ -208,6 +209,36 @@ def run_cai_cli(starting_agent, context_variables=None, stream=False, max_turns=
while turn_count < max_turns: while turn_count < max_turns:
try: try:
idle_start_time = time.time() idle_start_time = time.time()
# Check if model has changed and update if needed
current_model = os.getenv('CAI_MODEL', 'qwen2.5:14b')
if current_model != last_model and hasattr(agent, 'model'):
# Update the model in the agent
if hasattr(agent.model, 'model'):
agent.model.model = current_model
last_model = current_model
# Check if agent type has changed and recreate agent if needed
current_agent_type = os.getenv('CAI_AGENT_TYPE', 'one_tool_agent')
if current_agent_type != last_agent_type:
try:
# Import is already at the top level
agent = get_agent_by_name(current_agent_type)
last_agent_type = current_agent_type
# Configure the new agent's model flags
if hasattr(agent, 'model'):
if hasattr(agent.model, 'disable_rich_streaming'):
agent.model.disable_rich_streaming = True
if hasattr(agent.model, 'suppress_final_output'):
agent.model.suppress_final_output = True
# Apply current model to the new agent
if hasattr(agent.model, 'model'):
agent.model.model = current_model
except Exception as e:
console.print(f"[red]Error switching agent: {str(e)}[/red]")
# Get user input with command completion and history # Get user input with command completion and history
user_input = get_user_input( user_input = get_user_input(
command_completer, command_completer,
@ -280,8 +311,9 @@ def run_cai_cli(starting_agent, context_variables=None, stream=False, max_turns=
if commands_handle_command(command, args): if commands_handle_command(command, args):
continue # Command was handled, continue to next iteration continue # Command was handled, continue to next iteration
# If command wasn't recognized, show error # If command wasn't recognized, show error (skip for /shell or /s)
console.print(f"[red]Unknown command: {command}[/red]") if command not in ("/shell", "/s"):
console.print(f"[red]Unknown command: {command}[/red]")
continue continue
# Process the conversation with the agent # Process the conversation with the agent

View File

@ -25,7 +25,6 @@ from cai.repl.commands.base import (
# Import all command modules # Import all command modules
# These imports will register the commands with the registry # These imports will register the commands with the registry
from cai.repl.commands import ( # pylint: disable=import-error,unused-import,line-too-long,redefined-builtin # noqa: E501,F401 from cai.repl.commands import ( # pylint: disable=import-error,unused-import,line-too-long,redefined-builtin # noqa: E501,F401
memory,
help, help,
graph, graph,
exit, exit,
@ -34,10 +33,11 @@ from cai.repl.commands import ( # pylint: disable=import-error,unused-import,li
platform, platform,
kill, kill,
model, model,
turns,
agent, agent,
history, history,
config config,
flush,
workspace,
) )
# Define helper functions # Define helper functions

View File

@ -155,45 +155,37 @@ class AgentCommand(Command):
table = Table(title="Available Agents") table = Table(title="Available Agents")
table.add_column("#", style="dim") table.add_column("#", style="dim")
table.add_column("Name", style="cyan") table.add_column("Name", style="cyan")
table.add_column("Module", style="magenta") table.add_column("Key", style="magenta")
table.add_column("Module", style="green")
table.add_column("Description", style="green") table.add_column("Description", style="green")
table.add_column("Pattern", style="blue")
table.add_column("Model", style="yellow")
# Scan all agents from the agents folder # Retrieve all registered agents
agents_to_display = get_available_agents() agents_to_display = get_available_agents()
# Display all agents for idx, (agent_key, agent) in enumerate(agents_to_display.items(), start=1):
for i, (name, agent) in enumerate(agents_to_display.items(), 1): # Human-friendly name (falls back to the dict key)
description = agent.description display_name = getattr(agent, "name", agent_key)
if not description and hasattr(agent, 'instructions'):
if callable(agent.instructions): # Use provided description, otherwise derive from instructions
description = agent.instructions(context_variables={}) description = getattr(agent, "description", "") or ""
else: if not description and hasattr(agent, "instructions"):
description = agent.instructions instr = agent.instructions
# Clean up description - remove newlines and strip spaces description = instr(context_variables={}) if callable(instr) else instr
if isinstance(description, str): if isinstance(description, str):
description = " ".join(description.split()) description = " ".join(description.split())
if len(description) > 50: if len(description) > 50:
description = description[:47] + "..." description = description[:47] + "..."
# Get the module name for the agent # Module where this agent lives
module_name = get_agent_module(name) module_name = get_agent_module(agent_key)
# Get the pattern if it exists # Add a row with all collected info
pattern = getattr(agent, 'pattern', '')
if pattern:
pattern = pattern.capitalize()
# Handle model display based on agent type
model_display = self._get_model_display(name, agent)
table.add_row( table.add_row(
str(i), str(idx),
name, display_name,
agent_key,
module_name, module_name,
description, description
pattern,
model_display
) )
console.print(table) console.print(table)
@ -210,76 +202,45 @@ class AgentCommand(Command):
""" """
if not args: if not args:
console.print("[red]Error: No agent specified[/red]") console.print("[red]Error: No agent specified[/red]")
console.print("Usage: /agent select <name|number>") console.print("Usage: /agent select <agent_key|number>")
return False return False
agent_id = args[0] agent_id = args[0]
# Get the list of available agents
agents_to_display = get_available_agents() agents_to_display = get_available_agents()
agent_list = list(agents_to_display.items())
# Check if agent_id is a number # Check if agent_id is a number
if agent_id.isdigit(): if agent_id.isdigit():
index = int(agent_id) index = int(agent_id)
if 1 <= index <= len(agents_to_display): if 1 <= index <= len(agent_list):
agent_name = list(agents_to_display.keys())[index - 1] # Get the agent tuple from the list
selected_agent_key, selected_agent = agent_list[index - 1]
agent_name = getattr(selected_agent, "name", selected_agent_key)
agent = selected_agent
else: else:
console.print( console.print(f"[red]Error: Invalid agent number: {agent_id}[/red]")
f"[red]Error: Invalid agent number: {agent_id}[/red]")
return False return False
else: else:
# Treat as agent name # Treat as agent key
agent_name = agent_id selected_agent_key = None
if agent_name not in agents_to_display: for key, agent_obj in agents_to_display.items():
console.print(f"[red]Error: Unknown agent: {agent_name}[/red]") if key == agent_id:
agent = agent_obj
selected_agent_key = key
agent_name = getattr(agent_obj, "name", key)
break
else:
console.print(f"[red]Error: Unknown agent key: {agent_id}[/red]")
return False return False
# Get the agent # Set the agent key in environment variable (not the agent name)
agent = agents_to_display[agent_name] os.environ["CAI_AGENT_TYPE"] = selected_agent_key
# Set the agent as the current agent in the REPL console.print(
# We need to avoid circular imports, so we'll use a different approach f"[green]Switched to agent: {agent_name}[/green]")
# to access the client and current_agent variables visualize_agent_graph(agent)
return True
# Import the module dynamically to avoid circular imports
if 'cai.repl.repl' in sys.modules:
repl_module = sys.modules['cai.repl.repl']
# Check if client is initialized
if hasattr(repl_module, 'client') and repl_module.client:
# Update the active_agent in the client
repl_module.client.active_agent = agent
# Update the global current_agent variable if it exists
if hasattr(repl_module, 'current_agent'):
repl_module.current_agent = agent
# Update the global agent variable if it exists
if hasattr(repl_module, 'agent'):
repl_module.agent = agent
# Also update the agent variable in the run_demo_loop
# function's frame if possible
try:
for frame_info in inspect.stack():
frame = frame_info.frame
if ('run_demo_loop' in frame.f_code.co_name and
'agent' in frame.f_locals):
frame.f_locals['agent'] = agent
break
except Exception: # pylint: disable=broad-except # nosec
# If this fails, we still have the global current_agent as
# a fallback
pass
console.print(
f"[green]Switched to agent: {agent_name}[/green]")
visualize_agent_graph(agent)
return True
console.print("[red]Error: CAI client not initialized[/red]")
return False
console.print("[red]Error: REPL module not initialized[/red]")
return False
def handle_info(self, args: Optional[List[str]] = None) -> bool: def handle_info(self, args: Optional[List[str]] = None) -> bool:
"""Handle /agent info command. """Handle /agent info command.
@ -292,57 +253,74 @@ class AgentCommand(Command):
""" """
if not args: if not args:
console.print("[red]Error: No agent specified[/red]") console.print("[red]Error: No agent specified[/red]")
console.print("Usage: /agent info <name|number>") console.print("Usage: /agent info <agent_key|number>")
return False return False
agent_id = args[0] agent_id = args[0]
# Get the list of available agents # Get available agents
agents_to_display = get_available_agents() agents_to_display = get_available_agents()
# Check if agent_id is a number # Resolve agent_id to an agent key (by index or name)
if agent_id.isdigit(): if agent_id.isdigit():
index = int(agent_id) idx = int(agent_id)
if 1 <= index <= len(agents_to_display): if not (1 <= idx <= len(agents_to_display)):
agent_name = list(agents_to_display.keys())[index - 1] console.print(f"[red]Error: Invalid agent number: {agent_id}[/red]")
else:
console.print(
f"[red]Error: Invalid agent number: {agent_id}[/red]")
return False return False
agent_key = list(agents_to_display.keys())[idx - 1]
else: else:
# Treat as agent name agent_key = None
agent_name = agent_id for key, ag in agents_to_display.items():
if agent_name not in agents_to_display: if key == agent_id or getattr(ag, "name", "").lower() == agent_id.lower():
console.print(f"[red]Error: Unknown agent: {agent_name}[/red]") agent_key = key
break
if agent_key is None:
console.print(f"[red]Error: Unknown agent key: {agent_id}[/red]")
return False return False
# Get the agent agent = agents_to_display[agent_key]
agent = agents_to_display[agent_name]
# Display agent information # Display agent information
instructions = agent.instructions instructions = agent.instructions
if callable(instructions): if callable(instructions):
instructions = instructions() instructions = instructions()
# Prepare agent properties
name = agent.name or agent_key
description = getattr(agent, "description", None) or "N/A"
clean_description = " ".join(line.strip() for line in description.splitlines())
functions = getattr(agent, "functions", [])
parallel = getattr(agent, "parallel_tool_calls", False)
handoff_desc = getattr(agent, "handoff_description", None) or "N/A"
handoffs = getattr(agent, "handoffs", [])
tools = getattr(agent, "tools", [])
guardrails_in = getattr(agent, "input_guardrails", [])
guardrails_out = getattr(agent, "output_guardrails", [])
output_type = getattr(agent, "output_type", None) or "N/A"
hooks = getattr(agent, "hooks", []) or []
# Handle model display based on agent type # Build markdown content for agent info
model_display = self._get_model_display_for_info(agent_name, agent)
# Create a markdown table for agent details
markdown_content = f""" markdown_content = f"""
# Agent: {agent_name} # Agent Info: {name}
| Property | Value | | Property | Value |
|----------|-------| |------------------------|-------------------------------|
| Name | {agent.name} | | Key | {agent_key} |
| Model | {model_display} | | Name | {name} |
| Functions | {len(agent.functions)} | | Description | {clean_description} |
| Parallel Tool Calls | {'Yes' if agent.parallel_tool_calls else 'No'} | | Functions | {len(functions)} |
| Parallel Tool Calls | {"Yes" if parallel else "No"} |
| Handoff Description | {handoff_desc} |
| Handoffs | {len(handoffs)} |
| Tools | {len(tools)} |
| Input Guardrails | {len(guardrails_in)} |
| Output Guardrails | {len(guardrails_out)} |
| Output Type | {output_type} |
| Hooks | {len(hooks)} |
## Instructions ## Instructions
{instructions} {instructions}
"""
"""
console.print(Markdown(markdown_content)) console.print(Markdown(markdown_content))
return True return True

View File

@ -122,6 +122,16 @@ ENV_VARS = {
"name": "CAI_SUPPORT_INTERVAL", "name": "CAI_SUPPORT_INTERVAL",
"description": "Number of turns between support agent executions", "description": "Number of turns between support agent executions",
"default": "5" "default": "5"
},
22: {
"name": "CAI_STREAM",
"description": "Boolean to enable real-time, chunked responses instead of full messages.",
"default": "True"
},
23: {
"name": "CAI_WORKSPACE",
"description": "Name of the current workspace (affects log file naming)",
"default": None
}, },
} }

View File

@ -0,0 +1,120 @@
"""
Flush command for CAI REPL.
This module provides commands for clear the context.
"""
import os
from typing import (
Dict,
List,
Optional
)
from rich.console import Console # pylint: disable=import-error
from rich.panel import Panel # pylint: disable=import-error
from cai.util import get_model_input_tokens
from cai.repl.commands.base import Command, register_command
from cai.sdk.agents.models.openai_chatcompletions import message_history
console = Console()
class FlushCommand(Command):
"""Command to flush the conversation history."""
def __init__(self):
"""Initialize the flush command."""
super().__init__(
name="/flush",
description="Clear the current conversation history.",
aliases=["/clear"]
)
def handle_no_args(self, messages: Optional[List[Dict]] = None) -> bool:
"""Handle the flush command when no args are provided.
Args:
messages: The conversation history messages
Returns:
True if the command was handled successfully
"""
# Use both the local messages parameter and the global message_history
local_messages = messages or []
global_history_length = len(message_history)
# Get token usage information before clearing
token_info = ""
context_usage = ""
# Access client through a function to avoid circular imports
# We can use globals() to get the client at runtime
client = self._get_client()
if client and hasattr(client, 'interaction_input_tokens') and hasattr(
client, 'total_input_tokens'):
model = os.getenv('CAI_MODEL', "qwen2.5:14b")
input_tokens = client.interaction_input_tokens if hasattr(
client, 'interaction_input_tokens') else 0
total_tokens = client.total_input_tokens if hasattr(
client, 'total_input_tokens') else 0
max_tokens = get_model_input_tokens(model)
context_pct = (input_tokens / max_tokens) * \
100 if max_tokens > 0 else 0
token_info = f"Current tokens: {input_tokens}, Total tokens: {total_tokens}"
context_usage = f"Context usage: {context_pct:.1f}% of {max_tokens} tokens"
# Clear both the local messages list and the global message_history
if local_messages:
local_messages.clear()
# Always clear the global message history
message_history.clear()
# Determine which length to report (use the greater of the two)
initial_length = max(len(local_messages) if messages else 0, global_history_length)
# Display information about the cleared messages
if initial_length > 0:
content = [
f"Conversation history cleared. Removed {initial_length} messages."
]
if token_info:
content.append(token_info)
if context_usage:
content.append(context_usage)
console.print(Panel(
"\n".join(content),
title="[bold cyan]Context Flushed[/bold cyan]",
border_style="blue",
padding=(1, 2)
))
else:
console.print(Panel(
"No conversation history to clear.",
title="[bold cyan]Context Flushed[/bold cyan]",
border_style="blue",
padding=(1, 2)
))
return True
def _get_client(self):
"""Get the CAI client from the global namespace.
This function avoids circular imports by accessing the client
at runtime instead of import time.
Returns:
The global CAI client instance or None if not available
"""
try:
# Import here to avoid circular import
from cai.repl.repl import client as global_client # pylint: disable=import-outside-toplevel # noqa: E501
return global_client
except (ImportError, AttributeError):
return None
# Register the /flush command
register_command(FlushCommand())

View File

@ -1,20 +1,64 @@
""" """
Graph command for CAI REPL. Graph command for CAI cli.
This module provides commands for visualizing the agent interaction graph. This module provides commands for visualizing the agent interaction graph.
It allows users to display a simple directed graph of the conversation history,
showing the sequence of user and agent interactions, including tool calls.
""" """
from typing import ( from typing import List, Optional
List,
Optional
)
from rich.console import Console # pylint: disable=import-error from rich.console import Console # pylint: disable=import-error
from rich.panel import Panel
from cai.repl.commands.base import Command, register_command from cai.repl.commands.base import Command, register_command
import os
import importlib.util
console = Console() console = Console()
def find_agent_name_by_instructions(target_instructions: str, agents_dir: str) -> Optional[str]:
"""
Search all Python files in the agents directory for an agent whose 'instructions'
attribute matches the given target_instructions (ignoring leading/trailing whitespace).
Returns the agent's 'name' attribute if found, otherwise None.
Args:
target_instructions (str): The instructions string to match.
agents_dir (str): The directory containing agent files.
Returns:
Optional[str]: The agent name if found, else None.
"""
for filename in os.listdir(agents_dir):
if not filename.endswith(".py") or filename.startswith("__"):
continue
filepath = os.path.join(agents_dir, filename)
try:
spec = importlib.util.spec_from_file_location("agent_mod", filepath)
agent_mod = importlib.util.module_from_spec(spec)
spec.loader.exec_module(agent_mod)
for attr_name in dir(agent_mod):
attr = getattr(agent_mod, attr_name)
if hasattr(attr, "instructions"):
agent_instructions = getattr(attr, "instructions", None)
if agent_instructions and agent_instructions.strip() == target_instructions.strip():
agent_name = getattr(attr, "name", None)
if agent_name:
return agent_name
except Exception:
continue
return None
class GraphCommand(Command): class GraphCommand(Command):
"""Command for visualizing the agent interaction graph.""" """
Command for visualizing the agent interaction graph.
This command displays a directed graph of the conversation history,
showing the sequence of user and agent messages, and highlighting
tool calls made by the agent.
"""
def __init__(self): def __init__(self):
"""Initialize the graph command.""" """Initialize the graph command."""
@ -25,31 +69,104 @@ class GraphCommand(Command):
) )
def handle(self, args: Optional[List[str]] = None) -> bool: def handle(self, args: Optional[List[str]] = None) -> bool:
"""Handle the graph command. """
Handle the /graph command.
Args: Args:
args: Optional list of command arguments args: Optional list of command arguments
Returns: Returns:
True if the command was handled successfully, False otherwise bool: True if the command was handled successfully, False otherwise.
""" """
return self.handle_graph_show() return self.handle_graph_show()
def handle_graph_show(self) -> bool: def handle_graph_show(self) -> bool:
"""Handle /graph show command""" """Handle /graph show command"""
from cai.repl.repl import client # pylint: disable=import-error from cai.sdk.agents.models.openai_chatcompletions import message_history
if not message_history:
# Import here to avoid circular imports
if not client or not client._graph: # pylint: disable=protected-access
console.print("[yellow]No conversation graph available.[/yellow]") console.print("[yellow]No conversation graph available.[/yellow]")
return True return True
try: try:
agents_dir = os.path.join(os.path.dirname(__file__), "../../agents")
agents_dir = os.path.abspath(agents_dir)
import networkx as nx
G = nx.DiGraph()
last_agent_name = None
prev_node_idx = None # Track the last node actually added (not system)
for idx, msg in enumerate(message_history):
role = msg.get("role", "unknown")
# If the message is from the system, update last_agent_name but do not add a node
if role == "system":
system_msg = msg.get("content", "").strip()
agent_name = find_agent_name_by_instructions(system_msg, agents_dir)
if agent_name:
last_agent_name = agent_name
continue
label = role
extra_info = ""
if role == "assistant":
if last_agent_name:
label = last_agent_name
else:
label = "assistant"
if msg.get("tool_calls"):
tool_call = msg["tool_calls"][0]
if tool_call.get("function"):
func_name = tool_call["function"].get("name", "")
func_args = tool_call["function"].get("arguments", "")
extra_info = f"\n[cyan]Tool:[/cyan] [bold]{func_name}[/bold]\n[cyan]Args:[/cyan] {func_args}"
elif role == "user":
user_content = msg.get("content", "")
if user_content:
extra_info = f"\n{user_content}"
label = role
else:
label = role
G.add_node(idx, role=label, extra_info=extra_info)
if prev_node_idx is not None:
G.add_edge(prev_node_idx, idx)
prev_node_idx = idx
def ascii_graph(G):
"""
Render the conversation graph as a sequence of panels with arrows.
Args:
G (networkx.DiGraph): The conversation graph.
Returns:
List: List of rich Panel objects and arrow strings.
"""
lines = []
node_list = list(G.nodes(data=True))
for i, (idx, data) in enumerate(node_list):
role = data.get("role", "unknown")
extra_info = data.get("extra_info", "")
role_fmt = f"[bold][blue]{role[:1].upper()}{role[1:]}[/blue][/bold]"
panel_content = f"{role_fmt}"
if extra_info:
panel_content += f"{extra_info}"
panel = Panel(
panel_content,
expand=False,
border_style="cyan"
)
lines.append(panel)
if i < len(node_list) - 1:
lines.append("[cyan] │\n\n ▼[/cyan]")
return lines
console.print("\n[bold]Conversation Graph:[/bold]") console.print("\n[bold]Conversation Graph:[/bold]")
console.print("------------------") console.print("------------------")
console.print( if len(G.nodes) == 0:
client._graph.ascii()) # pylint: disable=protected-access console.print("[yellow]No messages to display in graph.[/yellow]")
else:
for item in ascii_graph(G):
console.print(item)
console.print() console.print()
return True return True
except Exception as e: # pylint: disable=broad-except except Exception as e: # pylint: disable=broad-except

View File

@ -2,8 +2,12 @@
History command for CAI REPL. History command for CAI REPL.
This module provides commands for displaying conversation history. This module provides commands for displaying conversation history.
""" """
import json
from typing import Any, Dict, List, Optional
from rich.console import Console # pylint: disable=import-error from rich.console import Console # pylint: disable=import-error
from rich.table import Table # pylint: disable=import-error from rich.table import Table # pylint: disable=import-error
from rich.text import Text # pylint: disable=import-error
from cai.repl.commands.base import Command, register_command from cai.repl.commands.base import Command, register_command
@ -20,6 +24,20 @@ class HistoryCommand(Command):
description="Display the conversation history", description="Display the conversation history",
aliases=["/h"] aliases=["/h"]
) )
def handle(self, args: Optional[List[str]] = None,
messages: Optional[List[Dict]] = None) -> bool:
"""Handle the history command.
Args:
args: Optional list of command arguments
messages: Optional list of conversation messages
Returns:
True if the command was handled successfully, False otherwise
"""
# Currently, the history command doesn't take any arguments
return self.handle_no_args()
def handle_no_args(self) -> bool: def handle_no_args(self) -> bool:
"""Handle the command when no arguments are provided. """Handle the command when no arguments are provided.
@ -27,15 +45,15 @@ class HistoryCommand(Command):
Returns: Returns:
True if the command was handled successfully, False otherwise True if the command was handled successfully, False otherwise
""" """
# Access messages directly from repl.py's global scope # Access messages directly from openai_chatcompletions.py
try: try:
from cai.repl.repl import messages # pylint: disable=import-outside-toplevel # noqa: E501 from cai.sdk.agents.models.openai_chatcompletions import message_history # pylint: disable=import-outside-toplevel # noqa: E501
except ImportError: except ImportError:
console.print( console.print(
"[red]Error: Could not access conversation history[/red]") "[red]Error: Could not access conversation history[/red]")
return False return False
if not messages: if not message_history:
console.print("[yellow]No conversation history available[/yellow]") console.print("[yellow]No conversation history available[/yellow]")
return True return True
@ -50,34 +68,87 @@ class HistoryCommand(Command):
table.add_column("Content", style="green") table.add_column("Content", style="green")
# Add messages to the table # Add messages to the table
for idx, msg in enumerate(messages, 1): for idx, msg in enumerate(message_history, 1):
role = msg.get("role", "unknown") try:
content = msg.get("content", "") role = msg.get("role", "unknown")
content = msg.get("content", "")
tool_calls = msg.get("tool_calls", None)
# Truncate long content for better display # Create formatted content based on message type
if len(content) > 100: formatted_content = self._format_message_content(
content = content[:97] + "..." content, tool_calls)
# Color the role based on type # Color the role based on type
if role == "user": if role == "user":
role_style = "cyan" role_style = "cyan"
elif role == "assistant": elif role == "assistant":
role_style = "yellow" role_style = "yellow"
else: else:
role_style = "red" role_style = "red"
# Add a newline between each role for better readability # Add a newline between each role for better readability
if idx > 1: if idx > 1:
table.add_row("", "", "") table.add_row("", "", "")
table.add_row( table.add_row(
str(idx), str(idx),
f"[{role_style}]{role}[/{role_style}]", f"[{role_style}]{role}[/{role_style}]",
content formatted_content
) )
except Exception as e:
# Log error but continue with next message
console.print(f"[red]Error displaying message {idx}: {e}[/red]")
continue
console.print(table) console.print(table)
return True return True
def _format_message_content(
self, content: Any, tool_calls: List[Dict[str, Any]]
) -> str:
"""Format message content for display, handling both text and tool calls.
Args:
content: Text content of the message
tool_calls: List of tool calls if present
Returns:
Formatted string representation of the message content
"""
if tool_calls:
# Format tool calls into a readable string
result = []
for tc in tool_calls:
func_details = tc.get("function", {})
func_name = func_details.get("name", "unknown_function")
# Format arguments (pretty-print JSON if possible)
args_str = func_details.get("arguments", "{}")
try:
# Parse and re-format JSON for better readability
args_dict = json.loads(args_str)
args_formatted = json.dumps(args_dict, indent=2)
# Limit to first 100 chars for display
if len(args_formatted) > 100:
args_formatted = args_formatted[:97] + "..."
except (json.JSONDecodeError, TypeError):
# If not valid JSON, use as is
args_formatted = args_str
if len(args_formatted) > 100:
args_formatted = args_formatted[:97] + "..."
result.append(f"Function: [bold blue]{func_name}[/bold blue]")
result.append(f"Args: {args_formatted}")
return "\n".join(result)
elif content:
# Regular text content (truncate if too long)
if len(content) > 100:
return content[:97] + "..."
return content
else:
# No content or tool calls (empty message)
return "[dim italic]Empty message[/dim italic]"
# Register the command # Register the command

View File

@ -1,4 +0,0 @@
"""
Memory command for CAI REPL.
This module provides commands for managing memory collections.
"""

View File

@ -145,6 +145,9 @@ class ModelCommand(Command):
{"name": "gpt-4-turbo", {"name": "gpt-4-turbo",
"description": "Fast and powerful GPT-4 model"} "description": "Fast and powerful GPT-4 model"}
], ],
"OpenAI GPT-4o-mini": [
{"name": "gpt-4o-mini", "description": " GPT-4o mini model"}
],
"OpenAI GPT-4.5": [ "OpenAI GPT-4.5": [
{ {
"name": "gpt-4.5-preview", "name": "gpt-4.5-preview",

View File

@ -1,96 +0,0 @@
"""
Turns command for CAI REPL.
This module provides commands for viewing and changing the maximum number
of turns.
"""
import os
from typing import (
List,
Optional
)
from rich.console import Console # pylint: disable=import-error
from rich.panel import Panel # pylint: disable=import-error
from cai.repl.commands.base import Command, register_command
console = Console()
class TurnsCommand(Command):
"""Command for viewing and changing the maximum number of turns."""
def __init__(self):
"""Initialize the turns command."""
super().__init__(
name="/turns",
description="View or change the maximum number of turns",
aliases=["/t"]
)
def handle(self, args: Optional[List[str]] = None) -> bool:
"""Handle the turns command.
Args:
args: Optional list of command arguments
Returns:
True if the command was handled successfully, False otherwise
"""
return self.handle_turns_command(args)
def handle_turns_command(self, args: List[str]) -> bool:
"""Change the maximum number of turns for CAI.
Args:
args: List containing the number of turns
Returns:
bool: True if the max turns was changed successfully
"""
if not args:
# Display current max turns
max_turns_info = os.getenv("CAI_MAX_TURNS", "inf")
console.print(Panel(
f"Current maximum turns: [bold green]{
max_turns_info}[/bold green]",
border_style="green",
title="Max Turns Setting"
))
# Usage instructions
console.print(
"\n[cyan]Usage:[/cyan] [bold]/turns <number_of_turns>[/bold]")
console.print("[cyan]Examples:[/cyan]")
console.print(" [bold]/turns 10[/bold] - Limit to 10 turns")
console.print(" [bold]/turns inf[/bold] - Unlimited turns")
return True
try:
turns = args[0]
# Check if it's a number or 'inf'
if turns.lower() == 'inf':
turns = 'inf'
else:
turns = int(turns)
# Set the max turns in environment variable
os.environ["CAI_MAX_TURNS"] = turns
console.print(Panel(
f"Maximum turns changed to: [bold green]{turns}[/bold green]\n"
"[yellow]Note: This will take effect on the next run[/yellow]",
border_style="green",
title="Max Turns Changed"
))
return True
except ValueError:
console.print(Panel(
"Error: Max turns must be a number or 'inf'",
border_style="red",
title="Invalid Input"
))
return False
# Register the command
register_command(TurnsCommand())

View File

@ -0,0 +1,687 @@
"""
Virtualization command for CAI REPL.
This module provides commands for setting up and managing Docker virtualization
environments.
"""
# Standard library imports
import os
import json
import subprocess
import datetime
import time
from typing import List, Optional, Dict, Any, Tuple
# Third-party imports
from rich.console import Console
from rich.table import Table
from rich.panel import Panel
from rich.markdown import Markdown
import rich.box
# Local imports
from cai.repl.commands.base import Command, register_command
console = Console()
class WorkspaceCommand(Command):
"""Command for workspace management within Docker containers or locally."""
def __init__(self):
"""Initialize the workspace command."""
super().__init__(
name="/workspace",
description=(
"Set or display the current workspace name and manage files."
" Affects log file naming and where files are stored."
),
aliases=["/ws"]
)
# Add subcommands
self.add_subcommand(
"set",
"Set the current workspace name",
self.handle_set
)
self.add_subcommand(
"get",
"Display the current workspace name",
self.handle_get
)
self.add_subcommand(
"ls",
"List files in the workspace",
self.handle_ls_subcommand
)
self.add_subcommand(
"exec",
"Execute a command in the workspace",
self.handle_exec_subcommand
)
self.add_subcommand(
"copy",
"Copy files between host and container",
self.handle_copy_subcommand
)
def handle(self, args: Optional[List[str]] = None) -> bool:
"""Handle the workspace command.
Args:
args: Optional list of command arguments
Returns:
True if the command was handled successfully, False otherwise
"""
# If there are subcommands, process them
if args and args[0] in self.subcommands:
return super().handle(args)
# No arguments means show workspace info (same as get)
return self.handle_get()
def handle_no_args(self) -> bool:
"""Handle the command when no arguments are provided."""
return self.handle_get()
def handle_get(self, _: Optional[List[str]] = None) -> bool:
"""Display the current workspace name and directory information."""
# Get workspace info
workspace_name = os.getenv("CAI_WORKSPACE", None)
# Check if a container is active
active_container = os.getenv("CAI_ACTIVE_CONTAINER", "")
# Determine environment (container or host)
if active_container:
try:
# Get container details
result = subprocess.run(
["docker", "inspect", active_container],
capture_output=True,
text=True,
check=False
)
if result.returncode == 0:
container_info = json.loads(result.stdout)
if container_info:
image = container_info[0].get("Config", {}).get("Image", "unknown")
env_type = "container"
env_name = f"Container ({image})"
# For containers, if workspace is set, use container workspace path
# otherwise use root directory
if workspace_name:
# This will create the workspace in the container if it doesn't exist
workspace_dir = f"/workspace/workspaces/{workspace_name}"
# Ensure the directory exists in the container
subprocess.run(
["docker", "exec", active_container, "mkdir", "-p", workspace_dir],
capture_output=True,
check=False
)
else:
workspace_dir = "/"
else:
env_type = "host"
env_name = "Host System (container not running)"
# Use common._get_workspace_dir() for consistency
try:
from cai.tools.common import _get_workspace_dir as get_common_workspace_dir
workspace_dir = get_common_workspace_dir()
except ImportError:
workspace_dir = os.getcwd() # Basic fallback
except Exception:
env_type = "host"
env_name = "Host System (error inspecting container)"
# Use common._get_workspace_dir() for consistency
try:
from cai.tools.common import _get_workspace_dir as get_common_workspace_dir
workspace_dir = get_common_workspace_dir()
except ImportError:
workspace_dir = os.getcwd() # Basic fallback
else:
env_type = "host"
env_name = "Host System"
# Use common._get_workspace_dir() for consistency
try:
from cai.tools.common import _get_workspace_dir as get_common_workspace_dir
workspace_dir = get_common_workspace_dir()
except ImportError:
workspace_dir = os.getcwd() # Basic fallback
# Show workspace information
console.print(
Panel(
f"Current workspace: [bold green]{workspace_name or 'None'}[/bold green]\n"
f"Working in environment: [bold]{env_name}[/bold]\n"
f"Workspace directory: [bold]{workspace_dir}[/bold]",
title="Workspace Information",
border_style="green"
)
)
# Show available workspace commands
console.print("\n[cyan]Workspace Commands:[/cyan]")
console.print(
" [bold]/workspace set <name>[/bold] - "
"Set the current workspace name")
console.print(
" [bold]/workspace ls[/bold] - "
"List files in the workspace")
console.print(
" [bold]/workspace exec <cmd>[/bold] - "
"Execute a command in the workspace")
if active_container:
console.print(
" [bold]/workspace copy <src> <dst>[/bold] - "
"Copy files between host and container")
# List contents of the workspace
self._list_workspace_contents(env_type, workspace_dir)
return True
def handle_set(self, args: Optional[List[str]] = None) -> bool:
"""Set the current workspace name """
if not args or len(args) != 1:
console.print(
"[yellow]Usage: /workspace set <workspace_name>[/yellow]"
)
return False
workspace_name = args[0]
# Allow alphanumeric, underscores, hyphens
if not all(c.isalnum() or c in ['_', '-'] for c in workspace_name):
console.print(
"[red]Invalid workspace name. "
"Use alphanumeric, underscores, or hyphens only.[/red]"
)
return False
# Import the necessary modules for setting environment variables
# And for getting workspace dir consistently
try:
from cai.repl.commands.config import set_env_var
from cai.tools.common import _get_workspace_dir as get_common_workspace_dir
from cai.tools.common import _get_container_workspace_path as get_common_container_path
# Set the environment variable
if not set_env_var("CAI_WORKSPACE", workspace_name):
console.print(
"[red]Failed to set workspace environment variable.[/red]"
)
return False
except ImportError:
# Fallback if import fails
os.environ["CAI_WORKSPACE"] = workspace_name
# Define basic fallbacks for path functions if import failed
def get_common_workspace_dir():
base = os.getenv("CAI_WORKSPACE_DIR", ".") # Default to current dir base
name = os.getenv("CAI_WORKSPACE")
if name:
return os.path.abspath(os.path.join(base, name))
return os.path.abspath(base) # Use base dir if no name
def get_common_container_path():
name = os.getenv("CAI_WORKSPACE")
if name:
return f"/workspace/workspaces/{name}"
return "/" # Default container path
# Get the new workspace directory using the common function
new_workspace_dir = get_common_workspace_dir()
# Create the directory if it doesn't exist on host
try: # Add try-except for robustness
os.makedirs(new_workspace_dir, exist_ok=True)
except OSError as e:
console.print(f"[red]Error creating host directory {new_workspace_dir}: {e}[/red]")
# Decide if this is fatal or just a warning
# If container is active, also create the directory in the container
active_container = os.getenv("CAI_ACTIVE_CONTAINER", "")
if active_container:
# Check if container is running
check_process = subprocess.run(
["docker", "inspect", "--format", "{{.State.Running}}", active_container],
capture_output=True,
text=True,
check=False
)
if check_process.returncode == 0 and "true" in check_process.stdout.lower():
# Get container workspace path using the common function
container_workspace_path = get_common_container_path()
try:
mkdir_cmd = ["docker", "exec", active_container, "mkdir", "-p", container_workspace_path]
mkdir_result = subprocess.run(
mkdir_cmd,
capture_output=True,
text=True,
check=False
)
if mkdir_result.returncode == 0:
console.print(
f"[dim]Created workspace directory in container: {container_workspace_path}[/dim]"
)
else:
console.print(
f"[yellow]Warning: Could not create workspace directory in container: {mkdir_result.stderr}[/yellow]"
)
except Exception as e:
console.print(
f"[yellow]Warning: Failed to setup workspace in container: {str(e)}[/yellow]"
)
# Use a different panel style to indicate success
console.print(
Panel(
f"Workspace changed to: [bold green]{workspace_name}[/bold green]\n"
f"New workspace directory: [bold]{new_workspace_dir}[/bold]",
title="Workspace Updated",
border_style="green"
)
)
return True
def _get_workspace_dir(self) -> str:
"""Get the host workspace directory using the common utility.
Returns:
The host workspace directory path.
"""
try:
# Use the centralized function from common.py
from cai.tools.common import _get_workspace_dir as get_common_workspace_dir
return get_common_workspace_dir()
except ImportError:
# Provide a basic fallback if import fails, mirroring common.py logic
# without 'cai_default'
base_dir = os.getenv("CAI_WORKSPACE_DIR")
workspace_name = os.getenv("CAI_WORKSPACE")
if base_dir and workspace_name:
# Basic validation
if not all(c.isalnum() or c in ['_', '-'] for c in workspace_name):
print(f"[yellow]Warning: Invalid CAI_WORKSPACE name '{workspace_name}' in fallback.[/yellow]")
# Fallback to base directory if name is invalid
return os.path.abspath(base_dir)
target_dir = os.path.join(base_dir, workspace_name)
return os.path.abspath(target_dir)
elif base_dir:
# If only base dir is set, use that
return os.path.abspath(base_dir)
else:
# Default to current working directory if nothing else is set
return os.getcwd()
def _list_workspace_contents(self, env_type: str, workspace_dir: str) -> None:
"""List the contents of the workspace.
Args:
env_type: The environment type (container or host)
workspace_dir: The workspace directory
"""
console.print("\n[bold]Workspace Contents:[/bold]")
if env_type == "container":
active_container = os.getenv("CAI_ACTIVE_CONTAINER", "")
# For containers, use the workspace path provided
# This should already be the correct path from handle_get
# First ensure the workspace directory exists in the container
try:
mkdir_cmd = ["docker", "exec", active_container, "mkdir", "-p", workspace_dir]
subprocess.run(
mkdir_cmd,
capture_output=True,
text=True,
check=False
)
# Now list the contents
result = subprocess.run(
["docker", "exec", active_container, "ls", "-la", workspace_dir],
capture_output=True,
text=True,
check=False
)
if result.returncode == 0:
console.print(result.stdout)
else:
console.print(f"[yellow]Error listing container files: {result.stderr}[/yellow]")
# Fallback to host
self._list_host_files(workspace_dir)
except Exception as e:
console.print(f"[yellow]Error accessing container: {str(e)}[/yellow]")
# Fallback to host
self._list_host_files(workspace_dir)
else:
# List files in host
self._list_host_files(workspace_dir)
def _list_host_files(self, workspace_dir: str) -> None:
"""List files in the host workspace.
Args:
workspace_dir: The workspace directory
"""
# Ensure the directory exists
os.makedirs(workspace_dir, exist_ok=True)
try:
result = subprocess.run(
["ls", "-la", workspace_dir],
capture_output=True,
text=True,
check=False
)
if result.returncode == 0:
console.print(result.stdout)
else:
console.print(f"[yellow]Error listing files: {result.stderr}[/yellow]")
except Exception as e:
console.print(f"[yellow]Error: {str(e)}[/yellow]")
def handle_ls_subcommand(self, args: Optional[List[str]] = None) -> bool:
"""Handle the ls subcommand.
Args:
args: Optional list of subcommand arguments
Returns:
True if the subcommand was handled successfully, False otherwise
"""
# Get workspace info using common functions
try:
from cai.tools.common import _get_workspace_dir as get_common_workspace_dir
from cai.tools.common import _get_container_workspace_path as get_common_container_path
except ImportError:
# Define basic fallbacks if import fails
def get_common_workspace_dir():
base = os.getenv("CAI_WORKSPACE_DIR", ".")
name = os.getenv("CAI_WORKSPACE")
if name: return os.path.abspath(os.path.join(base, name))
return os.path.abspath(base)
def get_common_container_path():
name = os.getenv("CAI_WORKSPACE")
if name: return f"/workspace/workspaces/{name}"
return "/"
host_workspace_dir = get_common_workspace_dir()
active_container = os.getenv("CAI_ACTIVE_CONTAINER", "")
# Execute command in the appropriate environment
if active_container:
# Use the container workspace path from common function
container_workspace_path = get_common_container_path()
# Determine the target path within the container
target_path_in_container = container_workspace_path
if args:
# Ensure args[0] is treated as relative to the workspace
target_path_in_container = os.path.join(container_workspace_path, args[0])
# Ensure the base workspace directory exists in the container
mkdir_cmd = ["docker", "exec", active_container, "mkdir", "-p", container_workspace_path]
subprocess.run(
mkdir_cmd,
capture_output=True,
text=True,
check=False
)
# Try in container
result = subprocess.run(
["docker", "exec", active_container, "ls", "-la", target_path_in_container], # Use target path
capture_output=True,
text=True,
check=False
)
if result.returncode == 0:
console.print(result.stdout)
return True
# If failed, try on host
console.print(f"[yellow]Failed to list files in container: {result.stderr}[/yellow]")
console.print("[yellow]Falling back to host system...[/yellow]")
# List on host
# Determine target path on host relative to host workspace dir
target_path_on_host = host_workspace_dir
if args:
# Ensure args[0] is treated as relative to the workspace
target_path_on_host = os.path.join(host_workspace_dir, args[0])
# Ensure the target directory exists on host before listing
# Use os.path.dirname if target is potentially a file path
dir_to_ensure = os.path.dirname(target_path_on_host) if '.' in os.path.basename(target_path_on_host) else target_path_on_host
try:
os.makedirs(dir_to_ensure, exist_ok=True)
except OSError as e:
console.print(f"[red]Error creating directory {dir_to_ensure} on host: {e}[/red]")
# Potentially return False or handle error appropriately
try:
result = subprocess.run(
["ls", "-la", target_path_on_host], # Use target path
capture_output=True,
text=True,
check=False
)
if result.returncode == 0:
console.print(result.stdout)
return True
else:
console.print(f"[red]Error listing files: {result.stderr}[/red]")
return False
except Exception as e:
console.print(f"[red]Error: {str(e)}[/red]")
return False
return True
def handle_exec_subcommand(self, args: Optional[List[str]] = None) -> bool:
"""Handle the exec subcommand.
Args:
args: Optional list of subcommand arguments
Returns:
True if the subcommand was handled successfully, False otherwise
"""
if not args:
console.print("[yellow]Please specify a command to execute.[/yellow]")
return False
command = " ".join(args)
# Get workspace info using common functions
try:
from cai.tools.common import _get_workspace_dir as get_common_workspace_dir
from cai.tools.common import _get_container_workspace_path as get_common_container_path
except ImportError:
# Define basic fallbacks if import fails
def get_common_workspace_dir():
base = os.getenv("CAI_WORKSPACE_DIR", ".")
name = os.getenv("CAI_WORKSPACE")
if name: return os.path.abspath(os.path.join(base, name))
return os.path.abspath(base)
def get_common_container_path():
name = os.getenv("CAI_WORKSPACE")
if name: return f"/workspace/workspaces/{name}"
return "/"
host_workspace_dir = get_common_workspace_dir()
active_container = os.getenv("CAI_ACTIVE_CONTAINER", "")
# Execute in container if active
if active_container:
try:
# Use the container workspace path from common function
container_workspace_path = get_common_container_path()
# First ensure the workspace directory exists in the container
mkdir_cmd = ["docker", "exec", active_container, "mkdir", "-p", container_workspace_path]
subprocess.run(
mkdir_cmd,
capture_output=True,
text=True,
check=False
)
# Execute the command in the container's workspace directory
result = subprocess.run(
["docker", "exec", "-w", container_workspace_path, active_container, "sh", "-c", command],
capture_output=True,
text=True,
check=False
)
console.print(f"[dim]$ {command}[/dim]")
if result.stdout:
console.print(result.stdout)
if result.stderr:
console.print(f"[yellow]{result.stderr}[/yellow]")
if result.returncode != 0:
console.print("[yellow]Command failed in container. Trying on host...[/yellow]")
return self._exec_on_host(command, host_workspace_dir) # Pass host_workspace_dir
return True
except Exception as e:
console.print(f"[yellow]Error executing in container: {str(e)}[/yellow]")
console.print("[yellow]Falling back to host execution...[/yellow]")
# Execute on host
return self._exec_on_host(command, host_workspace_dir) # Pass host_workspace_dir
def _exec_on_host(self, command: str, workspace_dir: str) -> bool:
"""Execute a command on the host.
Args:
command: The command to execute
workspace_dir: The workspace directory
Returns:
True if the command was executed successfully, False otherwise
"""
# Ensure the directory exists
os.makedirs(workspace_dir, exist_ok=True)
try:
result = subprocess.run(
command,
shell=True, # nosec B602
capture_output=True,
text=True,
check=False,
cwd=workspace_dir
)
console.print(f"[dim]$ {command}[/dim]")
if result.stdout:
console.print(result.stdout)
if result.stderr:
console.print(f"[yellow]{result.stderr}[/yellow]")
return result.returncode == 0
except Exception as e:
console.print(f"[red]Error executing command: {str(e)}[/red]")
return False
def handle_copy_subcommand(self, args: Optional[List[str]] = None) -> bool:
"""Handle the copy subcommand.
Args:
args: Optional list of subcommand arguments
Returns:
True if the subcommand was handled successfully, False otherwise
"""
if not args or len(args) < 2:
console.print("[yellow]Please specify source and destination for copy.[/yellow]")
console.print("Usage: /workspace copy <source> <destination>")
return False
active_container = os.getenv("CAI_ACTIVE_CONTAINER", "")
if not active_container:
console.print("[yellow]No active container. Copy only works with containers.[/yellow]")
return False
source = args[0]
destination = args[1]
# Check if copying from container to host or vice versa
if source.startswith("container:"):
# Copy from container to host
container_path = source[10:] # Remove "container:" prefix
host_path = destination
if not container_path.startswith("/"):
container_path = f"/workspace/{container_path}"
try:
result = subprocess.run(
["docker", "cp", f"{active_container}:{container_path}", host_path],
capture_output=True,
text=True,
check=False
)
if result.returncode == 0:
console.print(f"[green]Copied from container:{container_path} to {host_path}[/green]")
return True
else:
console.print(f"[red]Error copying from container: {result.stderr}[/red]")
return False
except Exception as e:
console.print(f"[red]Error: {str(e)}[/red]")
return False
elif destination.startswith("container:"):
# Copy from host to container
host_path = source
container_path = destination[10:] # Remove "container:" prefix
if not container_path.startswith("/"):
container_path = f"/workspace/{container_path}"
try:
result = subprocess.run(
["docker", "cp", host_path, f"{active_container}:{container_path}"],
capture_output=True,
text=True,
check=False
)
if result.returncode == 0:
console.print(f"[green]Copied from {host_path} to container:{container_path}[/green]")
return True
else:
console.print(f"[red]Error copying to container: {result.stderr}[/red]")
return False
except Exception as e:
console.print(f"[red]Error: {str(e)}[/red]")
return False
else:
# Ambiguous copy - show help
console.print("[yellow]Ambiguous copy direction. Please specify container: prefix.[/yellow]")
console.print("Examples:")
console.print(" /workspace copy file.txt container:file.txt # Host to container")
console.print(" /workspace copy container:file.txt file.txt # Container to host")
return False
# Register the commands
register_command(WorkspaceCommand())

View File

@ -18,7 +18,7 @@ toolbar_last_refresh = [datetime.datetime.now()]
toolbar_cache = { toolbar_cache = {
'html': "", 'html': "",
'last_update': datetime.datetime.now(), 'last_update': datetime.datetime.now(),
'refresh_interval': 60 # Refresh every 60 seconds 'refresh_interval': 5 # Refresh every 60 seconds
} }
# Cache for system information that rarely changes # Cache for system information that rarely changes
@ -57,7 +57,22 @@ def update_toolbar_in_background():
ip_address = sys_info['ip_address'] ip_address = sys_info['ip_address']
os_name = sys_info['os_name'] os_name = sys_info['os_name']
os_version = sys_info['os_version'] os_version = sys_info['os_version']
# Get the current workspace and base directory
workspace_name = os.getenv("CAI_WORKSPACE")
base_dir = os.getenv("CAI_WORKSPACE_DIR", "workspaces")
# Construct the workspace path
standard_path = os.path.join(base_dir, workspace_name) if workspace_name else ""
workspace_path = ""
if workspace_name:
if os.path.isdir(standard_path):
workspace_path = standard_path
elif os.path.isdir(workspace_name):
workspace_path = os.path.abspath(workspace_name)
else:
workspace_path = standard_path
# Get Ollama information # Get Ollama information
ollama_status = "unavailable" ollama_status = "unavailable"
try: try:

View File

@ -7,6 +7,9 @@ import os
import litellm import litellm
import tiktoken import tiktoken
import inspect import inspect
import hashlib
import re
import asyncio
from collections.abc import AsyncIterator, Iterable from collections.abc import AsyncIterator, Iterable
from dataclasses import dataclass, field from dataclasses import dataclass, field
@ -93,9 +96,40 @@ if TYPE_CHECKING:
# Suppress debug info from litellm # Suppress debug info from litellm
litellm.suppress_debug_info = True litellm.suppress_debug_info = True
if os.getenv('CAI_MODEL') == "o3-mini" or os.getenv('CAI_MODEL') == "gemini-1.5-pro":
litellm.drop_params = True
_USER_AGENT = f"Agents/Python {__version__}" _USER_AGENT = f"Agents/Python {__version__}"
_HEADERS = {"User-Agent": _USER_AGENT} _HEADERS = {"User-Agent": _USER_AGENT}
message_history = []
# Function to add a message to history if it's not a duplicate
def add_to_message_history(msg):
"""Add a message to history if it's not a duplicate."""
if not message_history:
message_history.append(msg)
return
is_duplicate = False
if msg.get("role") in ["system", "user"]:
is_duplicate = any(
existing.get("role") == msg.get("role") and
existing.get("content") == msg.get("content")
for existing in message_history
)
elif msg.get("role") == "assistant" and msg.get("tool_calls"):
is_duplicate = any(
existing.get("role") == "assistant" and
existing.get("tool_calls") and
existing["tool_calls"][0].get("id") == msg["tool_calls"][0].get("id")
for existing in message_history
)
if not is_duplicate:
message_history.append(msg)
@dataclass @dataclass
class _StreamingState: class _StreamingState:
@ -235,7 +269,32 @@ class OpenAIChatCompletionsModel(Model):
"role": "system", "role": "system",
}, },
) )
# --- Add to message_history: user, system, and assistant tool call messages ---
# Add system prompt to message_history
if system_instructions:
sys_msg = {
"role": "system",
"content": system_instructions
}
add_to_message_history(sys_msg)
# Add user prompt(s) to message_history
if isinstance(input, str):
user_msg = {
"role": "user",
"content": input
}
add_to_message_history(user_msg)
elif isinstance(input, list):
for item in input:
# Try to extract user messages
if isinstance(item, dict):
if item.get("role") == "user":
user_msg = {
"role": "user",
"content": item.get("content", "")
}
add_to_message_history(user_msg)
# Get token count estimate before API call for consistent counting # Get token count estimate before API call for consistent counting
estimated_input_tokens, _ = count_tokens_with_tiktoken(converted_messages) estimated_input_tokens, _ = count_tokens_with_tiktoken(converted_messages)
@ -335,6 +394,36 @@ class OpenAIChatCompletionsModel(Model):
tool_output=None, # Don't pass tool output here, we're using direct display tool_output=None, # Don't pass tool output here, we're using direct display
) )
# --- Add assistant tool call to message_history if present ---
# If the response contains tool_calls, add them to message_history as assistant messages
assistant_msg = response.choices[0].message
if hasattr(assistant_msg, "tool_calls") and assistant_msg.tool_calls:
for tool_call in assistant_msg.tool_calls:
# Compose a message for the tool call
tool_call_msg = {
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": tool_call.id,
"type": tool_call.type,
"function": {
"name": tool_call.function.name,
"arguments": tool_call.function.arguments
}
}
]
}
add_to_message_history(tool_call_msg)
# If the assistant message is just text, add it as well
elif hasattr(assistant_msg, "content") and assistant_msg.content:
asst_msg = {
"role": "assistant",
"content": assistant_msg.content
}
add_to_message_history(asst_msg)
usage = ( usage = (
Usage( Usage(
requests=1, requests=1,
@ -404,7 +493,29 @@ class OpenAIChatCompletionsModel(Model):
"role": "system", "role": "system",
}, },
) )
# --- Add to message_history: user, system prompts ---
if system_instructions:
sys_msg = {
"role": "system",
"content": system_instructions
}
add_to_message_history(sys_msg)
if isinstance(input, str):
user_msg = {
"role": "user",
"content": input
}
add_to_message_history(user_msg)
elif isinstance(input, list):
for item in input:
if isinstance(item, dict):
if item.get("role") == "user":
user_msg = {
"role": "user",
"content": item.get("content", "")
}
add_to_message_history(user_msg)
# Get token count estimate before API call for consistent counting # Get token count estimate before API call for consistent counting
estimated_input_tokens, _ = count_tokens_with_tiktoken(converted_messages) estimated_input_tokens, _ = count_tokens_with_tiktoken(converted_messages)
@ -429,7 +540,25 @@ class OpenAIChatCompletionsModel(Model):
# Initialize a streaming text accumulator for rich display # Initialize a streaming text accumulator for rich display
streaming_text_buffer = "" streaming_text_buffer = ""
# For tool call streaming, accumulate tool_calls to add to message_history at the end
streamed_tool_calls = []
# Ollama specific: accumulate full content to check for function calls at the end
# Some Ollama models output the function call as JSON in the text content
ollama_full_content = ""
is_ollama = False
model_str = str(self.model).lower()
is_ollama = self.is_ollama or "ollama" in model_str or ":" in model_str or "qwen" in model_str
# Add visual separation before agent output
if streaming_context and should_show_rich_stream:
# If we're using rich context, we'll add separation through that
pass
else:
# Print clear visual separator
print("\n")
async for chunk in stream: async for chunk in stream:
if not state.started: if not state.started:
state.started = True state.started = True
@ -453,6 +582,10 @@ class OpenAIChatCompletionsModel(Model):
choices = [{"delta": chunk.delta}] choices = [{"delta": chunk.delta}]
elif isinstance(chunk, dict) and 'choices' in chunk: elif isinstance(chunk, dict) and 'choices' in chunk:
choices = chunk['choices'] choices = chunk['choices']
# Special handling for Qwen/Ollama chunks
elif isinstance(chunk, dict) and ('content' in chunk or 'function_call' in chunk):
# Qwen direct delta format - convert to standard
choices = [{"delta": chunk}]
else: else:
# Skip chunks that don't contain choice data # Skip chunks that don't contain choice data
continue continue
@ -478,6 +611,10 @@ class OpenAIChatCompletionsModel(Model):
content = delta['content'] content = delta['content']
if content: if content:
# For Ollama, we need to accumulate the full content to check for function calls
if is_ollama:
ollama_full_content += content
# Add to the streaming text buffer # Add to the streaming text buffer
streaming_text_buffer += content streaming_text_buffer += content
@ -589,11 +726,7 @@ class OpenAIChatCompletionsModel(Model):
# Handle tool calls # Handle tool calls
# Because we don't know the name of the function until the end of the stream, we'll # 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 # save everything and yield events at the end
tool_calls = None tool_calls = self._detect_and_format_function_calls(delta)
if hasattr(delta, 'tool_calls') and delta.tool_calls:
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: if tool_calls:
for tc_delta in tool_calls: for tc_delta in tool_calls:
@ -636,9 +769,146 @@ class OpenAIChatCompletionsModel(Model):
call_id = tc_delta.id or "" call_id = tc_delta.id or ""
elif isinstance(tc_delta, dict) and 'id' in tc_delta: elif isinstance(tc_delta, dict) and 'id' in tc_delta:
call_id = tc_delta.get('id', "") or "" call_id = tc_delta.get('id', "") or ""
else:
# For Qwen models, generate a predictable ID if none is provided
if state.function_calls[tc_index].name:
# Generate a stable ID from the function name and arguments
call_id = f"call_{hashlib.md5(state.function_calls[tc_index].name.encode()).hexdigest()[:8]}"
state.function_calls[tc_index].call_id += call_id state.function_calls[tc_index].call_id += call_id
# --- Accumulate tool call for message_history ---
# Only add if not already present (avoid duplicates in streaming)
tool_call_msg = {
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": state.function_calls[tc_index].call_id,
"type": "function",
"function": {
"name": state.function_calls[tc_index].name,
"arguments": state.function_calls[tc_index].arguments
}
}
]
}
# Only add if not already in streamed_tool_calls
if tool_call_msg not in streamed_tool_calls:
streamed_tool_calls.append(tool_call_msg)
add_to_message_history(tool_call_msg)
# Special handling for Ollama - check if accumulated text contains a valid function call
if is_ollama and ollama_full_content and len(state.function_calls) == 0:
# Look for JSON object that might be a function call
try:
# Try to extract a JSON object from the content
json_start = ollama_full_content.find('{')
json_end = ollama_full_content.rfind('}') + 1
if json_start >= 0 and json_end > json_start:
json_str = ollama_full_content[json_start:json_end]
# Try to parse the JSON
parsed = json.loads(json_str)
# Check if it looks like a function call
if ('name' in parsed and 'arguments' in parsed):
logger.debug(f"Found valid function call in Ollama output: {json_str}")
# Create a tool call ID
tool_call_id = f"call_{hashlib.md5((parsed['name'] + str(time.time())).encode()).hexdigest()[:8]}"
# Ensure arguments is a valid JSON string
arguments_str = ""
if isinstance(parsed['arguments'], dict):
# Remove 'ctf' field if it exists
if 'ctf' in parsed['arguments']:
del parsed['arguments']['ctf']
arguments_str = json.dumps(parsed['arguments'])
elif isinstance(parsed['arguments'], str):
# If it's already a string, check if it's valid JSON
try:
# Try parsing to validate and remove 'ctf' if present
args_dict = json.loads(parsed['arguments'])
if isinstance(args_dict, dict) and 'ctf' in args_dict:
del args_dict['ctf']
arguments_str = json.dumps(args_dict)
except:
# If not valid JSON, encode it as a JSON string
arguments_str = json.dumps(parsed['arguments'])
else:
# For any other type, convert to string and then JSON
arguments_str = json.dumps(str(parsed['arguments']))
# Add it to our function_calls state
state.function_calls[0] = ResponseFunctionToolCall(
id=FAKE_RESPONSES_ID,
arguments=arguments_str,
name=parsed['name'],
type="function_call",
call_id=tool_call_id,
)
# Display the tool call in CLI
from cai.util import cli_print_agent_messages
try:
# Create a message-like object to display the function call
tool_msg = type('ToolCallWrapper', (), {
'content': None,
'tool_calls': [
type('ToolCallDetail', (), {
'function': type('FunctionDetail', (), {
'name': parsed['name'],
'arguments': arguments_str
}),
'id': tool_call_id,
'type': 'function'
})
]
})
# Print the tool call using the CLI utility
cli_print_agent_messages(
agent_name=getattr(self, 'agent_name', 'Agent'),
message=tool_msg,
counter=getattr(self, 'interaction_counter', 0),
model=str(self.model),
debug=False,
interaction_input_tokens=estimated_input_tokens,
interaction_output_tokens=estimated_output_tokens,
interaction_reasoning_tokens=0, # Not available for Ollama
total_input_tokens=getattr(self, 'total_input_tokens', 0) + estimated_input_tokens,
total_output_tokens=getattr(self, 'total_output_tokens', 0) + estimated_output_tokens,
total_reasoning_tokens=getattr(self, 'total_reasoning_tokens', 0),
interaction_cost=None,
total_cost=None,
tool_output=None # Will be shown once the tool is executed
)
except Exception as e:
logger.error(f"Error displaying tool call in CLI: {e}")
# Add to message history
tool_call_msg = {
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": tool_call_id,
"type": "function",
"function": {
"name": parsed['name'],
"arguments": arguments_str
}
}
]
}
streamed_tool_calls.append(tool_call_msg)
add_to_message_history(tool_call_msg)
logger.debug(f"Added function call: {parsed['name']} with args: {arguments_str}")
except Exception as e:
pass
function_call_starting_index = 0 function_call_starting_index = 0
if state.text_content_index_and_output: if state.text_content_index_and_output:
function_call_starting_index += 1 function_call_starting_index += 1
@ -817,6 +1087,9 @@ class OpenAIChatCompletionsModel(Model):
direct_stats["total_cost"] = float(total_cost) direct_stats["total_cost"] = float(total_cost)
# Use the direct copy with guaranteed float costs # Use the direct copy with guaranteed float costs
finish_agent_streaming(streaming_context, direct_stats) finish_agent_streaming(streaming_context, direct_stats)
# Add visual separation after agent output completes
print("\n")
# If we're not using rich streaming and not suppressing output, use old method # If we're not using rich streaming and not suppressing output, use old method
elif not self.suppress_final_output and final_response.output and any(isinstance(item, ResponseOutputMessage) for item in final_response.output): elif not self.suppress_final_output and final_response.output and any(isinstance(item, ResponseOutputMessage) for item in final_response.output):
# Find the assistant message to print # Find the assistant message to print
@ -837,8 +1110,22 @@ class OpenAIChatCompletionsModel(Model):
interaction_cost=interaction_cost, interaction_cost=interaction_cost,
total_cost=total_cost, total_cost=total_cost,
) )
# Add visual separation after message
print("\n")
break break
# --- Add assistant tool call(s) to message_history at the end of streaming ---
for tool_call_msg in streamed_tool_calls:
add_to_message_history(tool_call_msg)
# If there was only text output, add that as an assistant message
if (not streamed_tool_calls) and state.text_content_index_and_output and state.text_content_index_and_output[1].text:
asst_msg = {
"role": "assistant",
"content": state.text_content_index_and_output[1].text
}
add_to_message_history(asst_msg)
if tracing.include_data(): if tracing.include_data():
span_generation.span_data.output = [final_response.model_dump()] span_generation.span_data.output = [final_response.model_dump()]
@ -964,36 +1251,69 @@ class OpenAIChatCompletionsModel(Model):
"extra_headers": _HEADERS, "extra_headers": _HEADERS,
} }
# Determine provider based on model string
model_str = str(self.model).lower()
# Error encountered: Error code: 400 - {'error': {'code': 'invalid_request_error',
# 'message': "'tool_choice' is only allowed when 'tools' are specified", # Provider-specific adjustments
# 'type': 'invalid_request_error', 'param': None}} if "/" in model_str:
# # Handle provider/model format
# Only remove tool_choice if model starts with "gpt" and has no tools provider = model_str.split("/")[0]
if self.model.startswith("gpt") and not converted_tools:
kwargs.pop("tool_choice", None)
# TODO: review this. Remove tool_choice for Anthropic/Claude models when no tools are provided # Apply provider-specific configurations
if ("claude" in str(self.model).lower() or "anthropic" in str(self.model).lower()) and not converted_tools: if provider == "deepseek":
kwargs.pop("tool_choice", None) litellm.drop_params = True
kwargs.pop("parallel_tool_calls", None)
# Model adjustments # Remove tool_choice if no tools are specified
if any(x in self.model for x in ["claude"]): if not converted_tools:
litellm.drop_params = True kwargs.pop("tool_choice", None)
# BadRequestError encountered: litellm.BadRequestError: AnthropicException - elif provider == "claude":
# b'{"type":"error","error": litellm.drop_params = True
# {"type":"invalid_request_error","message":"store: Extra inputs are not permitted"}}' kwargs.pop("store", None)
# # Remove tool_choice if no tools are specified
kwargs.pop("store", None) if not converted_tools:
kwargs.pop("tool_choice", None)
# Filter out NotGiven values to avoid JSON serialization issues elif provider == "gemini":
filtered_kwargs = {} kwargs.pop("parallel_tool_calls", None)
for key, value in kwargs.items(): # Add any specific gemini settings if needed
if value is not NOT_GIVEN: else:
filtered_kwargs[key] = value # Handle models without provider prefix
kwargs = filtered_kwargs if "claude" in model_str:
litellm.drop_params = True
# Remove store parameter which isn't supported by Anthropic
kwargs.pop("store", None)
# Remove tool_choice if no tools are specified
if not converted_tools:
kwargs.pop("tool_choice", None)
elif "gemini" in model_str:
kwargs.pop("parallel_tool_calls", None)
elif "qwen" in model_str or ":" in model_str:
# Handle Ollama-served models with custom formats (e.g., qwen2.5:14b)
# These typically need the Ollama provider
litellm.drop_params = True
kwargs.pop("parallel_tool_calls", None)
# These models may not support certain parameters
if not converted_tools:
kwargs.pop("tool_choice", None)
# Don't add custom_llm_provider here to avoid duplication with Ollama provider
if self.is_ollama:
# Clean kwargs for ollama to avoid parameter conflicts
for param in ["custom_llm_provider"]:
kwargs.pop(param, None)
elif any(x in model_str for x in ["o1", "o3", "o4"]):
# Handle OpenAI reasoning models (o1, o3, o4)
kwargs.pop("parallel_tool_calls", None)
# Add reasoning effort if provided
if hasattr(model_settings, "reasoning_effort"):
kwargs["reasoning_effort"] = model_settings.reasoning_effort
# 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
try: try:
if self.is_ollama: if self.is_ollama:
return await self._fetch_response_litellm_ollama(kwargs, model_settings, tool_choice, stream, parallel_tool_calls) return await self._fetch_response_litellm_ollama(kwargs, model_settings, tool_choice, stream, parallel_tool_calls)
@ -1003,21 +1323,71 @@ class OpenAIChatCompletionsModel(Model):
except litellm.exceptions.BadRequestError as e: except litellm.exceptions.BadRequestError as e:
# print(color("BadRequestError encountered: " + str(e), fg="yellow")) # print(color("BadRequestError encountered: " + str(e), fg="yellow"))
if "LLM Provider NOT provided" in str(e): if "LLM Provider NOT provided" in str(e):
# Create a copy of params to avoid overwriting the original model_str = str(self.model).lower()
# ones provider = None
try: is_qwen = "qwen" in model_str or ":" in model_str
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 # Special handling for Qwen models
# if is_qwen:
# CTRL-C handler for ollama models try:
# # Use the specialized Qwen approach first
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) return await self._fetch_response_litellm_ollama(kwargs, model_settings, tool_choice, stream, parallel_tool_calls)
except Exception as qwen_e:
print(qwen_e)
# If that fails, try our direct OpenAI approach
qwen_params = kwargs.copy()
qwen_params["api_base"] = get_ollama_api_base()
qwen_params["custom_llm_provider"] = "openai" # Use openai provider
# Make sure tools are passed
if "tools" in kwargs and kwargs["tools"]:
qwen_params["tools"] = kwargs["tools"]
if "tool_choice" in kwargs and kwargs["tool_choice"] is not NOT_GIVEN:
qwen_params["tool_choice"] = kwargs["tool_choice"]
try:
if stream:
# Streaming case
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,
)
stream_obj = await litellm.acompletion(**qwen_params)
return response, stream_obj
else:
# Non-streaming case
ret = litellm.completion(**qwen_params)
return ret
except Exception as direct_e:
# All approaches failed, log and raise the original error
print(f"All Qwen approaches failed. Original error: {str(e)}, Direct error: {str(direct_e)}")
raise e
# Try to detect provider from model string
if "/" in model_str:
provider = model_str.split("/")[0]
if provider:
# Add provider-specific settings based on detected provider
provider_kwargs = kwargs.copy()
if provider == "deepseek":
provider_kwargs["custom_llm_provider"] = "deepseek"
elif provider == "claude" or "claude" in model_str:
provider_kwargs["custom_llm_provider"] = "anthropic"
elif provider == "gemini":
provider_kwargs["custom_llm_provider"] = "gemini"
else: else:
raise e # For unknown providers, try ollama as fallback
return await self._fetch_response_litellm_ollama(kwargs, model_settings, tool_choice, stream, parallel_tool_calls)
elif ("An assistant message with 'tool_calls'" in str(e) or 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 "`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)}") print(f"Error: {str(e)}")
@ -1121,77 +1491,165 @@ class OpenAIChatCompletionsModel(Model):
# Standard OpenAI handling for non-streaming # Standard OpenAI handling for non-streaming
ret = litellm.completion(**kwargs) ret = litellm.completion(**kwargs)
return ret return ret
async def _fetch_response_litellm_ollama( async def _fetch_response_litellm_ollama(
self, self,
kwargs: dict, kwargs: dict,
model_settings: ModelSettings, model_settings: ModelSettings,
tool_choice: ChatCompletionToolChoiceOptionParam | NotGiven, tool_choice: ChatCompletionToolChoiceOptionParam | NotGiven,
stream: bool, stream: bool,
parallel_tool_calls: bool parallel_tool_calls: bool,
provider="ollama"
) -> ChatCompletion | tuple[Response, AsyncStream[ChatCompletionChunk]]: ) -> ChatCompletion | tuple[Response, AsyncStream[ChatCompletionChunk]]:
# Filter out parameters not supported by Ollama # Extract only supported parameters for Ollama
ollama_supported_params = { ollama_supported_params = {
"model": kwargs["model"], "model": kwargs.get("model", ""),
"messages": kwargs["messages"], "messages": kwargs.get("messages", []),
"temperature": kwargs["temperature"] if kwargs["temperature"] is not NOT_GIVEN else None, "stream": kwargs.get("stream", False)
"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 # Add optional parameters if they exist and are not NOT_GIVEN
if ollama_supported_params["messages"] and ollama_supported_params["messages"][0].get("role") == "system": for param in ["temperature", "top_p", "max_tokens"]:
# Extract the system message if param in kwargs and kwargs[param] is not NOT_GIVEN:
system_content = ollama_supported_params["messages"][0].get("content", "") ollama_supported_params[param] = kwargs[param]
# Remove it from the messages
ollama_supported_params["messages"] = ollama_supported_params["messages"][1:] # Add extra headers if available
# If there are user messages, prepend system to first user if "extra_headers" in kwargs:
if ollama_supported_params["messages"] and ollama_supported_params["messages"][0].get("role") == "user": ollama_supported_params["extra_headers"] = kwargs["extra_headers"]
# Prepend the system instruction to the first user message, with a separator
user_content = ollama_supported_params["messages"][0].get("content", "") # Add tools and tool_choice for compatibility with Qwen
if isinstance(user_content, str): if "tools" in kwargs and kwargs.get("tools") and kwargs.get("tools") is not NOT_GIVEN:
ollama_supported_params["messages"][0]["content"] = f"System: {system_content}\n\nUser: {user_content}" ollama_supported_params["tools"] = kwargs.get("tools")
if "tool_choice" in kwargs and kwargs.get("tool_choice") is not NOT_GIVEN:
ollama_supported_params["tool_choice"] = kwargs.get("tool_choice")
# Remove None values # Remove None values
ollama_kwargs = {k: v for k, v in ollama_supported_params.items() if v is not None} ollama_kwargs = {k: v for k, v in ollama_supported_params.items() if v is not None}
# Check if this is a Qwen model
model_str = str(self.model).lower()
is_qwen = "qwen" in model_str
api_base = get_ollama_api_base()
if "ollama" in provider:
api_base = api_base.rstrip('/v1')
# Create response object for streaming
if stream: if stream:
# For streaming with Ollama, we need to create a Response object first
response = Response( response = Response(
id=FAKE_RESPONSES_ID, id=FAKE_RESPONSES_ID,
created_at=time.time(), created_at=time.time(),
model=self.model, model=self.model,
object="response", object="response",
output=[], output=[],
tool_choice="auto" if tool_choice is None or tool_choice == NOT_GIVEN else cast(Literal["auto", "required", "none"], tool_choice), 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, top_p=model_settings.top_p,
temperature=model_settings.temperature, temperature=model_settings.temperature,
tools=[], tools=[],
parallel_tool_calls=parallel_tool_calls or False, parallel_tool_calls=parallel_tool_calls or False,
) )
# Get the streaming object # Get streaming response
stream_obj = await litellm.acompletion( stream_obj = await litellm.acompletion(
**ollama_kwargs, **ollama_kwargs,
api_base=get_ollama_api_base().rstrip('/v1'), api_base=api_base,
custom_llm_provider="ollama" custom_llm_provider=provider,
) )
return response, stream_obj return response, stream_obj
else: else:
# Non-streaming mode
ret = litellm.completion(
# Get completion response
return litellm.completion(
**ollama_kwargs, **ollama_kwargs,
api_base=get_ollama_api_base().rstrip('/v1'), api_base=api_base,
custom_llm_provider="ollama" custom_llm_provider=provider,
) )
return ret
def _get_client(self) -> AsyncOpenAI: def _get_client(self) -> AsyncOpenAI:
if self._client is None: if self._client is None:
self._client = AsyncOpenAI() self._client = AsyncOpenAI()
return self._client return self._client
# Helper function to detect and format function calls from various models
def _detect_and_format_function_calls(self, delta):
"""
Helper to detect function calls in different formats and normalize them.
Handles Qwen specifics where function calls may be formatted differently.
Returns: List of normalized tool calls or None
"""
# Standard OpenAI-style tool_calls format
if hasattr(delta, 'tool_calls') and delta.tool_calls:
return delta.tool_calls
elif isinstance(delta, dict) and 'tool_calls' in delta and delta['tool_calls']:
return delta['tool_calls']
# Qwen/Ollama function_call format
if isinstance(delta, dict) and 'function_call' in delta:
function_call = delta['function_call']
return [{
'index': 0,
'id': f"call_{time.time_ns()}", # Generate a unique ID
'type': 'function',
'function': {
'name': function_call.get('name', ''),
'arguments': function_call.get('arguments', '')
}
}]
if isinstance(delta, dict) and 'content' in delta:
content = delta['content']
# Try to detect if the content is a JSON string with function call format
try:
if isinstance(content, str) and '{' in content and '}' in content:
# Try to extract JSON from the content (it might be embedded in text)
json_start = content.find('{')
json_end = content.rfind('}') + 1
if json_start >= 0 and json_end > json_start:
json_str = content[json_start:json_end]
parsed = json.loads(json_str)
if 'name' in parsed and 'arguments' in parsed:
# This looks like a function call in JSON format
return [{
'index': 0,
'id': f"call_{time.time_ns()}", # Generate a unique ID
'type': 'function',
'function': {
'name': parsed['name'],
'arguments': json.dumps(parsed['arguments']) if isinstance(parsed['arguments'], dict) else parsed['arguments']
}
}]
except Exception:
# If JSON parsing fails, just continue with normal processing
pass
# Anthropic-style tool_use format
if hasattr(delta, 'tool_use') and delta.tool_use:
tool_use = delta.tool_use
return [{
'index': 0,
'id': tool_use.get('id', f"tool_{time.time_ns()}"),
'type': 'function',
'function': {
'name': tool_use.get('name', ''),
'arguments': tool_use.get('input', '{}')
}
}]
elif isinstance(delta, dict) and 'tool_use' in delta and delta['tool_use']:
tool_use = delta['tool_use']
return [{
'index': 0,
'id': tool_use.get('id', f"tool_{time.time_ns()}"),
'type': 'function',
'function': {
'name': tool_use.get('name', ''),
'arguments': tool_use.get('input', '{}')
}
}]
return None
class _Converter: class _Converter:
@classmethod @classmethod

View File

@ -778,6 +778,7 @@ class Runner:
run_config: RunConfig, run_config: RunConfig,
tool_use_tracker: AgentToolUseTracker, tool_use_tracker: AgentToolUseTracker,
) -> SingleStepResult: ) -> SingleStepResult:
processed_response = RunImpl.process_model_response( processed_response = RunImpl.process_model_response(
agent=agent, agent=agent,
all_tools=all_tools, all_tools=all_tools,
@ -785,6 +786,41 @@ class Runner:
output_schema=output_schema, output_schema=output_schema,
handoffs=handoffs, handoffs=handoffs,
) )
# Log tools used with robust type checking
if hasattr(processed_response, 'tools_used') and processed_response.tools_used:
for i, tool_call in enumerate(processed_response.tools_used):
try:
# Safely extract tool name with multiple fallbacks
tool_name = "Unknown"
try:
if hasattr(tool_call, 'tool'):
if isinstance(tool_call.tool, str):
tool_name = tool_call.tool
elif hasattr(tool_call.tool, 'name'):
tool_name = tool_call.tool.name
else:
tool_name = str(tool_call.tool)
except Exception:
pass
# Safely extract call_id
call_id = "Unknown"
try:
if hasattr(tool_call, 'call_id'):
call_id = str(tool_call.call_id)
except Exception:
pass
# Safely extract parsed_args
parsed_args = "Unknown"
try:
if hasattr(tool_call, 'parsed_args'):
parsed_args = str(tool_call.parsed_args)
except Exception:
pass
except Exception:
pass
tool_use_tracker.add_tool_use(agent, processed_response.tools_used) tool_use_tracker.add_tool_use(agent, processed_response.tools_used)

View File

@ -24,14 +24,70 @@ except ImportError:
# Global dictionary to store active sessions # Global dictionary to store active sessions
ACTIVE_SESSIONS = {} ACTIVE_SESSIONS = {}
def _get_workspace_dir() -> str:
"""Determines the target workspace directory based on env vars for host."""
base_dir_env = os.getenv("CAI_WORKSPACE_DIR")
workspace_name = os.getenv("CAI_WORKSPACE")
# Determine the base directory
if base_dir_env:
base_dir = os.path.abspath(base_dir_env)
else: # Default base directory is 'workspaces'
if workspace_name:
base_dir = os.path.join(os.getcwd(), "workspaces")
else: # If no workspace name is set, the workspace IS the CWD.
return os.getcwd()
# If a workspace name is provided, append it to the base directory
if workspace_name:
if not all(c.isalnum() or c in ['_', '-'] for c in workspace_name):
print(color(f"Invalid CAI_WORKSPACE name '{workspace_name}'. "
f"Using directory '{base_dir}' instead.", fg="yellow"))
target_dir = base_dir
else:
target_dir = os.path.join(base_dir, workspace_name)
else:
target_dir = base_dir
# Ensure the final target directory exists on the host
try:
abs_target_dir = os.path.abspath(target_dir)
os.makedirs(abs_target_dir, exist_ok=True)
return abs_target_dir
except OSError as e:
print(color(f"Error creating/accessing host workspace directory '{abs_target_dir}': {e}",
fg="red"))
print(color(f"Falling back to current directory: {os.getcwd()}", fg="yellow"))
return os.getcwd()
def _get_container_workspace_path() -> str:
"""Determines the target workspace path inside the container."""
workspace_name = os.getenv("CAI_WORKSPACE")
if workspace_name:
if not all(c.isalnum() or c in ['_', '-'] for c in workspace_name):
print(color(f"Invalid CAI_WORKSPACE name '{workspace_name}' for container. "
f"Using '/workspace'.", fg="yellow"))
return "/"
# Standard path inside CAI containers
return f"/workspace/workspaces/{workspace_name}"
else:
return "/"
class ShellSession: # pylint: disable=too-many-instance-attributes class ShellSession: # pylint: disable=too-many-instance-attributes
"""Class to manage interactive shell sessions""" """Class to manage interactive shell sessions"""
def __init__(self, command, session_id=None, ctf=None): def __init__(self, command, session_id=None, ctf=None, workspace_dir=None, container_id=None): # noqa E501
self.session_id = session_id or str(uuid.uuid4())[:8] self.session_id = session_id or str(uuid.uuid4())[:8]
self.command = command self.command = command
self.ctf = ctf self.ctf = ctf
self.container_id = container_id
# Determine workspace based on context (container, ctf or local host)
if self.container_id:
self.workspace_dir = _get_container_workspace_path()
elif self.ctf:
self.workspace_dir = workspace_dir or _get_workspace_dir()
else:
self.workspace_dir = _get_workspace_dir()
self.process = None self.process = None
self.master = None self.master = None
self.slave = None self.slave = None
@ -40,9 +96,40 @@ class ShellSession: # pylint: disable=too-many-instance-attributes
self.last_activity = time.time() self.last_activity = time.time()
def start(self): def start(self):
"""Start the shell session""" """Start the shell session in the appropriate environment."""
start_message_cmd = self.command
# --- Start in Container ---
if self.container_id:
try:
self.master, self.slave = pty.openpty()
docker_cmd_list = [
"docker", "exec", "-i",
"-w", self.workspace_dir,
self.container_id,
"sh", "-c", # Use shell to handle complex commands if needed
self.command # The actual command to run
]
self.process = subprocess.Popen(
docker_cmd_list,
stdin=self.slave,
stdout=self.slave,
stderr=self.slave,
preexec_fn=os.setsid,
universal_newlines=True
)
self.is_running = True
self.output_buffer.append(
f"[Session {self.session_id}] Started in container {self.container_id[:12]}: "
f"{start_message_cmd} in {self.workspace_dir}")
threading.Thread(target=self._read_output, daemon=True).start()
except Exception as e:
self.output_buffer.append(f"Error starting container session: {str(e)}")
self.is_running = False
return
# --- Start in CTF ---
if self.ctf: if self.ctf:
# For CTF environments
self.is_running = True self.is_running = True
self.output_buffer.append( self.output_buffer.append(
f"[Session { f"[Session {
@ -52,81 +139,100 @@ class ShellSession: # pylint: disable=too-many-instance-attributes
output = self.ctf.get_shell(self.command) output = self.ctf.get_shell(self.command)
self.output_buffer.append(output) self.output_buffer.append(output)
except Exception as e: # pylint: disable=broad-except except Exception as e: # pylint: disable=broad-except
self.output_buffer.append(f"Error: {str(e)}") self.output_buffer.append(f"Error executing CTF command: {str(e)}")
self.is_running = False self.is_running = False
return return
# For local environment # --- Start Locally (Host) ---
try: try:
# Create a pseudo-terminal
self.master, self.slave = pty.openpty() self.master, self.slave = pty.openpty()
# Start the process
self.process = subprocess.Popen( # pylint: disable=subprocess-popen-preexec-fn, consider-using-with # noqa: E501 self.process = subprocess.Popen( # pylint: disable=subprocess-popen-preexec-fn, consider-using-with # noqa: E501
self.command, self.command,
shell=True, # nosec B602 shell=True, # nosec B602
stdin=self.slave, stdin=self.slave,
stdout=self.slave, stdout=self.slave,
stderr=self.slave, stderr=self.slave,
preexec_fn=os.setsid, # Create a new process group cwd=self.workspace_dir,
preexec_fn=os.setsid,
universal_newlines=True universal_newlines=True
) )
self.is_running = True self.is_running = True
self.output_buffer.append( self.output_buffer.append(
f"[Session { f"[Session {
self.session_id}] Started: { self.session_id}] Started: {
self.command}") self.command}")
# Start a thread to read output # Start a thread to read output
threading.Thread(target=self._read_output, daemon=True).start() threading.Thread(target=self._read_output, daemon=True).start()
except Exception as e: # pylint: disable=broad-except except Exception as e: # pylint: disable=broad-except
self.output_buffer.append(f"Error starting session: {str(e)}") self.output_buffer.append(f"Error starting local session: {str(e)}")
self.is_running = False self.is_running = False
def _read_output(self): def _read_output(self):
"""Read output from the process""" """Read output from the process"""
try: try:
while self.is_running: while self.is_running and self.master is not None:
try: try:
# Check if process has exited before reading
if self.process and self.process.poll() is not None:
self.is_running = False
break
# Read the output
output = os.read(self.master, 1024).decode() output = os.read(self.master, 1024).decode()
if output: if output:
self.output_buffer.append(output) self.output_buffer.append(output)
self.last_activity = time.time() self.last_activity = time.time()
except OSError: else:
# No data available or terminal closed
time.sleep(0.1)
if not self.is_process_running():
self.is_running = False self.is_running = False
break break
except Exception as e: # pylint: disable=broad-except except Exception as e:
self.output_buffer.append(f"Error reading output: {str(e)}") self.output_buffer.append(f"Error reading output buffer: {str(read_err)}")
self.is_running = False
break
# Add a small sleep to prevent busy-waiting if no output
if is_process_running(self):
time.sleep(0.05)
except Exception as e:
self.output_buffer.append(f"Error in read_output loop: {str(e)}")
self.is_running = False self.is_running = False
def is_process_running(self): def is_process_running(self):
"""Check if the process is still running""" """Check if the process is still running"""
# For CTF or container
if self.container_id or self.ctf:
return self.is_running
# For local host
if not self.process: if not self.process:
return False return False
return self.process.poll() is None return self.process.poll() is None
def send_input(self, input_data): def send_input(self, input_data):
"""Send input to the process""" """Send input to the process (local or container)"""
if not self.is_running: if not self.is_running: # For CTF or container
return "Session is not running" if self.process and self.process.poll() is None:
self.is_running = True
else: # For local host
return "Session is not running"
try: try:
# --- Send to CTF ---
if self.ctf: if self.ctf:
# For CTF environments
output = self.ctf.get_shell(input_data) output = self.ctf.get_shell(input_data)
self.output_buffer.append(output) self.output_buffer.append(output)
return "Input sent to CTF session" return "Input sent to CTF session"
# For local environment # --- Send to Local or Container PTY ---
input_data = input_data.rstrip() + "\n" if self.master is not None:
os.write(self.master, input_data.encode()) input_data_bytes = (input_data.rstrip() + "\n").encode()
self.last_activity = time.time() bytes_written = os.write(self.master, input_data_bytes)
return "Input sent to session" if bytes_written != len(input_data_bytes):
self.output_buffer.append(f"[Session {self.session_id}] Warning: Partial input write.")
self.last_activity = time.time()
return "Input sent to session"
else:
return "Session PTY not available for input"
except Exception as e: # pylint: disable=broad-except except Exception as e: # pylint: disable=broad-except
self.output_buffer.append(f"Error sending input: {str(e)}")
return f"Error sending input: {str(e)}" return f"Error sending input: {str(e)}"
def get_output(self, clear=True): def get_output(self, clear=True):
@ -138,8 +244,12 @@ class ShellSession: # pylint: disable=too-many-instance-attributes
def terminate(self): def terminate(self):
"""Terminate the session""" """Terminate the session"""
session_id_short = self.session_id[:8]
if not self.is_running: if not self.is_running:
return "Session already terminated" if self.process and self.process.poll() is None:
pass # Process is running, proceed with termination
else:
return f"Session {session_id_short} already terminated or finished."
try: try:
self.is_running = False self.is_running = False
@ -147,28 +257,64 @@ class ShellSession: # pylint: disable=too-many-instance-attributes
if self.process: if self.process:
# Try to terminate the process group # Try to terminate the process group
try: try:
os.killpg(os.getpgid(self.process.pid), signal.SIGTERM) os.killpg(os.getpgid(self.process.pid), signal.SIGTERM)
except BaseException: # pylint: disable=bare-except,broad-except # noqa: E501 except ProcessLookupError:
# If that fails, try to terminate just the process pass # Process already gone
self.process.terminate() except subprocess.TimeoutExpired:
print(color(f"Session {session_id_short} did not terminate gracefully, sending SIGKILL...", fg="yellow")) # noqa E501
try:
if pgid:
os.killpg(pgid, signal.SIGKILL) # Force kill
else:
self.process.kill()
except ProcessLookupError:
pass # Already gone
except Exception as kill_err:
termination_message = f" (Error during SIGKILL: {kill_err})"
except Exception as term_err: # Catch other errors during SIGTERM
termination_message = f" (Error during SIGTERM: {term_err})"
try:
self.process.kill()
except Exception: pass # Ignore nested errors
# Clean up resources
if self.master:
os.close(self.master)
if self.slave:
os.close(self.slave)
return f"Session {self.session_id} terminated" # Final check
if self.process.poll() is None:
print(color(f"Session {session_id_short} process {self.process.pid} may still be running after termination attempts.", fg="red")) # noqa E501
termination_message += " (Warning: Process may still be running)"
# Clean up PTY resources if they exist
if self.master:
try: os.close(self.master)
except OSError: pass
self.master = None
if self.slave:
try: os.close(self.slave)
except OSError: pass
self.slave = None
return termination_message or f"Session {self.session_id} terminated"
except Exception as e: # pylint: disable=broad-except except Exception as e: # pylint: disable=broad-except
return f"Error terminating session: {str(e)}" return f"Error terminating session {session_id_short}: {str(e)}"
def create_shell_session(command, ctf=None): def create_shell_session(command, ctf=None, container_id=None, **kwargs):
"""Create a new shell session""" """Create a new shell session in the correct workspace/environment."""
session = ShellSession(command, ctf=ctf) if container_id:
session = ShellSession(command, ctf=ctf, container_id=container_id)
else:
workspace_dir = _get_workspace_dir()
session = ShellSession(command, ctf=ctf, workspace_dir=workspace_dir)
session.start() session.start()
ACTIVE_SESSIONS[session.session_id] = session if session.is_running or (ctf and not session.is_running):
return session.session_id ACTIVE_SESSIONS[session.session_id] = session
return session.session_id
else:
error_msg = session.get_output(clear=True)
print(color(f"Failed to start session: {error_msg}", fg="red"))
return f"Failed to start session: {error_msg}"
def list_shell_sessions(): def list_shell_sessions():
@ -212,47 +358,100 @@ def get_session_output(session_id, clear=True):
def terminate_session(session_id): def terminate_session(session_id):
"""Terminate a specific session""" """Terminate a specific session"""
if session_id not in ACTIVE_SESSIONS: if session_id not in ACTIVE_SESSIONS:
return f"Session {session_id} not found" return f"Session {session_id} not found or already terminated."
session = ACTIVE_SESSIONS[session_id] session = ACTIVE_SESSIONS[session_id]
result = session.terminate() result = session.terminate()
del ACTIVE_SESSIONS[session_id] if session_id in ACTIVE_SESSIONS:
del ACTIVE_SESSIONS[session_id]
return result return result
def _run_ctf(ctf, command, stdout=False, timeout=100, stream=False, call_id=None): def _run_ctf(ctf, command, stdout=False, timeout=100, workspace_dir=None):
"""Runs command in CTF env, changing to workspace_dir first."""
target_dir = workspace_dir or _get_workspace_dir()
full_command = f"cd '{target_dir}' && {command}"
original_cmd_for_msg = command # For logging
context_msg = f"(ctf:{target_dir})"
try: try:
# Ensure the command is executed in a shell that supports command output = ctf.get_shell(full_command, timeout=timeout)
# chaining
output = ctf.get_shell(command, timeout=timeout)
# exploit_logger.log_ok()
if stdout: if stdout:
print("\033[32m" + output + "\033[0m") print(f"\033[32m{context_msg} $ {original_cmd_for_msg}\n{output}\033[0m") # noqa E501
return output # output if output else result.stder return output
except Exception as e: # pylint: disable=broad-except except Exception as e: # pylint: disable=broad-except
print(color(f"Error executing CTF command: {e}", fg="red")) error_msg = f"Error executing CTF command '{original_cmd_for_msg}' in '{target_dir}': {e}" # noqa E501
# exploit_logger.log_error(str(e)) print(color(error_msg, fg="red"))
return f"Error executing CTF command: {str(e)}" return error_msg
def _run_ssh(command, stdout=False, timeout=100, workspace_dir=None):
"""Runs command via SSH. Assumes SSH agent or passwordless setup unless sshpass is used externally.""" # noqa E501
ssh_user = os.environ.get('SSH_USER')
ssh_host = os.environ.get('SSH_HOST')
ssh_pass = os.environ.get('SSH_PASS')
remote_command = command
original_cmd_for_msg = command
context_msg = f"({ssh_user}@{ssh_host})"
# Construct base SSH command list
if ssh_pass:
ssh_cmd_list = ["sshpass", "-p", ssh_pass, "ssh", f"{ssh_user}@{ssh_host}"] # noqa E501
else:
ssh_cmd_list = ["ssh", f"{ssh_user}@{ssh_host}"]
ssh_cmd_list.append(remote_command)
try:
# Use subprocess.run with list of args for better security than shell=True
result = subprocess.run(
ssh_cmd_list,
capture_output=True,
text=True,
check=False, # Don't raise exception on non-zero exit code
timeout=timeout
)
output = result.stdout if result.stdout else result.stderr
if stdout:
print(f"\033[32m{context_msg} $ {original_cmd_for_msg}\n{output}\033[0m") # noqa E501
# Return combined output, potentially including errors
return output.strip()
except subprocess.TimeoutExpired as e:
error_output = e.stdout if e.stdout else str(e)
timeout_msg = f"Timeout executing SSH command: {error_output}"
if stdout:
print(f"\033[33m{context_msg} $ {original_cmd_for_msg}\nTIMEOUT\n{error_output}\033[0m") # noqa E501
return timeout_msg
except FileNotFoundError:
# Handle case where ssh or sshpass isn't installed
error_msg = f"'sshpass' or 'ssh' command not found. Ensure they are installed and in PATH." # noqa E501
print(color(error_msg, fg="red"))
return error_msg
except Exception as e: # pylint: disable=broad-except
error_msg = f"Error executing SSH command '{original_cmd_for_msg}' on {ssh_host}: {e}" # noqa E501
print(color(error_msg, fg="red"))
return error_msg
def _run_local(command, stdout=False, timeout=100, stream=False, call_id=None, tool_name=None): def _run_local(command, stdout=False, timeout=100, stream=False, call_id=None, tool_name=None, workspace_dir=None):
"""Runs command locally in the specified workspace_dir."""
# If streaming is enabled and we have a call_id # If streaming is enabled and we have a call_id
if stream and call_id: if stream and call_id:
return _run_local_streamed(command, call_id, timeout, tool_name) return _run_local_streamed(command, call_id, timeout, tool_name, workspace_dir)
target_dir = workspace_dir or _get_workspace_dir()
original_cmd_for_msg = command # For logging
context_msg = f"(local:{target_dir})"
try: try:
# nosec B602 - shell=True is required for command chaining
result = subprocess.run( result = subprocess.run(
command, command,
shell=True, # nosec B602 shell=True, # nosec B602
capture_output=True, capture_output=True,
text=True, text=True,
check=False, check=False,
timeout=timeout) timeout=timeout,
cwd=target_dir
)
output = result.stdout if result.stdout else result.stderr output = result.stdout if result.stdout else result.stderr
if stdout: if stdout:
print("\033[32m" + output + "\033[0m") print(f"\033[32m{context_msg} $ {original_cmd_for_msg}\n{output}\033[0m") # noqa E501
# Skip passing output to cli_print_tool_output when CAI_STREAM=true # Skip passing output to cli_print_tool_output when CAI_STREAM=true
# This prevents duplicate output in streaming mode # This prevents duplicate output in streaming mode
@ -261,20 +460,21 @@ def _run_local(command, stdout=False, timeout=100, stream=False, call_id=None, t
# Optional: Add cli_print_tool_output call here if needed for non-streaming # Optional: Add cli_print_tool_output call here if needed for non-streaming
pass pass
return output return output.strip()
except subprocess.TimeoutExpired as e: except subprocess.TimeoutExpired as e:
error_output = e.stdout.decode() if e.stdout else str(e) error_output = e.stdout if e.stdout else str(e)
if stdout: if stdout:
print("\033[32m" + error_output + "\033[0m") print("\033[32m" + error_output + "\033[0m")
return error_output return error_output
except Exception as e: # pylint: disable=broad-except except Exception as e: # pylint: disable=broad-except
error_msg = f"Error executing local command: {e}" error_msg = f"Error executing local command: {e}"
print(color(error_msg, fg="red")) print(color(error_msg, fg="red"))
return error_msg return error_msg
def _run_local_streamed(command, call_id, timeout=100, tool_name=None): def _run_local_streamed(command, call_id, timeout=100, tool_name=None, workspace_dir=None):
"""Run a local command with streaming output to the Tool output panel""" """Run a local command with streaming output to the Tool output panel."""
target_dir = workspace_dir or _get_workspace_dir()
try: try:
# Try to import Rich for nice display # Try to import Rich for nice display
try: try:
@ -298,7 +498,8 @@ def _run_local_streamed(command, call_id, timeout=100, tool_name=None):
stdout=subprocess.PIPE, stdout=subprocess.PIPE,
stderr=subprocess.PIPE, stderr=subprocess.PIPE,
text=True, text=True,
bufsize=1 bufsize=1,
cwd=target_dir # Set CWD for local process
) )
# If tool_name is not provided, derive it from the command # If tool_name is not provided, derive it from the command
@ -328,7 +529,7 @@ def _run_local_streamed(command, call_id, timeout=100, tool_name=None):
header.append(")", style="yellow") header.append(")", style="yellow")
tool_time = 0 tool_time = 0
start_time = time.time() start_time = time.time()
total_time = time.time() - START_TIME total_time = time.time() - START_TIME
timing_info = [] timing_info = []
if total_time: if total_time:
timing_info.append(f"Total: {format_time(total_time)}") timing_info.append(f"Total: {format_time(total_time)}")
@ -505,7 +706,8 @@ def run_command(command, ctf=None, stdout=False, # pylint: disable=too-many-arg
async_mode=False, session_id=None, async_mode=False, session_id=None,
timeout=100, stream=False, call_id=None, tool_name=None): timeout=100, stream=False, call_id=None, tool_name=None):
""" """
Run command either in CTF container or on the local attacker machine Run command in the appropriate environment (Docker, CTF, SSH, Local)
and workspace.
Args: Args:
command: The command to execute command: The command to execute
@ -520,34 +722,168 @@ def run_command(command, ctf=None, stdout=False, # pylint: disable=too-many-arg
If None, the tool name will be derived from the command. If None, the tool name will be derived from the command.
Returns: Returns:
str: Command output, status message, or session ID str: Command output, status message, or session ID.
""" """
# If session_id is provided, send command to that session # If session_id is provided, send command to that session
if session_id: if session_id:
if session_id not in ACTIVE_SESSIONS: if session_id not in ACTIVE_SESSIONS:
return f"Session {session_id} not found" return f"Session {session_id} not found"
session = ACTIVE_SESSIONS[session_id]
result = send_to_session(session_id, command) result = session.send_input(command) # Send the raw command string
if stdout: if stdout:
output = get_session_output(session_id, clear=False) output = get_session_output(session_id, clear=False)
print("\033[32m" + output + "\033[0m") env_type = "Local"
return result if session.container_id:
env_type = f"Container({session.container_id[:12]})"
# If async_mode, create a new session elif session.ctf:
if async_mode: env_type = "CTF"
session_id = create_shell_session(command, ctf) print(f"\033[32m(Session {session_id} in {env_type}:{session.workspace_dir}) >> {command}\n{output}\033[0m") # noqa E501
if stdout: return result # Return the result of sending input ("Input sent..." or error)
# Wait a moment for initial output
time.sleep(0.5)
output = get_session_output(session_id, clear=False)
print("\033[32m" + output + "\033[0m")
return f"Created session {session_id}. Use this ID to interact with the session."
# Generate a call_id if we're streaming and one wasn't provided # Generate a call_id if we're streaming and one wasn't provided
if stream and not call_id: if stream and not call_id:
call_id = str(uuid.uuid4())[:8] call_id = str(uuid.uuid4())[:8]
# Otherwise, run command normally # 2. Determine Execution Environment (Container > CTF > SSH > Local)
active_container = os.getenv("CAI_ACTIVE_CONTAINER", "")
is_ssh_env = all(os.getenv(var) for var in ['SSH_USER', 'SSH_HOST'])
# --- Docker Container Execution ---
if active_container and not ctf and not is_ssh_env:
container_id = active_container
container_workspace = _get_container_workspace_path()
context_msg = f"(docker:{container_id[:12]}:{container_workspace})"
# Handle Async Session Creation in Container
if async_mode:
# Create a session specifically for the container environment
new_session_id = create_shell_session(command, container_id=container_id) # noqa E501
if "Failed" in new_session_id: # Check if session creation failed
return new_session_id
if stdout:
# Wait a moment for initial output
time.sleep(0.2)
output = get_session_output(new_session_id, clear=False)
print(f"\033[32m(Started Session {new_session_id} in {context_msg})\n{output}\033[0m") # noqa E501
return f"Started async session {new_session_id} in container {container_id[:12]}. Use this ID to interact." # noqa E501
# Handle Streaming Container Execution - not yet implemented for containers
if stream:
# For now, display that streaming isn't supported for containers
from cai.util import cli_print_tool_output
if call_id and tool_name:
tool_args = {"command": command, "container": container_id[:12]}
cli_print_tool_output(
tool_name,
tool_args,
"Streaming not yet supported for container execution. Running normally...",
call_id=call_id
)
# Handle Synchronous Execution in Container
try:
# Ensure container workspace exists (best effort)
# Consider moving this to workspace set/container activation
mkdir_cmd = ["docker", "exec", container_id, "mkdir", "-p", container_workspace] # noqa E501
subprocess.run(mkdir_cmd, capture_output=True, text=True, check=False, timeout=10) # noqa E501
# Construct the docker exec command with workspace context
cmd_list = [
"docker", "exec",
"-w", container_workspace, # Set working directory
container_id,
"sh", "-c", command # Execute command via shell
]
result = subprocess.run(
cmd_list,
capture_output=True,
text=True,
check=False, # Don't raise exception on non-zero exit
timeout=timeout
)
output = result.stdout if result.stdout else result.stderr
output = output.strip() # Clean trailing newline
if stdout:
print(f"\033[32m{context_msg} $ {command}\n{output}\033[0m") # noqa E501
# Check if command failed specifically because container isn't running
if result.returncode != 0 and "is not running" in result.stderr:
print(color(f"{context_msg} Container is not running. Attempting execution on host instead.", fg="yellow")) # noqa E501
# Fallback to local execution, preserving workspace context
return _run_local(command, stdout, timeout, stream, call_id, tool_name, _get_workspace_dir()) # noqa E501
return output # Return combined stdout/stderr
except subprocess.TimeoutExpired:
timeout_msg = "Timeout executing command in container."
if stdout:
print(f"\033[33m{context_msg} $ {command}\nTIMEOUT\033[0m") # noqa E501
print(color("Attempting execution on host instead.", fg="yellow"))
# Fallback to local execution on timeout
return _run_local(command, stdout, timeout, stream, call_id, tool_name, _get_workspace_dir()) # noqa E501
except Exception as e: # pylint: disable=broad-except
error_msg = f"Error executing command in container: {str(e)}"
print(color(f"{context_msg} {error_msg}", fg="red"))
print(color("Attempting execution on host instead.", fg="yellow"))
# Fallback to local execution on other errors
return _run_local(command, stdout, timeout, stream, call_id, tool_name, _get_workspace_dir()) # noqa E501
# --- CTF Execution ---
if ctf: if ctf:
return _run_ctf(ctf, command, stdout, timeout, stream, call_id) # Handling streaming for CTF - not fully implemented yet
return _run_local(command, stdout, timeout, stream, call_id, tool_name) if stream:
from cai.util import cli_print_tool_output
if call_id and tool_name:
tool_args = {"command": command, "ctf": True}
cli_print_tool_output(
tool_name,
tool_args,
"Streaming not yet supported for CTF execution. Running normally...",
call_id=call_id
)
# _run_ctf handles workspace internally using _get_workspace_dir() default
return _run_ctf(ctf, command, stdout, timeout) # Pass None for workspace_dir
# --- SSH Execution ---
if is_ssh_env:
# Async for SSH would require session management via SSH client features
if async_mode:
return "Async mode not fully supported for SSH environment via this function yet."
# Handling streaming for SSH - not fully implemented yet
if stream:
from cai.util import cli_print_tool_output
if call_id and tool_name:
tool_args = {"command": command, "ssh": True}
cli_print_tool_output(
tool_name,
tool_args,
"Streaming not yet supported for SSH execution. Running normally...",
call_id=call_id
)
# _run_ssh handles command execution, workspace is relative to remote home
return _run_ssh(command, stdout, timeout) # Workspace dir less relevant here
# --- Local Execution (Default Fallback) ---
# Let _run_local handle determining the host workspace
# Handle Async Session Creation Locally
if async_mode:
# create_shell_session uses _get_workspace_dir() when container_id is None
new_session_id = create_shell_session(command)
if isinstance(new_session_id, str) and "Failed" in new_session_id: # Check failure
return new_session_id
# Retrieve the actual workspace dir the session is using
session = ACTIVE_SESSIONS.get(new_session_id)
actual_workspace = session.workspace_dir if session else "unknown"
if stdout:
time.sleep(0.2) # Allow session buffer to populate
output = get_session_output(new_session_id, clear=False)
print(f"\033[32m(Started Session {new_session_id} in local:{actual_workspace})\n{output}\033[0m")
return f"Started async session {new_session_id} locally. Use this ID to interact."
# Handle Synchronous Execution Locally using _run_local default with streaming support
return _run_local(command, stdout, timeout, stream, call_id, tool_name, None)

View File

@ -268,87 +268,72 @@ def load_prompt_template(template_path):
except Exception as e: except Exception as e:
raise ValueError(f"Failed to load template '{template_path}': {str(e)}") raise ValueError(f"Failed to load template '{template_path}': {str(e)}")
# Start of Selection
def visualize_agent_graph(start_agent): def visualize_agent_graph(start_agent):
""" """
Visualize agent graph showing all bidirectional connections between agents. Visualize agent graph showing all bidirectional connections between agents.
Uses Rich library for pretty printing. Uses Rich library for pretty printing.
""" """
console = Console() # pylint: disable=redefined-outer-name console = Console()
if start_agent is None: if start_agent is None:
console.print("[red]No agent provided to visualize.[/red]") console.print("[red]No agent provided to visualize.[/red]")
return return
tree = Tree( tree = Tree(f"🤖 {start_agent.name} (Current Agent)", guide_style="bold blue")
f"🤖 {
start_agent.name} (Current Agent)",
guide_style="bold blue")
# Track visited agents and their nodes to handle cross-connections visited = set()
visited = {}
agent_nodes = {} agent_nodes = {}
agent_positions = {} # Track positions in tree agent_positions = {}
position_counter = 0 # Counter for tracking positions position_counter = 0
def add_agent_node(agent, parent=None, is_transfer=False): # pylint: disable=too-many-branches # noqa: E501 def add_agent_node(agent, parent=None, is_transfer=False):
"""Add agent node and track for cross-connections""" """Add an agent node and track for cross-connections."""
nonlocal position_counter nonlocal position_counter
if agent is None: if agent is None:
return None return None
aid = id(agent)
if aid in visited:
if is_transfer and parent:
original_pos = agent_positions.get(aid)
parent.add(f"[cyan]↩ Return to {agent.name} (Agent #{original_pos})[/cyan]")
return agent_nodes.get(aid)
# Create or get existing node for this agent visited.add(aid)
if id(agent) in visited:
if is_transfer:
# Add reference with position for repeated agents
original_pos = agent_positions[id(agent)]
parent.add(
f"[cyan]↩ Return to {
agent.name} (Top Level Agent #{original_pos})[/cyan]")
return agent_nodes[id(agent)]
visited[id(agent)] = True
position_counter += 1 position_counter += 1
agent_positions[id(agent)] = position_counter agent_positions[aid] = position_counter
# Create node for current agent if is_transfer and parent:
if is_transfer:
node = parent node = parent
elif parent:
node = parent.add(f"[green]{agent.name} (#{position_counter})[/green]")
else: else:
node = parent.add( node = tree
f"[green]{agent.name} (#{position_counter})[/green]") if parent else tree # noqa: E501 pylint: disable=line-too-long agent_nodes[aid] = node
agent_nodes[id(agent)] = node
# Add tools as children # Add tools
tools_node = node.add("[yellow]Tools[/yellow]") tools_node = node.add("[yellow]Tools[/yellow]")
for fn in getattr(agent, "functions", []): for tool in getattr(agent, "tools", []):
if callable(fn): tool_name = getattr(tool, "name", None) or getattr(tool, "__name__", "")
fn_name = getattr(fn, "__name__", "") tools_node.add(f"[blue]{tool_name}[/blue]")
if ("handoff" not in fn_name.lower() and
not fn_name.startswith("transfer_to")):
tools_node.add(f"[blue]{fn_name}[/blue]")
# Add Handoffs section # Add handoffs
transfers_node = node.add("[magenta]Handoffs[/magenta]") transfers_node = node.add("[magenta]Handoffs[/magenta]")
for handoff_fn in getattr(agent, "handoffs", []):
if callable(handoff_fn):
try:
next_agent = handoff_fn()
if next_agent:
transfer_node = transfers_node.add(f"🤖 {next_agent.name}")
add_agent_node(next_agent, transfer_node, True)
except Exception:
continue
# Process handoff functions
for fn in getattr(agent, "functions", []): # pylint: disable=too-many-nested-blocks # noqa: E501
if callable(fn):
fn_name = getattr(fn, "__name__", "")
if ("handoff" in fn_name.lower() or
fn_name.startswith("transfer_to")):
try:
next_agent = fn()
if next_agent:
# Show bidirectional connection
transfer = transfers_node.add(
f"🤖 {next_agent.name}") # noqa: E501
add_agent_node(next_agent, transfer, True)
except Exception: # nosec: B112 # pylint: disable=broad-exception-caught # noqa: E501
continue
return node return node
# Start recursive traversal from root agent
# Start traversal from the root agent
add_agent_node(start_agent) add_agent_node(start_agent)
console.print(tree) console.print(tree)
# End of Selectio
def fix_litellm_transcription_annotations(): def fix_litellm_transcription_annotations():
""" """
@ -1005,6 +990,7 @@ def update_agent_streaming_content(context, text_delta):
# Force an update with the new panel # Force an update with the new panel
context["live"].update(updated_panel) context["live"].update(updated_panel)
context["panel"] = updated_panel context["panel"] = updated_panel
context["live"].refresh()
def finish_agent_streaming(context, final_stats=None): def finish_agent_streaming(context, final_stats=None):
"""Finish the streaming session and display final stats if available.""" """Finish the streaming session and display final stats if available."""