diff --git a/src/cai/cli.py b/src/cai/cli.py index 827f3733..fec52b12 100644 --- a/src/cai/cli.py +++ b/src/cai/cli.py @@ -58,6 +58,10 @@ Environment Variables CAI_STREAM: Enable/disable streaming output in rich panel (default: "true") CAI_TELEMETRY: Enable/disable telemetry (default: "true") + CAI_PARALLEL: Number of parallel agent instances to run + (default: "1"). When set to values greater than 1, + executes multiple instances of the same agent in + parallel and displays all results. Extensions (only applicable if the right extension is installed): @@ -97,6 +101,10 @@ Usage Examples: CAI_MODEL="o3-mini" CAI_MEMORY_ONLINE_INTERVAL="3" \ CAI_MEMORY_ONLINE="False" CTF_INSIDE="False" \ CTF_HINTS="False" python3 cai/cli.py + + # Run with parallel agents (3 instances) + CTF_NAME="hackableII" CAI_AGENT_TYPE="redteam_agent" \ + CAI_MODEL="gpt-4o" CAI_PARALLEL="3" python3 cai/cli.py """ import os @@ -104,6 +112,7 @@ import time import asyncio from dotenv import load_dotenv from rich.console import Console +from rich.panel import Panel # OpenAI imports from openai import AsyncOpenAI @@ -140,6 +149,9 @@ from cai.repl.ui.toolbar import get_toolbar_with_refresh from cai.agents import get_agent_by_name from cai.internal.components.metrics import process_metrics +# Add import for parallel configs at the top of the file +from cai.repl.commands.parallel import PARALLEL_CONFIGS, ParallelConfig + # Load environment variables from .env file load_dotenv() @@ -196,6 +208,7 @@ def run_cai_cli(starting_agent, context_variables=None, stream=False, max_turns= console = Console() last_model = os.getenv('CAI_MODEL', 'qwen2.5:14b') last_agent_type = os.getenv('CAI_AGENT_TYPE', 'one_tool_agent') + parallel_count = int(os.getenv('CAI_PARALLEL', '1')) # Initialize command completer and key bindings command_completer = FuzzyCommandCompleter() @@ -386,6 +399,58 @@ def run_cai_cli(starting_agent, context_variables=None, stream=False, max_turns= break try: + # Check if we have parallel configurations to run + if PARALLEL_CONFIGS and not user_input.startswith('/') and not user_input.startswith('$'): + # Use parallel configurations instead of normal processing + console.print(f"[bold cyan]Running {len(PARALLEL_CONFIGS)} parallel agents...[/bold cyan]") + + async def run_agent_instance(config: ParallelConfig, input_text: str): + """Run a single agent instance with its own configuration.""" + try: + # Create a fresh agent instance + instance_agent = get_agent_by_name(config.agent_name) + + # Override model if specified + if config.model and hasattr(instance_agent, 'model'): + if hasattr(instance_agent.model, 'model'): + instance_agent.model.model = config.model + + # Override prompt if specified + instance_input = config.prompt or input_text + + # Run the agent with its own isolated context + result = await Runner.run(instance_agent, instance_input) + + return (config, result) + except Exception as e: + import traceback + console.print(f"[bold red]Error in {config.agent_name}: {str(e)}[/bold red]") + return (config, None) + + async def run_parallel_agents(): + """Run all configured agents in parallel.""" + # Create tasks for each agent + tasks = [run_agent_instance(config, user_input) + for config in PARALLEL_CONFIGS] + + # Wait for all to complete + results = await asyncio.gather(*tasks, return_exceptions=True) + + # Filter out exceptions and failed results + valid_results = [] + for item in results: + if isinstance(item, tuple) and len(item) == 2 and item[1] is not None: + valid_results.append(item) + + return valid_results + + # Run in asyncio event loop + results = asyncio.run(run_parallel_agents()) + turn_count += 1 + stop_active_timer() + start_idle_timer() + continue + # Handle special commands if user_input.startswith('/') or user_input.startswith('$'): parts = user_input.strip().split() @@ -448,7 +513,7 @@ def run_cai_cli(starting_agent, context_variables=None, stream=False, max_turns= from cai.util import fix_message_list history_context = fix_message_list(history_context) except Exception as e: - console.print(f"[yellow]Warning: Could not preprocess message history: {e}[/yellow]") + pass # Append the current user input as the last message in the list. conversation_input: list | str @@ -458,79 +523,137 @@ def run_cai_cli(starting_agent, context_variables=None, stream=False, max_turns= else: conversation_input = user_input - # Process the conversation with the agent. - if stream: - async def process_streamed_response(): + # Process the conversation with the agent - with parallel execution if enabled + if parallel_count > 1: + # Parallel execution mode (always non-streaming) + async def run_agent_instance(instance_number): + """Run a single agent instance with its own complete context""" try: - result = Runner.run_streamed(agent, conversation_input) - - # Consume events so the async generator is executed. - async for event in result.stream_events(): - if isinstance(event, RunItemStreamEvent) and event.name == "tool_output": - # Ensure item is a ToolCallOutputItem before accessing attributes - if isinstance(event.item, ToolCallOutputItem): - tool_msg = { - "role": "tool", - "tool_call_id": event.item.raw_item["call_id"], # Changed to dictionary access - "content": event.item.output, - } - add_to_message_history(tool_msg) - # pass # Original logic was just pass - - return result + # Create a fresh agent instance to ensure complete isolation + instance_agent = get_agent_by_name(last_agent_type) + + # Configure agent instance to match main agent settings + if hasattr(instance_agent, 'model') and hasattr(agent, 'model'): + if hasattr(instance_agent.model, 'model') and hasattr(agent.model, 'model'): + instance_agent.model.model = agent.model.model + + # Create a fresh input for this instance - use just the user input directly + # This ensures each instance has its own completely independent context + instance_input = user_input + + # Run the agent with its own isolated context + result = await Runner.run(instance_agent, instance_input) + + return (instance_number, result) except Exception as e: import traceback - tb = traceback.format_exc() - print( - f"\n[Error occurred during streaming: {str(e)}]" - f"\nLocation: {tb}" - ) - return None - - asyncio.run(process_streamed_response()) - else: - # Use non-streamed response - response = asyncio.run(Runner.run(agent, conversation_input)) + console.print(f"[bold red]Error in instance {instance_number}: {str(e)}[/bold red]") + return (instance_number, None) - # Process the response items - for item in response.new_items: - # Handle tool call output items (tool results) - if isinstance(item, ToolCallOutputItem): - # First, ensure there's a corresponding assistant message with tool_calls - # before adding the tool response to prevent the OpenAI error - assistant_with_tool_call_exists = False - tool_call_id = item.raw_item["call_id"] - - for msg in message_history: - if (msg.get("role") == "assistant" and - msg.get("tool_calls") and - any(tc.get("id") == tool_call_id for tc in msg.get("tool_calls", []))): - assistant_with_tool_call_exists = True - break - - # If no matching assistant message exists, create one first - if not assistant_with_tool_call_exists: - tool_call_msg = { - "role": "assistant", - "content": None, - "tool_calls": [{ - "id": tool_call_id, - "type": "function", - "function": { - "name": "unknown_function", - "arguments": "{}" - } - }] + async def process_parallel_responses(): + """Process multiple parallel agent executions""" + # Create tasks for each instance + tasks = [run_agent_instance(i) for i in range(parallel_count)] + + # Wait for all to complete, no matter if some fail + results = await asyncio.gather(*tasks, return_exceptions=True) + + # Filter out exceptions and failed results + valid_results = [] + for idx, result in results: + if result is not None and not isinstance(result, Exception): + valid_results.append((idx, result)) + + return valid_results + + # Execute all parallel instances + results = asyncio.run(process_parallel_responses()) + + # Print summary info about the results + console.print(f"[bold cyan]Completed {len(results)}/{parallel_count} parallel agent executions[/bold cyan]") + + # Display the results + for idx, result in results: + if result and hasattr(result, 'final_output') and result.final_output: + # Add to main message history for context + add_to_message_history({ + "role": "assistant", + "content": f"{result.final_output}" + }) + else: + # Single agent execution (original behavior) + if stream: + async def process_streamed_response(agent, conversation_input): + try: + result = Runner.run_streamed(agent, conversation_input) + + # Consume events so the async generator is executed. + async for event in result.stream_events(): + if isinstance(event, RunItemStreamEvent) and event.name == "tool_output": + # Ensure item is a ToolCallOutputItem before accessing attributes + if isinstance(event.item, ToolCallOutputItem): + tool_msg = { + "role": "tool", + "tool_call_id": event.item.raw_item["call_id"], # Changed to dictionary access + "content": event.item.output, + } + add_to_message_history(tool_msg) + # pass # Original logic was just pass + + return result + except Exception as e: + import traceback + tb = traceback.format_exc() + print( + f"\n[Error occurred during streaming: {str(e)}]" + f"\nLocation: {tb}" + ) + return None + + asyncio.run(process_streamed_response(agent, conversation_input)) + else: + # Use non-streamed response + response = asyncio.run(Runner.run(agent, conversation_input)) + + # Process the response items + for item in response.new_items: + # Handle tool call output items (tool results) + if isinstance(item, ToolCallOutputItem): + # First, ensure there's a corresponding assistant message with tool_calls + # before adding the tool response to prevent the OpenAI error + assistant_with_tool_call_exists = False + tool_call_id = item.raw_item["call_id"] + + for msg in message_history: + if (msg.get("role") == "assistant" and + msg.get("tool_calls") and + any(tc.get("id") == tool_call_id for tc in msg.get("tool_calls", []))): + assistant_with_tool_call_exists = True + break + + # If no matching assistant message exists, create one first + if not assistant_with_tool_call_exists: + tool_call_msg = { + "role": "assistant", + "content": None, + "tool_calls": [{ + "id": tool_call_id, + "type": "function", + "function": { + "name": "unknown_function", + "arguments": "{}" + } + }] + } + add_to_message_history(tool_call_msg) + + # Now add the tool response + tool_msg = { + "role": "tool", + "tool_call_id": tool_call_id, + "content": item.output, } - add_to_message_history(tool_call_msg) - - # Now add the tool response - tool_msg = { - "role": "tool", - "tool_call_id": tool_call_id, - "content": item.output, - } - add_to_message_history(tool_msg) + add_to_message_history(tool_msg) # Make sure that assistant messages with tool calls are also added to message_history # This is especially important for non-streaming mode diff --git a/src/cai/repl/commands/__init__.py b/src/cai/repl/commands/__init__.py index 41bac88a..1ac26681 100644 --- a/src/cai/repl/commands/__init__.py +++ b/src/cai/repl/commands/__init__.py @@ -39,7 +39,8 @@ from cai.repl.commands import ( # pylint: disable=import-error,unused-import,li flush, workspace, virtualization, - load + load, + parallel # Add the new parallel command ) # Define helper functions diff --git a/src/cai/repl/commands/parallel.py b/src/cai/repl/commands/parallel.py new file mode 100644 index 00000000..6c6195c3 --- /dev/null +++ b/src/cai/repl/commands/parallel.py @@ -0,0 +1,190 @@ +""" +Parallel command for CAI CLI abstraction. + +Provides commands for managing parallel agent configurations. +Different agents can be configured with specific models and prompts, +which will then be executed in parallel through the CLI. +""" + +# Standard library imports +import os +from typing import List, Optional, Dict + +# Third-party imports +from rich.console import Console +from rich.table import Table + +# Local imports +from cai.repl.commands.base import Command, register_command +from cai.agents import get_available_agents + +console = Console() + +# Store configured parallel runs (made global so CLI can access it) +PARALLEL_CONFIGS = [] + + +class ParallelConfig: + """Configuration for a parallel agent run.""" + + def __init__(self, agent_name, model=None, prompt=None): + """Initialize a parallel agent configuration. + + Args: + agent_name: Name of the agent to use + model: Optional model to use (overrides default) + prompt: Optional specific prompt to use + """ + self.agent_name = agent_name + self.model = model + self.prompt = prompt + + def __str__(self): + """String representation of the configuration.""" + model_str = f", model: {self.model}" if self.model else "" + prompt_str = f", prompt: '{self.prompt[:20]}...'" if self.prompt and len(self.prompt) > 20 else \ + f", prompt: '{self.prompt}'" if self.prompt else "" + return f"Agent: {self.agent_name}{model_str}{prompt_str}" + + +class ParallelCommand(Command): + """Command for managing parallel agent configurations.""" + + def __init__(self): + """Initialize the parallel command.""" + super().__init__( + name="/parallel", + description="Configure multiple agents to run in parallel with different settings", + aliases=["/par", "/p"] + ) + + # Add subcommands for configuration management + self.add_subcommand("add", "Add a new agent to the parallel config", self.handle_add) + self.add_subcommand("list", "List configured parallel agents", self.handle_list) + self.add_subcommand("clear", "Clear all configured parallel agents", self.handle_clear) + self.add_subcommand("remove", "Remove a specific parallel agent by index", self.handle_remove) + + def handle_add(self, args: Optional[List[str]] = None) -> bool: + """Handle the add subcommand. + + Args: + args: Command arguments [agent_name] [--model MODEL] [--prompt PROMPT] + + Returns: + True if successful + """ + if not args: + console.print("[red]Error: Agent name required[/red]") + console.print("Usage: /parallel add [--model MODEL] [--prompt PROMPT]") + return False + + agent_name = args[0] + + # Check if agent exists + available_agents = get_available_agents() + if agent_name not in available_agents: + console.print(f"[red]Error: Unknown agent '{agent_name}'[/red]") + console.print("Available agents:") + for idx, name in enumerate(available_agents.keys(), 1): + console.print(f" {idx}. {name}") + return False + + # Parse optional arguments + model = None + prompt = None + i = 1 + while i < len(args): + if args[i] == "--model" and i + 1 < len(args): + model = args[i + 1] + i += 2 + elif args[i] == "--prompt" and i + 1 < len(args): + prompt = args[i + 1] + i += 2 + else: + i += 1 + + # Add configuration + config = ParallelConfig(agent_name, model, prompt) + PARALLEL_CONFIGS.append(config) + + console.print(f"[green]Added parallel configuration:[/green] {config}") + console.print("[cyan]Use your next query to run all parallel agents[/cyan]") + return True + + def handle_list(self, args: Optional[List[str]] = None) -> bool: + """Handle the list subcommand. + + Args: + args: Command arguments (unused) + + Returns: + True if successful + """ + if not PARALLEL_CONFIGS: + console.print("[yellow]No parallel configurations defined[/yellow]") + console.print("Use '/parallel add ' to add a configuration") + return True + + table = Table(title="Configured Parallel Agents") + table.add_column("#", style="dim") + table.add_column("Agent", style="cyan") + table.add_column("Model", style="green") + table.add_column("Custom Prompt", style="yellow") + + for idx, config in enumerate(PARALLEL_CONFIGS, 1): + prompt_display = (config.prompt[:30] + "...") if config.prompt and len(config.prompt) > 30 else config.prompt or "" + table.add_row( + str(idx), + config.agent_name, + config.model or "Default", + prompt_display + ) + + console.print(table) + console.print("[cyan]Your next query will run all these agents in parallel[/cyan]") + return True + + def handle_clear(self, args: Optional[List[str]] = None) -> bool: + """Handle the clear subcommand. + + Args: + args: Command arguments (unused) + + Returns: + True if successful + """ + global PARALLEL_CONFIGS + count = len(PARALLEL_CONFIGS) + PARALLEL_CONFIGS = [] + + console.print(f"[green]Cleared {count} parallel configurations[/green]") + return True + + def handle_remove(self, args: Optional[List[str]] = None) -> bool: + """Handle the remove subcommand. + + Args: + args: Command arguments [index] + + Returns: + True if successful + """ + if not args: + console.print("[red]Error: Index required[/red]") + console.print("Usage: /parallel remove ") + return False + + try: + idx = int(args[0]) + if idx < 1 or idx > len(PARALLEL_CONFIGS): + raise ValueError("Index out of range") + + removed = PARALLEL_CONFIGS.pop(idx - 1) + console.print(f"[green]Removed configuration:[/green] {removed}") + return True + except ValueError: + console.print(f"[red]Error: Invalid index '{args[0]}'[/red]") + return False + +# Register the command +register_command(ParallelCommand()) \ No newline at end of file