From 027db7cef37500f059acd19b3e7678be5e3371b1 Mon Sep 17 00:00:00 2001 From: luijait Date: Sun, 11 May 2025 21:36:09 +0200 Subject: [PATCH 1/2] Handoffs --- src/cai/agents/patterns/red_team.py | 42 +- src/cai/util.py | 590 ++++++++++++++++++---------- 2 files changed, 390 insertions(+), 242 deletions(-) diff --git a/src/cai/agents/patterns/red_team.py b/src/cai/agents/patterns/red_team.py index 00b408ad..bf950f67 100644 --- a/src/cai/agents/patterns/red_team.py +++ b/src/cai/agents/patterns/red_team.py @@ -1,4 +1,3 @@ - """ Implementation of a Cyclic Swarm Pattern for Red Team Operations @@ -11,36 +10,29 @@ for comprehensive security analysis. from cai.agents.red_teamer import redteam_agent from cai.agents.thought import thought_agent from cai.agents.mail import dns_smtp_agent +from cai.sdk.agents import handoff -def transfer_to_dns_agent(): - """ - Use THIS always for DNS scans and domain reconnaissance - about dmarc and dkim registers - """ - return dns_smtp_agent +# Create handoffs using the SDK handoff function +dns_smtp_handoff = handoff( + agent=dns_smtp_agent, + tool_description_override="Use for DNS scans and domain reconnaissance about DMARC and DKIM records" +) +redteam_handoff = handoff( + agent=redteam_agent, + tool_description_override="Transfer to Red Team Agent for security assessment and exploitation tasks" +) -def redteam_agent_handoff(ctf=None): # pylint: disable=unused-argument - """ - Red Team Agent, call this function - empty to transfer to redteam_agent - """ - return redteam_agent - - -def thought_agent_handoff(ctf=None): # pylint: disable=unused-argument - """ - Thought Agent, call this function empty - to transfer to thought_agent - """ - return thought_agent - +thought_handoff = handoff( + agent=thought_agent, + tool_description_override="Transfer to Thought Agent for analysis and planning" +) # Register handoff to enable inter-agent communication pathways -redteam_agent.handoffs.append(transfer_to_dns_agent) -dns_smtp_agent.handoffs.append(redteam_agent_handoff) -thought_agent.handoffs.append(redteam_agent_handoff) +redteam_agent.handoffs.append(dns_smtp_handoff) +dns_smtp_agent.handoffs.append(redteam_handoff) +thought_agent.handoffs.append(redteam_handoff) # Initialize the swarm pattern with the thought agent as the entry point redteam_swarm_pattern = thought_agent diff --git a/src/cai/util.py b/src/cai/util.py index 00e4bbd4..094b177c 100644 --- a/src/cai/util.py +++ b/src/cai/util.py @@ -28,6 +28,7 @@ from rich.console import Group from rich.box import ROUNDED from rich.table import Table import re +import uuid # Global timing variables for tracking active and idle time _active_timer_start = None @@ -496,8 +497,10 @@ def visualize_agent_graph(start_agent): # Add handoffs transfers_node = node.add("[magenta]Handoffs[/magenta]") + + # First, handle old-style handoffs through handoffs list for handoff_fn in getattr(agent, "handoffs", []): - if callable(handoff_fn): + if callable(handoff_fn) and not hasattr(handoff_fn, "agent_name"): try: next_agent = handoff_fn() if next_agent: @@ -505,13 +508,44 @@ def visualize_agent_graph(start_agent): add_agent_node(next_agent, transfer_node, True) except Exception: continue + elif hasattr(handoff_fn, "agent_name"): + # Handle SDK handoff objects + try: + handoff_name = handoff_fn.agent_name + # Find the actual agent instance if available + next_agent = None + + # Try to find the agent by name in the global namespace + # This is a heuristic and might not always work + import sys + for module_name, module in sys.modules.items(): + if module_name.startswith('cai.agents'): + agent_var_name = handoff_name.lower().replace(' ', '_') + '_agent' + if hasattr(module, agent_var_name): + next_agent = getattr(module, agent_var_name) + break + + if next_agent: + transfer_node = transfers_node.add( + f"🤖 {handoff_name} via {handoff_fn.tool_name}") + add_agent_node(next_agent, transfer_node, True) + else: + # If we can't find the agent, just show the name + transfers_node.add( + f"[yellow]🤖 {handoff_name} via {handoff_fn.tool_name}[/yellow]") + except Exception as e: + transfers_node.add(f"[red]Error: {str(e)}[/red]") + elif isinstance(handoff_fn, dict) and "agent_name" in handoff_fn: + # Handle dictionary handoff objects + handoff_name = handoff_fn["agent_name"] + tool_name = handoff_fn.get("tool_name", f"transfer_to_{handoff_name}") + transfers_node.add(f"[yellow]🤖 {handoff_name} via {tool_name}[/yellow]") return node # Start traversal from the root agent add_agent_node(start_agent) console.print(tree) -# End of Selectio def fix_litellm_transcription_annotations(): """ @@ -1573,6 +1607,36 @@ def finish_agent_streaming(context, final_stats=None): interaction_cost, total_cost ) + + # Crear una línea de tokens compacta para el streaming + compact_tokens = Text() + compact_tokens.append(" | ", style="dim") + compact_tokens.append(f"I:{interaction_input_tokens} O:{interaction_output_tokens} ", style="green") + compact_tokens.append(f"(${interaction_cost:.4f}) ", style="bold cyan") + + # Añadir un indicador de uso de contexto + context_pct = interaction_input_tokens / get_model_input_tokens(model_name) * 100 + if context_pct < 50: + indicator = "🟩" + elif context_pct < 80: + indicator = "🟨" + else: + indicator = "🟥" + compact_tokens.append(f"{indicator} {context_pct:.1f}%", style="bold") + + # Add the compact token info to the footer + if 'footer' in context and final_stats: + # Clear the existing footer + context['footer'] = Text() + # Add timestamp and model + context['footer'].append(f"\n[{context['timestamp']}", style="dim") + if context['model']: + context['footer'].append(f" ({context['model']})", style="bold magenta") + context['footer'].append("]", style="dim") + + # Add the compact token info if available + if final_stats and 'compact_tokens' in locals(): + context['footer'].append(compact_tokens) final_panel = Panel( Text.assemble( @@ -1832,17 +1896,59 @@ def cli_print_tool_output(tool_name="", args="", output="", call_id=None, execut status = execution_info.get('status', 'completed') if status == "completed": border_style = "green" - title = f"[bold green]{tool_name}({args_str}) [Completed][/bold green]" elif status == "error": border_style = "red" - title = f"[bold red]{tool_name}({args_str}) [Error][/bold red]" elif status == "timeout": border_style = "red" - title = f"[bold red]{tool_name}({args_str}) [Timeout][/bold red]" + + # Check if this is a handoff (transfer to another agent) + is_handoff = tool_name.startswith("transfer_to_") + + # Create the title based on whether it's a handoff or regular tool + if is_handoff: + # Extract agent name for the handoff title + agent_name = None + if tool_name.startswith("transfer_to_"): + # Remove 'transfer_to_' prefix and convert to a nicer format + agent_name_raw = tool_name[len("transfer_to_"):] + # Convert underscores to spaces and capitalize words + agent_name = " ".join(word.capitalize() for word in agent_name_raw.split("_")) + + # Special case for acronyms like DNS or SMTP that might be in the agent name + # Convert words that are all uppercase to remain uppercase + parts = agent_name.split() + for i, part in enumerate(parts): + if part.upper() == part and len(part) > 1: # It's an acronym + parts[i] = part.upper() + agent_name = " ".join(parts) + + # For handoffs, include the agent name in the title + if execution_info: + status = execution_info.get('status', 'completed') + if status == "completed": + title = f"[bold green]Handoff: {agent_name} [Completed][/bold green]" + elif status == "error": + title = f"[bold red]Handoff: {agent_name} [Error][/bold red]" + elif status == "timeout": + title = f"[bold red]Handoff: {agent_name} [Timeout][/bold red]" + else: + title = f"[bold blue]Handoff: {agent_name}[/bold blue]" + else: + title = f"[bold blue]Handoff: {agent_name}[/bold blue]" + else: + # For regular tools, use the original format + if execution_info: + status = execution_info.get('status', 'completed') + if status == "completed": + title = f"[bold green]{tool_name}({args_str}) [Completed][/bold green]" + elif status == "error": + title = f"[bold red]{tool_name}({args_str}) [Error][/bold red]" + elif status == "timeout": + title = f"[bold red]{tool_name}({args_str}) [Timeout][/bold red]" + else: + title = f"[bold blue]{tool_name}({args_str})[/bold blue]" else: title = f"[bold blue]{tool_name}({args_str})[/bold blue]" - else: - title = f"[bold blue]{tool_name}({args_str})[/bold blue]" # Create the panel panel = Panel( @@ -1858,10 +1964,208 @@ def cli_print_tool_output(tool_name="", args="", output="", call_id=None, execut console.print(panel) except ImportError: - pass + # Fall back to simple output format without rich + _print_simple_tool_output(tool_name, args, output, execution_info, token_info) -# Helper function to create tool panel content +# Helper function to format tool arguments +def _format_tool_args(args, tool_name=None): + """Format tool arguments as a clean string.""" + # If the tool is execute_code, we don't want to show any args in the main header, + # as they are detailed in subsequent panels (either code or args string). + if tool_name == "execute_code": + return "" + + # If args is already a string, it might be pre-formatted or a simple arg string + if isinstance(args, str): + # If it looks like a JSON dict string, try to parse and format nicely + if args.strip().startswith('{') and args.strip().endswith('}'): + try: + parsed_dict = json.loads(args) + # Recursively call with the parsed dict for consistent formatting + return _format_tool_args(parsed_dict, tool_name=tool_name) + except json.JSONDecodeError: + # Not valid JSON, or not a dict; return as is + return args + else: + # Simple string arg, return as is + return args + + # Format arguments from a dictionary + if isinstance(args, dict): + # Only include non-empty values and exclude special flags + arg_parts = [] + for key, value in args.items(): + # Skip empty values + if value == "" or value == {} or value is None: + continue + # Skip special flags + if key in ["async_mode", "streaming"] and not value: + continue + + value_str = str(value) + + # Format the value + if isinstance(value, str): + # Truncate long string values + if len(value_str) > 70 and key not in ["code", "args"]: + value_str = value_str[:67] + "..." + arg_parts.append(f"{key}={value_str}") + else: + arg_parts.append(f"{key}={value_str}") + return ", ".join(arg_parts) + else: + return str(args) + +def print_message_history(messages, title="Message History"): + """ + Pretty-print a sequence of messages with enhanced debug information. + + Args: + messages (List[dict]): List of message dictionaries to display + title (str, optional): Title to display above the message history + """ + from rich.console import Console + from rich.panel import Panel + from rich.text import Text + from rich.table import Table + + console = Console() + + # Create a table for displaying messages + table = Table(show_header=True, header_style="bold magenta", expand=True) + table.add_column("#", style="dim", width=3) + table.add_column("Role", style="cyan", width=10) + table.add_column("Content", width=40) + table.add_column("Metadata", width=30) + + # Process each message + for i, msg in enumerate(messages): + # Get role with color based on type + role = msg.get("role", "unknown") + role_style = { + "user": "green", + "assistant": "blue", + "system": "yellow", + "tool": "magenta" + }.get(role, "white") + + # Get content preview + content = msg.get("content") + content_preview = "" + if content is None: + content_preview = "[dim]None[/dim]" + elif isinstance(content, str): + # Truncate and escape long content + content_preview = (content[:37] + "...") if len(content) > 40 else content + content_preview = content_preview.replace("\n", "\\n") + elif isinstance(content, list): + content_preview = f"[list with {len(content)} items]" + else: + content_preview = f"[{type(content).__name__}]" + + # Gather metadata + metadata = [] + if msg.get("tool_calls"): + tc_count = len(msg["tool_calls"]) + tc_info = [] + for tc in msg["tool_calls"]: + tc_id = tc.get("id", "unknown") + tc_name = tc.get("function", {}).get("name", "unknown") if "function" in tc else "unknown" + tc_info.append(f"{tc_name}({tc_id})") + metadata.append(f"tool_calls[{tc_count}]: {', '.join(tc_info)}") + + if msg.get("tool_call_id"): + metadata.append(f"tool_call_id: {msg['tool_call_id']}") + + metadata_str = ", ".join(metadata) + + # Add row to table + table.add_row( + str(i), + f"[{role_style}]{role}[/{role_style}]", + content_preview, + metadata_str + ) + + # Create the panel with the table + panel = Panel( + table, + title=f"[bold]{title}[/bold]", + expand=False + ) + + # Display the panel + console.print(panel) + + return len(messages) # Return message count for convenience + +def get_language_from_code_block(lang_identifier): + """ + Maps a language identifier from a markdown code block to a proper syntax + highlighting language name. Handles common aliases and defaults. + + Args: + lang_identifier (str): Language identifier from markdown code block + + Returns: + str: Proper language name for syntax highlighting + """ + # Convert to lowercase and strip whitespace + lang = lang_identifier.lower().strip() if lang_identifier else "" + + # Map common language aliases to their proper names + lang_map = { + # Empty strings or unknown + "": "text", + # Python variants + "py": "python", + "python3": "python", + # JavaScript variants + "js": "javascript", + "jsx": "jsx", + "ts": "typescript", + "tsx": "tsx", + "typescript": "typescript", + # Shell variants + "sh": "bash", + "shell": "bash", + "console": "bash", + "terminal": "bash", + # Web languages + "html": "html", + "css": "css", + "json": "json", + "xml": "xml", + "yml": "yaml", + "yaml": "yaml", + # C family + "c": "c", + "cpp": "cpp", + "c++": "cpp", + "csharp": "csharp", + "cs": "csharp", + "java": "java", + # Other common languages + "go": "go", + "golang": "go", + "ruby": "ruby", + "rb": "ruby", + "rust": "rust", + "php": "php", + "sql": "sql", + "diff": "diff", + "markdown": "markdown", + "md": "markdown", + # Default fallback + "text": "text", + "plaintext": "text", + "txt": "text", + } + + # Return mapped language or default to the original if not in map + return lang_map.get(lang, lang or "text") + def _create_tool_panel_content(tool_name, args, output, execution_info=None, token_info=None): """Create the header and content for a tool output panel.""" from rich.text import Text @@ -1870,6 +2174,9 @@ def _create_tool_panel_content(tool_name, args, output, execution_info=None, tok from rich.console import Group from rich.box import ROUNDED + # Check if this is a handoff (transfer to another agent) + is_handoff = tool_name.startswith("transfer_to_") + # Format arguments for display, passing tool_name for specific formatting args_str = _format_tool_args(args, tool_name=tool_name) @@ -1878,10 +2185,40 @@ def _create_tool_panel_content(tool_name, args, output, execution_info=None, tok # Create header header = Text() - header.append(tool_name, style="#00BCD4") - header.append("(", style="yellow") - header.append(args_str, style="yellow") - header.append(")", style="yellow") + if is_handoff: + # Extract agent name from transfer function name + agent_name = None + if tool_name.startswith("transfer_to_"): + # Remove 'transfer_to_' prefix and convert to a nicer format + agent_name_raw = tool_name[len("transfer_to_"):] + # Convert underscores to spaces and capitalize words + agent_name = " ".join(word.capitalize() for word in agent_name_raw.split("_")) + + # Special case for acronyms like DNS or SMTP that might be in the agent name + # Convert words that are all uppercase to remain uppercase + parts = agent_name.split() + for i, part in enumerate(parts): + if part.upper() == part and len(part) > 1: # It's an acronym + parts[i] = part.upper() + agent_name = " ".join(parts) + + # For handoffs, show "transfer_to_X → Agent Name" + header.append(tool_name, style="#00BCD4") + if agent_name: + header.append(" → ", style="bold yellow") + header.append(agent_name, style="bold green") + + # Add arguments if present + if args_str: + header.append("(", style="yellow") + header.append(args_str, style="yellow") + header.append(")", style="yellow") + else: + # For regular tools, use the original format + header.append(tool_name, style="#00BCD4") + header.append("(", style="yellow") + header.append(args_str, style="yellow") + header.append(")", style="yellow") # Add timing information if timing_info: @@ -2046,61 +2383,6 @@ def _create_tool_panel_content(tool_name, args, output, execution_info=None, tok return header, Group(*group_content) -# Helper function to format tool arguments -def _format_tool_args(args, tool_name=None): - """Format tool arguments as a clean string.""" - # If the tool is execute_code, we don't want to show any args in the main header, - # as they are detailed in subsequent panels (either code or args string). - if tool_name == "execute_code": - return "" - - # If args is already a string, it might be pre-formatted or a simple arg string - if isinstance(args, str): - # If it looks like a JSON dict string, try to parse and format nicely - if args.strip().startswith('{') and args.strip().endswith('}'): - try: - - parsed_dict = json.loads(args) - # Recursively call with the parsed dict for consistent formatting - return _format_tool_args(parsed_dict, tool_name=tool_name) - except json.JSONDecodeError: - # Not valid JSON, or not a dict; return as is - return args - else: - # Simple string arg, return as is - return args - - # Format arguments from a dictionary - if isinstance(args, dict): - # Only include non-empty values and exclude special flags - arg_parts = [] - for key, value in args.items(): - # For execute_code, if the 'code' key is present, its content is shown in a dedicated panel. - # So, skip adding 'code=...' to the header string to avoid redundancy and verbosity. - if tool_name == "execute_code" and key == "code": - continue - - # Skip empty values - if value == "" or value == {} or value is None: - continue - # Skip special flags - if key in ["async_mode", "streaming"] and not value: - continue - - value_str = str(value) - - # Format the value - if isinstance(value, str): - # Truncate long string values that are not specifically handled above - if len(value_str) > 70 and key not in ["code", "args"]: - value_str = value_str[:67] + "..." - arg_parts.append(f"{key}={value_str}") - else: - arg_parts.append(f"{key}={value_str}") - return ", ".join(arg_parts) - else: - return str(args) - # Helper function to get timing information def _get_timing_info(execution_info=None): """Get timing information for display.""" @@ -2126,6 +2408,7 @@ def _get_timing_info(execution_info=None): timing_info.append(f"Tool: {format_time(tool_time)}") return timing_info, tool_time + # Helper function to create token info display def _create_token_info_display(token_info=None): """Create token information display text.""" @@ -2361,6 +2644,28 @@ def finish_tool_streaming(tool_name, args, output, call_id, execution_info=None, if 'start_time' in session and 'tool_time' not in execution_info: execution_info["tool_time"] = time.time() - session['start_time'] + # Add compact token info for display + if token_info: + # Create compact token representation + input_tokens = token_info.get('interaction_input_tokens', 0) + output_tokens = token_info.get('interaction_output_tokens', 0) + interaction_cost = token_info.get('interaction_cost', 0) + + # Calculate cost if not provided + if not interaction_cost and input_tokens > 0: + model_name = token_info.get('model', os.environ.get('CAI_MODEL', 'qwen2.5:14b')) + interaction_cost = calculate_model_cost(model_name, input_tokens, output_tokens) + + # Add compact token info to output + if input_tokens > 0: + compact_tokens = f"\n[Tokens: I:{input_tokens} O:{output_tokens} | Cost: ${interaction_cost:.4f}]" + if output: + if not output.endswith("\n"): + output += "\n" + output += compact_tokens + else: + output = compact_tokens + # Show the final output cli_print_tool_output( tool_name=tool_name, @@ -2375,152 +2680,3 @@ def finish_tool_streaming(tool_name, args, output, call_id, execution_info=None, # Mark the streaming session as complete if hasattr(cli_print_tool_output, '_streaming_sessions') and call_id in cli_print_tool_output._streaming_sessions: cli_print_tool_output._streaming_sessions[call_id]['is_complete'] = True - -def print_message_history(messages, title="Message History"): - """ - Pretty-print a sequence of messages with enhanced debug information. - - Args: - messages (List[dict]): List of message dictionaries to display - title (str, optional): Title to display above the message history - """ - from rich.console import Console - from rich.panel import Panel - from rich.text import Text - from rich.table import Table - - console = Console() - - # Create a table for displaying messages - table = Table(show_header=True, header_style="bold magenta", expand=True) - table.add_column("#", style="dim", width=3) - table.add_column("Role", style="cyan", width=10) - table.add_column("Content", width=40) - table.add_column("Metadata", width=30) - - # Process each message - for i, msg in enumerate(messages): - # Get role with color based on type - role = msg.get("role", "unknown") - role_style = { - "user": "green", - "assistant": "blue", - "system": "yellow", - "tool": "magenta" - }.get(role, "white") - - # Get content preview - content = msg.get("content") - content_preview = "" - if content is None: - content_preview = "[dim]None[/dim]" - elif isinstance(content, str): - # Truncate and escape long content - content_preview = (content[:37] + "...") if len(content) > 40 else content - content_preview = content_preview.replace("\n", "\\n") - elif isinstance(content, list): - content_preview = f"[list with {len(content)} items]" - else: - content_preview = f"[{type(content).__name__}]" - - # Gather metadata - metadata = [] - if msg.get("tool_calls"): - tc_count = len(msg["tool_calls"]) - tc_info = [] - for tc in msg["tool_calls"]: - tc_id = tc.get("id", "unknown") - tc_name = tc.get("function", {}).get("name", "unknown") if "function" in tc else "unknown" - tc_info.append(f"{tc_name}({tc_id})") - metadata.append(f"tool_calls[{tc_count}]: {', '.join(tc_info)}") - - if msg.get("tool_call_id"): - metadata.append(f"tool_call_id: {msg['tool_call_id']}") - - metadata_str = ", ".join(metadata) - - # Add row to table - table.add_row( - str(i), - f"[{role_style}]{role}[/{role_style}]", - content_preview, - metadata_str - ) - - # Create the panel with the table - panel = Panel( - table, - title=f"[bold]{title}[/bold]", - expand=False - ) - - # Display the panel - console.print(panel) - - return len(messages) # Return message count for convenience - -def get_language_from_code_block(lang_identifier): - """ - Maps a language identifier from a markdown code block to a proper syntax - highlighting language name. Handles common aliases and defaults. - - Args: - lang_identifier (str): Language identifier from markdown code block - - Returns: - str: Proper language name for syntax highlighting - """ - # Convert to lowercase and strip whitespace - lang = lang_identifier.lower().strip() if lang_identifier else "" - - # Map common language aliases to their proper names - lang_map = { - # Empty strings or unknown - "": "text", - # Python variants - "py": "python", - "python3": "python", - # JavaScript variants - "js": "javascript", - "jsx": "jsx", - "ts": "typescript", - "tsx": "tsx", - "typescript": "typescript", - # Shell variants - "sh": "bash", - "shell": "bash", - "console": "bash", - "terminal": "bash", - # Web languages - "html": "html", - "css": "css", - "json": "json", - "xml": "xml", - "yml": "yaml", - "yaml": "yaml", - # C family - "c": "c", - "cpp": "cpp", - "c++": "cpp", - "csharp": "csharp", - "cs": "csharp", - "java": "java", - # Other common languages - "go": "go", - "golang": "go", - "ruby": "ruby", - "rb": "ruby", - "rust": "rust", - "php": "php", - "sql": "sql", - "diff": "diff", - "markdown": "markdown", - "md": "markdown", - # Default fallback - "text": "text", - "plaintext": "text", - "txt": "text", - } - - # Return mapped language or default to the original if not in map - return lang_map.get(lang, lang or "text") From 857f19d396d2e1a5841e50f315f26fd69dcd2b4c Mon Sep 17 00:00:00 2001 From: luijait Date: Sun, 11 May 2025 21:44:04 +0200 Subject: [PATCH 2/2] Details --- .../agents/models/openai_chatcompletions.py | 10 ++++- src/cai/util.py | 43 ++++++++++++++++++- 2 files changed, 50 insertions(+), 3 deletions(-) diff --git a/src/cai/sdk/agents/models/openai_chatcompletions.py b/src/cai/sdk/agents/models/openai_chatcompletions.py index afd1ffdf..a61ab895 100644 --- a/src/cai/sdk/agents/models/openai_chatcompletions.py +++ b/src/cai/sdk/agents/models/openai_chatcompletions.py @@ -14,7 +14,7 @@ import asyncio from collections.abc import AsyncIterator, Iterable from dataclasses import dataclass, field from typing import TYPE_CHECKING, Any, Literal, cast, overload -from cai.util import get_ollama_api_base, fix_message_list, cli_print_agent_messages, create_agent_streaming_context, update_agent_streaming_content, finish_agent_streaming +from cai.util import get_ollama_api_base, fix_message_list, cli_print_agent_messages, create_agent_streaming_context, update_agent_streaming_content, finish_agent_streaming, calculate_model_cost from cai.util import start_idle_timer, stop_idle_timer, start_active_timer, stop_active_timer from wasabi import color from cai.sdk.agents.run_to_jsonl import get_session_recorder @@ -778,7 +778,13 @@ class OpenAIChatCompletionsModel(Model): # Update streaming display if enabled - always do this for text content if streaming_context: - update_agent_streaming_content(streaming_context, content) + # Create token stats to pass with each update + token_stats = { + 'input_tokens': estimated_input_tokens, + 'output_tokens': estimated_output_tokens, + 'cost': calculate_model_cost(str(self.model), estimated_input_tokens, estimated_output_tokens) + } + update_agent_streaming_content(streaming_context, content, token_stats) # More accurate token counting for text content output_text += content diff --git a/src/cai/util.py b/src/cai/util.py index 094b177c..a33cb05d 100644 --- a/src/cai/util.py +++ b/src/cai/util.py @@ -1480,13 +1480,14 @@ def create_agent_streaming_context(agent_name, counter, model): print(f"Error creating streaming context: {e}", file=sys.stderr) return None -def update_agent_streaming_content(context, text_delta): +def update_agent_streaming_content(context, text_delta, token_stats=None): """ Update the streaming content with new text. Args: context: The streaming context created by create_agent_streaming_context text_delta: The new text to add + token_stats: Optional token statistics to show with each update """ if not context: return False @@ -1502,6 +1503,46 @@ def update_agent_streaming_content(context, text_delta): # Add the parsed text to the content context["content"].append(parsed_delta) + # Update the footer with token stats if provided + if token_stats: + # Create token stats display + from rich.text import Text + footer_stats = Text() + + # Add timestamp and model info + footer_stats.append(f"\n[{context['timestamp']}", style="dim") + if context['model']: + footer_stats.append(f" ({context['model']})", style="bold magenta") + footer_stats.append("]", style="dim") + + # Add token stats + input_tokens = token_stats.get('input_tokens', 0) + output_tokens = token_stats.get('output_tokens', 0) + total_cost = token_stats.get('cost', 0.0) + + if input_tokens > 0: + footer_stats.append(" | ", style="dim") + footer_stats.append(f"I:{input_tokens} O:{output_tokens}", style="green") + if total_cost > 0: + footer_stats.append(f" (${total_cost:.4f})", style="bold cyan") + + # Add context usage indicator + model_name = context.get("model", os.environ.get('CAI_MODEL', 'qwen2.5:14b')) + context_pct = input_tokens / get_model_input_tokens(model_name) * 100 + if context_pct < 50: + indicator = "🟩" + color = "green" + elif context_pct < 80: + indicator = "🟨" + color = "yellow" + else: + indicator = "🟥" + color = "red" + footer_stats.append(f" {indicator} {context_pct:.1f}%", style=f"bold {color}") + + # Update the footer + context["footer"] = footer_stats + # Update the live display with the latest content updated_panel = Panel( Text.assemble(context["header"], context["content"], context["footer"]),