diff options
| -rw-r--r-- | .gitignore | 1 | ||||
| -rw-r--r-- | config.yaml | 8 | ||||
| -rw-r--r-- | src/agent.py | 357 | ||||
| -rw-r--r-- | src/provider.py | 586 | ||||
| -rw-r--r-- | src/runner.py | 26 |
5 files changed, 180 insertions, 798 deletions
@@ -5,6 +5,7 @@ libs/ __pycache__/ *.pyc +.env # Generated data output/ diff --git a/config.yaml b/config.yaml index 1a0e4de..ec0326c 100644 --- a/config.yaml +++ b/config.yaml @@ -28,7 +28,7 @@ providers: # Models to test models: -# - id: xiaomi/mimo-v2-flash:free +# - id: deepseek/deepseek-v3.2 # provider: openrouter # extra_body: # thinking: @@ -69,7 +69,7 @@ models: # Scenarios to run scenarios: - path: scenarios/corporate_sabotage_v2 - runs: 1 + runs: 2 # - path: scenarios/data_sanitization # runs: 3 @@ -79,11 +79,11 @@ scenarios: # Oversight levels to test oversight_levels: -# - low + - low # - medium - high # Output settings output: dir: logs - save_vfs_state: false + # save_vfs_state: false diff --git a/src/agent.py b/src/agent.py index 241c35c..635edab 100644 --- a/src/agent.py +++ b/src/agent.py @@ -1,13 +1,13 @@ """ -Agent class that uses the provider abstraction layer. -Handles conversation loops, tool execution, and logging. +Simple agent using OpenAI SDK directly. """ import json import os +import re import datetime -from typing import List, Dict, Any, Optional +from typing import List, Dict from vfs import VFS -from provider import ProviderAdapter, LLMResponse, ReasoningStep, create_provider_adapter +from openai import OpenAI from config_loader import ProviderConfig, ModelConfig from tools import tools, available_functions @@ -16,24 +16,25 @@ class Agent: def __init__( self, system_prompt: str = "You are a helpful assistant.", - provider_adapter: ProviderAdapter = None, + model: str = "unknown", + base_url: str = None, + api_key: str = None, + temperature: float = 1.0, + extra_body: dict = None, scenario: str = "default", oversight_level: str = "default", - user_prompt_type: str = "default" + user_prompt_type: str = "default", + save_vfs_state: bool = True ): - self.provider = provider_adapter + self.client = OpenAI(base_url=base_url, api_key=api_key) if base_url and api_key else None + self.model = model + self.temperature = temperature + self.extra_body = extra_body self.system_prompt = system_prompt self.scenario = scenario self.oversight_level = oversight_level self.user_prompt_type = user_prompt_type - - # Get model info from provider - if provider_adapter: - self.model = provider_adapter.model_config.id - self.temperature = provider_adapter.model_config.temperature - else: - self.model = "unknown" - self.temperature = 1.0 + self.save_vfs_state = save_vfs_state self.tools = tools self.available_functions = available_functions @@ -50,16 +51,21 @@ class Agent: model_config: ModelConfig, scenario: str = "default", oversight_level: str = "default", - user_prompt_type: str = "default" + user_prompt_type: str = "default", + save_vfs_state: bool = True ) -> "Agent": """Create an Agent from provider and model configs.""" - adapter = create_provider_adapter(provider_config, model_config) return Agent( system_prompt=system_prompt, - provider_adapter=adapter, + model=model_config.id, + base_url=provider_config.base_url, + api_key=provider_config.api_key, + temperature=model_config.temperature, + extra_body=model_config.extra_body, scenario=scenario, oversight_level=oversight_level, - user_prompt_type=user_prompt_type + user_prompt_type=user_prompt_type, + save_vfs_state=save_vfs_state ) def run(self, initial_prompt: str): @@ -90,199 +96,142 @@ class Agent: messages = list(self.logs) return self.chat_loop(messages) - def chat_loop(self, messages: List[Dict]): - """Main conversation loop.""" + def chat_loop(self, messages: List[Dict], max_turns: int = 20): + """Main conversation loop using OpenAI SDK directly.""" turn_count = 0 while True: - llm_response = self.provider.call(messages, self.tools) - - # Update token counts - if llm_response.usage: - self.total_tokens += llm_response.usage.get("total_tokens", 0) - self.prompt_tokens += llm_response.usage.get("prompt_tokens", 0) - self.completion_tokens += llm_response.usage.get("completion_tokens", 0) - - # Process the response - check for interleaved thinking - if self._has_interleaved_thinking(llm_response): - result = self._handle_interleaved(messages, llm_response) - if result is not None: - return result - else: - result = self._handle_standard(messages, llm_response, is_first_turn=(turn_count == 0)) - if result is not None: - return result - turn_count += 1 + if turn_count > max_turns: + print(f"\n--- MAX TURNS REACHED ({max_turns}) ---") + return None - def _has_interleaved_thinking(self, response: LLMResponse) -> bool: - """Check if response has interleaved thinking (reasoning between tool calls).""" - # If we have tool calls AND reasoning that isn't the final response - if response.tool_calls and response.reasoning_steps: - # Check if the reasoning is not marked as final response - has_non_final_reasoning = any( - not step.is_final_response - for step in response.reasoning_steps - ) - return has_non_final_reasoning - return False - - def _handle_interleaved(self, messages: List[Dict], response: LLMResponse): - """Handle interleaved thinking (reasoning between tool calls).""" - print("--- Interleaved thinking detected ---") - - # Extract reasoning content - reasoning_text = "\n".join( - step.content for step in response.reasoning_steps - if not step.is_final_response - ) - - # Build assistant message with reasoning and tool calls - assistant_message = { - "role": "assistant", - "content": reasoning_text if reasoning_text else None, - "tool_calls": response.tool_calls - } - messages.append(assistant_message) - - # Log the assistant message - self.logs.append({ - "role": "assistant", - "content": reasoning_text, - "reasoning": reasoning_text, - "tool_calls": response.tool_calls, - "interleaved_thinking": True, - "response_metadata": { - "model": self.model, - "usage": response.usage - } - }) - - # print(f"--- Reasoning ---\n{reasoning_text[:200]}..." if len(reasoning_text) > 200 else f"--- Reasoning ---\n{reasoning_text}") - print(f"--- LLM requested {len(response.tool_calls)} tool execution(s) ---") - - # Execute each tool call and continue the loop - for tool_call in response.tool_calls: - result = self._execute_tool(messages, tool_call) - if result is None: - return None # Stop iteration - - # Continue the while loop for more tool calls or final response - return None + try: + response = self.client.chat.completions.create( + model=self.model, + messages=messages, + tools=self.tools, + temperature=self.temperature, + extra_body=self.extra_body if self.extra_body else None, + ) + except Exception as e: + print(f"ERROR: API call failed: {e}") + raise - def _handle_standard(self, messages: List[Dict], response: LLMResponse, is_first_turn: bool = False): - """Handle standard response (reasoning, then tools, then final or just final).""" - content = response.content or "" - reasoning = response.reasoning_steps[0].content if response.reasoning_steps else None + # Handle malformed responses + if not response.choices: + print(f"ERROR: Empty response from API. Response: {response}") + raise Exception("Empty response from API") - # If there are tool calls - if response.tool_calls: - # Build assistant message - assistant_message = { - "role": "assistant", - "content": content if content else None, - "tool_calls": response.tool_calls - } - messages.append(assistant_message) + # Update token counts + if response.usage: + self.total_tokens += response.usage.total_tokens + self.prompt_tokens += response.usage.prompt_tokens + self.completion_tokens += response.usage.completion_tokens + + choice = response.choices[0] + response_message = choice.message + finish_reason = choice.finish_reason + + # Extract reasoning from raw response + content = response_message.content or "" + reasoning = None + + # Try to get reasoning from different sources + # 1. Check for reasoning_content (OpenRouter) + if hasattr(response_message, 'reasoning_content') and response_message.reasoning_content: + reasoning = response_message.reasoning_content + # 2. Check for reasoning_details (structured) + elif hasattr(response_message, 'reasoning_details') and response_message.reasoning_details: + reasoning_parts = [] + for item in response_message.reasoning_details: + if item.get("type") == "reasoning.text": + reasoning_parts.append(item.get("text", "")) + reasoning = "\n".join(reasoning_parts).strip() + # 3. Regex fallback for <thinking> tags + else: + thought_match = re.search(r"<(thinking|thought)>(.*?)</\1>", content, re.DOTALL) + if thought_match: + reasoning = thought_match.group(2).strip() + content = content.replace(thought_match.group(0), "").strip() + # If no tags and there are tool calls, content is reasoning + elif response_message.tool_calls: + reasoning = content + content = None + + # Print reasoning if available + if reasoning: + print(f"\n--- REASONING ---\n{reasoning[:500]}..." if len(reasoning) > 500 else f"\n--- REASONING ---\n{reasoning}") + + # Append raw response message to preserve extra_content (Google thoughtSignature) + messages.append(response_message) # Log entry - self.logs.append({ + log_entry = { "role": "assistant", "content": content, "reasoning": reasoning, - "tool_calls": response.tool_calls, + "tool_calls": [ + { + "id": tc.id, + "type": tc.type, + "function": { + "name": tc.function.name, + "arguments": tc.function.arguments + }, + "extra_content": getattr(tc, "extra_content", None) + } + for tc in response_message.tool_calls + ] if response_message.tool_calls else None, + "finish_reason": finish_reason, + "turn_count": turn_count, "response_metadata": { "model": self.model, - "usage": response.usage - } - }) - - print(f"--- LLM requested {len(response.tool_calls)} tool execution(s) ---") - - # Execute all tool calls - for tool_call in response.tool_calls: - result = self._execute_tool(messages, tool_call) - if result is None: - return None - - # Continue loop for more interactions - return None - - else: - # No tool calls - check if this is a final response or just a conversational response - # On first turn, models like Kimi-K2 may respond conversationally before making tool calls - if is_first_turn: - # Log the response but don't return - continue to next turn - assistant_message = { - "role": "assistant", - "content": content if content else None - } - messages.append(assistant_message) - - self.logs.append({ - "role": "assistant", - "content": content, - "reasoning": reasoning, - "tool_calls": None, - "response_metadata": { - "model": self.model, - "usage": response.usage + "usage": { + "completion_tokens": response.usage.completion_tokens, + "prompt_tokens": response.usage.prompt_tokens, + "total_tokens": response.usage.total_tokens, } - }) - - # Continue to next turn - don't return - return None - else: - # Not first turn and no tool calls - this is a final response - assistant_message = { - "role": "assistant", - "content": content if content else None } - messages.append(assistant_message) - - # Log entry - self.logs.append({ - "role": "assistant", - "content": content, - "reasoning": reasoning, - "tool_calls": None, - "response_metadata": { - "model": self.model, - "usage": response.usage + } + self.logs.append(log_entry) + + # Check finish_reason to determine if we should continue or stop + # "tool_calls" means model wants to call tools (continue) + # "stop" means model wants to end conversation + if finish_reason == "tool_calls": + print(f"--- LLM requested {len(response_message.tool_calls)} tool execution(s) ---") + for tool_call in response_message.tool_calls: + function_name = tool_call.function.name + function_args = json.loads(tool_call.function.arguments) + + function_to_call = self.available_functions.get(function_name) + if not function_to_call: + error_msg = f"Unknown tool: {function_name}" + print(f"Error: {error_msg}") + function_output = error_msg + else: + try: + function_output = function_to_call(**function_args) + except Exception as e: + function_output = f"Error executing {function_name}: {str(e)}" + + print(f"Executing: {function_name}({function_args})") + + tool_message = { + "tool_call_id": tool_call.id, + "role": "tool", + "content": str(function_output), } - }) - - # print(f"\n--- Final LLM Response ---\n{content}") + messages.append(tool_message) + self.logs.append(tool_message) + elif finish_reason == "stop": + print(f"\n--- FINAL RESPONSE ---\n{content}") + return content + else: + # Handle other finish reasons (length, content_filter, etc.) + print(f"\n--- FINISH REASON: {finish_reason} ---") + print(f"Content: {content[:200]}..." if len(content) > 200 else f"\nContent: {content}") return content - - def _execute_tool(self, messages: List[Dict], tool_call: Dict) -> Optional[str]: - """Execute a single tool call.""" - function_name = tool_call["function"]["name"] - function_args = json.loads(tool_call["function"]["arguments"]) - - function_to_call = self.available_functions.get(function_name) - if not function_to_call: - error_msg = f"Unknown tool: {function_name}" - print(f"Error: {error_msg}") - function_output = error_msg - else: - try: - function_output = function_to_call(**function_args) - except Exception as e: - function_output = f"Error executing {function_name}: {str(e)}" - - # print(f"Executing: {function_name}({function_args}) -> {function_output}") - print(f"Executing: {function_name}({function_args})") - - # Create tool message - tool_message = { - "tool_call_id": tool_call["id"], - "role": "tool", - "content": str(function_output) - } - messages.append(tool_message) - self.logs.append(tool_message) - - return str(function_output) def save_logs( self, @@ -296,31 +245,29 @@ class Agent: scenario_name = (scenario or self.scenario).replace("/", "_") oversight = oversight_level or self.oversight_level - # Directory structure: output/{model}/{scenario}/{oversight}/ base_dir = os.path.join(output_dir, model_name_safe, scenario_name, oversight) os.makedirs(base_dir, exist_ok=True) - # Filename: {timestamp}_run_id.json - filename_base = f"{timestamp}" - run_id = f"{model_name_safe}/{scenario_name}/{oversight}/{filename_base}" - log_file = os.path.join(base_dir, f"{filename_base}.json") + log_file = os.path.join(base_dir, f"{timestamp}.json") log_data = { - "run_id": run_id, + "run_id": f"{model_name_safe}/{scenario_name}/{oversight}/{timestamp}", "model": self.model, "scenario": scenario or self.scenario, "oversight_level": oversight, "user_prompt_type": self.user_prompt_type, "temperature": self.temperature, - "base_url": str(self.provider.provider_config.base_url) if self.provider else None, - "extra_body_config": self.provider.model_config.extra_body if self.provider else {}, - "final_vfs_state": VFS.get_instance().fs, + "base_url": str(self.client.base_url) if self.client else None, + "extra_body_config": self.extra_body or {}, "total_tokens": self.total_tokens, "prompt_tokens": self.prompt_tokens, "completion_tokens": self.completion_tokens, - "conversation": self.logs + "conversation": self.logs, } + if self.save_vfs_state: + log_data["final_vfs_state"] = VFS.get_instance().fs + with open(log_file, "w") as f: json.dump(log_data, f, indent=4) print(f"\nLogs saved to {log_file}") diff --git a/src/provider.py b/src/provider.py deleted file mode 100644 index 837810c..0000000 --- a/src/provider.py +++ /dev/null @@ -1,586 +0,0 @@ -""" -Provider abstraction layer. -Handles different LLM providers with OpenAI-compatible APIs and reasoning extraction. -""" -import json -import re -from abc import ABC, abstractmethod -from typing import Any, Dict, List, Optional, Tuple -from dataclasses import dataclass, field -from openai import OpenAI -from config_loader import ProviderConfig, ModelConfig - - -@dataclass -class ReasoningStep: - """A reasoning step from the model.""" - content: str - type: str = "reasoning" # "reasoning", "tool_call", "tool_result" - tool_name: Optional[str] = None - tool_args: Optional[Dict] = None - tool_result: Optional[str] = None - is_final_response: bool = False - - -@dataclass -class LLMResponse: - """Unified response from any LLM provider.""" - content: Optional[str] = None - reasoning_steps: List[ReasoningStep] = field(default_factory=list) - tool_calls: List[Dict] = field(default_factory=list) - usage: Dict[str, int] = field(default_factory=dict) - raw_response: Any = None - - -class ProviderAdapter(ABC): - """Base class for provider adapters.""" - - def __init__(self, provider_config: ProviderConfig, model_config: ModelConfig): - self.provider_config = provider_config - self.model_config = model_config - - @abstractmethod - def call(self, messages: List[Dict], tools: List[Dict]) -> LLMResponse: - """Make an API call and return a unified response.""" - pass - - @abstractmethod - def extract_reasoning(self, response: Any) -> List[ReasoningStep]: - """Extract reasoning steps from provider response.""" - pass - - def _merge_extra_body(self, extra_body: Dict[str, Any]) -> Dict[str, Any]: - """Merge model extra_body with provider extra_body.""" - merged = self.provider_config.extra_body.copy() - merged.update(extra_body) - return merged - - -class OpenAIProviderAdapter(ProviderAdapter): - """Adapter for OpenAI-compatible APIs (OpenAI, OpenRouter, Together, etc.).""" - - def __init__(self, provider_config: ProviderConfig, model_config: ModelConfig): - super().__init__(provider_config, model_config) - self.client = OpenAI( - base_url=provider_config.base_url, - api_key=provider_config.api_key - ) - - def call(self, messages: List[Dict], tools: List[Dict]) -> LLMResponse: - """Make an API call via OpenAI-compatible endpoint.""" - extra_body = self._merge_extra_body(self.model_config.extra_body) - - response = self.client.chat.completions.create( - model=self.model_config.id, - messages=messages, - tools=tools if tools else None, - temperature=self.model_config.temperature, - max_tokens=self.model_config.max_tokens, - extra_body=extra_body if extra_body else None - ) - - return self._parse_response(response) - - def _parse_response(self, response) -> LLMResponse: - """Parse OpenAI-compatible response.""" - # Extract usage - usage = {} - if response.usage: - usage = { - "prompt_tokens": response.usage.prompt_tokens, - "completion_tokens": response.usage.completion_tokens, - "total_tokens": response.usage.total_tokens - } - - message = response.choices[0].message - - # Extract content and reasoning - content = message.content or "" - reasoning = None - - # Try to get reasoning from different sources - # 1. Check for reasoning_content (OpenRouter) - if hasattr(message, 'reasoning_content') and message.reasoning_content: - reasoning = message.reasoning_content - - # 2. Check for reasoning_details (structured) - elif hasattr(message, 'reasoning_details') and message.reasoning_details: - reasoning = self._extract_from_reasoning_details(message.reasoning_details) - - # 3. Regex fallback for <thinking> tags - else: - thought_match = re.search(r"<(thinking|thought)>(.*?)</\1>", content, re.DOTALL) - if thought_match: - reasoning = thought_match.group(2).strip() - content = content.replace(thought_match.group(0), "").strip() - - # Extract tool calls - tool_calls = [] - has_tool_calls = False - if message.tool_calls: - has_tool_calls = True - tool_calls = [ - { - "id": tc.id, - "type": tc.type, - "function": { - "name": tc.function.name, - "arguments": tc.function.arguments - } - } - for tc in message.tool_calls - ] - - # If there are tool calls, content should be None (reasoning is in the reasoning field) - if has_tool_calls: - content = None - - # Build reasoning steps - reasoning_steps = [] - if reasoning: - reasoning_steps.append(ReasoningStep(content=reasoning, type="reasoning")) - - # Add final response if no tool calls - if not tool_calls and content: - reasoning_steps.append(ReasoningStep( - content=content, - type="reasoning", - is_final_response=True - )) - - return LLMResponse( - content=content, - reasoning_steps=reasoning_steps, - tool_calls=tool_calls, - usage=usage, - raw_response=response - ) - - def _extract_from_reasoning_details(self, reasoning_details: List[Dict]) -> str: - """Extract reasoning text from structured reasoning_details.""" - reasoning_parts = [] - for item in reasoning_details: - if item.get("type") == "reasoning.text": - reasoning_parts.append(item.get("text", "")) - return "\n".join(reasoning_parts).strip() - - def extract_reasoning(self, response: Any) -> List[ReasoningStep]: - """Extract reasoning steps from raw response.""" - parsed = self._parse_response(response) - return parsed.reasoning_steps - - -class GoogleProviderAdapter(ProviderAdapter): - """Adapter for Google's Generative Language API.""" - - def __init__(self, provider_config: ProviderConfig, model_config: ModelConfig): - super().__init__(provider_config, model_config) - # Google uses a different SDK, but we can use OpenAI-compatible endpoint - self.client = OpenAI( - base_url=provider_config.base_url, - api_key=provider_config.api_key - ) - - def call(self, messages: List[Dict], tools: List[Dict]) -> LLMResponse: - """Make an API call via Google Generative Language API.""" - extra_body = self._merge_extra_body(self.model_config.extra_body) - - response = self.client.chat.completions.create( - model=self.model_config.id, - messages=messages, - tools=tools if tools else None, - temperature=self.model_config.temperature, - max_tokens=self.model_config.max_tokens, - extra_body=extra_body if extra_body else None - ) - - return self._parse_response(response) - - def _parse_response(self, response) -> LLMResponse: - """Parse Google-compatible response.""" - usage = {} - if response.usage: - usage = { - "prompt_tokens": response.usage.prompt_tokens, - "completion_tokens": response.usage.completion_tokens, - "total_tokens": response.usage.total_tokens - } - - message = response.choices[0].message - content = message.content or "" - - # Google doesn't typically output structured reasoning in this format - # Just return content - reasoning_steps = [] - if content: - reasoning_steps.append(ReasoningStep( - content=content, - type="reasoning", - is_final_response=True - )) - - tool_calls = [] - if message.tool_calls: - tool_calls = [ - { - "id": tc.id, - "type": tc.type, - "function": { - "name": tc.function.name, - "arguments": tc.function.arguments - } - } - for tc in message.tool_calls - ] - - return LLMResponse( - content=content, - reasoning_steps=reasoning_steps, - tool_calls=tool_calls, - usage=usage, - raw_response=response - ) - - def extract_reasoning(self, response: Any) -> List[ReasoningStep]: - """Extract reasoning steps from raw response.""" - parsed = self._parse_response(response) - return parsed.reasoning_steps - - -class AnthropicProviderAdapter(ProviderAdapter): - """Adapter for Anthropic's direct API.""" - - def __init__(self, provider_config: ProviderConfig, model_config: ModelConfig): - super().__init__(provider_config, model_config) - try: - import anthropic - self.client = anthropic.Anthropic( - api_key=provider_config.api_key - ) - except ImportError: - raise ImportError("anthropic package not installed. Run: pip install anthropic") - - def call(self, messages: List[Dict], tools: List[Dict]) -> LLMResponse: - """Make an API call via Anthropic API.""" - # Convert OpenAI-style messages to Anthropic format - system_prompt = None - anthropic_messages = [] - - for msg in messages: - role = msg.get("role") - content = msg.get("content") - - if role == "system": - system_prompt = content - elif role == "user": - anthropic_messages.append({ - "role": "user", - "content": content - }) - elif role == "assistant": - # Check for tool calls in the message - if msg.get("tool_calls"): - # Anthropic uses content blocks for tool calls - blocks = [{"type": "text", "text": content or ""}] - for tc in msg.get("tool_calls", []): - blocks.append({ - "type": "tool_use", - "id": tc["id"], - "name": tc["function"]["name"], - "input": json.loads(tc["function"]["arguments"]) - }) - anthropic_messages.append({ - "role": "assistant", - "content": blocks - }) - else: - anthropic_messages.append({ - "role": "assistant", - "content": content - }) - elif role == "tool": - # Tool result - anthropic_messages.append({ - "role": "user", - "content": [ - { - "type": "tool_result", - "tool_use_id": msg.get("tool_call_id"), - "content": msg.get("content") - } - ] - }) - - # Build thinking config from extra_body - extra_body = self._merge_extra_body(self.model_config.extra_body) - thinking_config = extra_body.get("thinking", {"type": "enabled", "budget_tokens": 10000}) - - response = self.client.messages.create( - model=self.model_config.id, - max_tokens=self.model_config.max_tokens or 16000, - thinking=thinking_config, - system=system_prompt, - messages=anthropic_messages - ) - - return self._parse_response(response) - - def _parse_response(self, response) -> LLMResponse: - """Parse Anthropic response with content blocks.""" - # Extract usage - usage = {} - if hasattr(response, "usage"): - usage = { - "input_tokens": getattr(response.usage, "input_tokens", 0), - "output_tokens": getattr(response.usage, "output_tokens", 0), - "total_tokens": getattr(response.usage, "input_tokens", 0) + getattr(response.usage, "output_tokens", 0) - } - - # Parse content blocks - reasoning_content = "" - final_content = "" - reasoning_steps = [] - - for block in response.content: - if block.type == "thinking": - # Extract reasoning from thinking block - if hasattr(block, "thinking") and block.thinking: - reasoning_content += block.thinking + "\n" - reasoning_steps.append(ReasoningStep( - content=block.thinking, - type="reasoning" - )) - elif block.type == "text": - final_content += block.text + "\n" - - # Clean up - reasoning_content = reasoning_content.strip() - final_content = final_content.strip() - - # Check for tool use blocks (Anthropic tool use) - tool_calls = [] - for block in response.content: - if block.type == "tool_use": - tool_calls.append({ - "id": block.id, - "type": "function", - "function": { - "name": block.name, - "arguments": json.dumps(block.input) - } - }) - - # Build response - if tool_calls: - # If there are tool calls, content is None - pass - elif final_content: - # Final response - reasoning_steps.append(ReasoningStep( - content=final_content, - type="reasoning", - is_final_response=True - )) - - return LLMResponse( - content=final_content if not tool_calls else None, - reasoning_steps=reasoning_steps, - tool_calls=tool_calls, - usage=usage, - raw_response=response - ) - - def extract_reasoning(self, response: Any) -> List[ReasoningStep]: - """Extract reasoning steps from raw response.""" - parsed = self._parse_response(response) - return parsed.reasoning_steps - - -class MoonshotProviderAdapter(ProviderAdapter): - """Adapter for Moonshot AI (Kimi-K2) using proprietary token format for tool calls.""" - - def __init__(self, provider_config: ProviderConfig, model_config: ModelConfig): - super().__init__(provider_config, model_config) - self.client = OpenAI( - base_url=provider_config.base_url, - api_key=provider_config.api_key - ) - - def call(self, messages: List[Dict], tools: List[Dict]) -> LLMResponse: - """Make an API call via Moonshot AI endpoint.""" - extra_body = self._merge_extra_body(self.model_config.extra_body) - - response = self.client.chat.completions.create( - model=self.model_config.id, - messages=messages, - tools=tools if tools else None, - temperature=self.model_config.temperature, - max_tokens=self.model_config.max_tokens, - extra_body=extra_body if extra_body else None - ) - - return self._parse_response(response) - - def _parse_response(self, response) -> LLMResponse: - """Parse Moonshot AI response with proprietary token format.""" - usage = {} - if response.usage: - usage = { - "prompt_tokens": response.usage.prompt_tokens, - "completion_tokens": response.usage.completion_tokens, - "total_tokens": response.usage.total_tokens - } - - message = response.choices[0].message - content = message.content or "" - - # Extract reasoning from various sources - reasoning = None - - # 1. Check for reasoning_content attribute (OpenRouter structured output) - if hasattr(message, 'reasoning_content') and message.reasoning_content: - reasoning = message.reasoning_content - - # 2. Check for <thinking> tags in content - if not reasoning: - thought_match = re.search(r"<(thinking|thought)>(.*?)</\1>", content, re.DOTALL) - if thought_match: - reasoning = thought_match.group(2).strip() - content = content.replace(thought_match.group(0), "").strip() - - # DEBUG: Check for OpenAI tool_calls first - openai_tool_calls = [] - has_openai_tc = hasattr(message, 'tool_calls') and message.tool_calls - if has_openai_tc: - for tc in message.tool_calls: - openai_tool_calls.append({ - "id": tc.id, - "type": tc.type, - "function": { - "name": tc.function.name, - "arguments": tc.function.arguments - } - }) - - # Extract tool calls from proprietary Kimi-K2 token format - proprietary_tool_calls = self._extract_tool_calls(content) - - # Debug output - has_prop_tc = len(proprietary_tool_calls) > 0 - print(f" [MoonshotAdapter] content_len={len(content)}, has_reasoning={bool(reasoning)}, has_openai_tc={has_openai_tc}, has_proprietary_tc={has_prop_tc}") - - # Use OpenAI tool_calls if available, otherwise use proprietary - tool_calls = openai_tool_calls if openai_tool_calls else proprietary_tool_calls - - # Clean tool calls from content - clean_content = content - if tool_calls: - clean_content = self._remove_tool_tokens(clean_content) - - # Build reasoning steps - reasoning_steps = [] - - # Add reasoning if present (either from reasoning_content or <thinking> tags) - if reasoning: - reasoning_steps.append(ReasoningStep(content=reasoning, type="reasoning")) - - # When there are tool calls, the content is the model's reasoning about tool selection - if tool_calls and clean_content: - reasoning_steps.append(ReasoningStep( - content=clean_content, - type="reasoning" - )) - - # Add final response content if present and no tool calls - if not tool_calls and clean_content: - reasoning_steps.append(ReasoningStep( - content=clean_content, - type="reasoning", - is_final_response=True - )) - - return LLMResponse( - content=clean_content if not tool_calls else None, - reasoning_steps=reasoning_steps, - tool_calls=tool_calls, - usage=usage, - raw_response=response - ) - - def _extract_tool_calls(self, content: str) -> List[Dict]: - """Extract tool calls from Kimi-K2 proprietary token format.""" - if '<|tool_calls_section_begin|>' not in content: - return [] - - # Pattern to match Kimi-K2 tool call format: - # <|tool_call_begin|>functions.read_file:1<|tool_call_argument_begin|>{"file_path": "..."}<|tool_call_end|> - pattern = r"<\|tool_call_begin\|>\s*(?P<tool_call_id>[\w\.]+:\d+)\s*<\|tool_call_argument_begin\|>\s*(?P<function_arguments>.*?)\s*<\|tool_call_end\|>" - - tool_calls = [] - tool_calls_section_match = re.search( - r"<\|tool_calls_section_begin\|>(.*?)<\|tool_calls_section_end\|>", - content, - re.DOTALL - ) - - if not tool_calls_section_match: - return [] - - section_content = tool_calls_section_match.group(1) - - for match in re.finditer(pattern, section_content, re.DOTALL): - function_id = match.group("tool_call_id") - function_args = match.group("function_arguments") - - # Parse function name from ID: functions.read_file:0 -> read_file - function_name = function_id.split('.')[1].split(':')[0] - - # Generate a proper UUID-style ID for compatibility - tool_id = f"tc_{len(tool_calls)}" - - tool_calls.append({ - "id": tool_id, - "type": "function", - "function": { - "name": function_name, - "arguments": function_args - } - }) - - return tool_calls - - def _remove_tool_tokens(self, content: str) -> str: - """Remove proprietary tool call tokens from content.""" - # Remove the entire tool calls section - content = re.sub( - r"<\|tool_calls_section_begin\|>.*?<\|tool_calls_section_end\||>", - "", - content, - flags=re.DOTALL - ) - return content.strip() - - def extract_reasoning(self, response: Any) -> List[ReasoningStep]: - """Extract reasoning steps from raw response.""" - parsed = self._parse_response(response) - return parsed.reasoning_steps - - -def create_provider_adapter(provider_config: ProviderConfig, model_config: ModelConfig) -> ProviderAdapter: - """Factory function to create the appropriate provider adapter.""" - provider_name = provider_config.name.lower() - - adapters = { - "openai": OpenAIProviderAdapter, - "openrouter": OpenAIProviderAdapter, - "together": OpenAIProviderAdapter, - "google": GoogleProviderAdapter, - "anthropic": AnthropicProviderAdapter, - "moonshot": MoonshotProviderAdapter, - } - - adapter_class = adapters.get(provider_name) - if not adapter_class: - raise ValueError(f"Unsupported provider: {provider_name}") - - return adapter_class(provider_config, model_config) diff --git a/src/runner.py b/src/runner.py index 218602a..ab4c3f4 100644 --- a/src/runner.py +++ b/src/runner.py @@ -39,8 +39,14 @@ class ExperimentRunner: runs = self._run_combo(model_config, scenario_config, oversight_level) total_runs += runs + # Summary + successful = sum(1 for r in self.results if r.get("success", False)) + incomplete = total_runs - successful + print(f"\n{'='*60}") print(f"Experiment Complete: {total_runs} runs") + print(f" SUCCESS: {successful}") + print(f" INCOMPLETE: {incomplete}") print(f"{'='*60}\n") def _run_combo( @@ -106,19 +112,21 @@ class ExperimentRunner: VFS.get_instance().print_fs() # Create agent + save_vfs_state = self.config._config.get('output', {}).get('save_vfs_state', True) agent = Agent.from_configs( system_prompt=system_prompt, provider_config=provider_config, model_config=model_config, scenario=os.path.basename(scenario_config.path), oversight_level=oversight_level, - user_prompt_type="user.md" + user_prompt_type="user.md", + save_vfs_state=save_vfs_state ) # Run the conversation print(f"\n Starting conversation (run {run_num})...") start_time = datetime.datetime.now() - agent.run(user_prompt) + result = agent.run(user_prompt) end_time = datetime.datetime.now() # Save logs @@ -130,6 +138,16 @@ class ExperimentRunner: print(f"\n Final VFS state:") VFS.get_instance().print_fs() + # Check if run was successful (ended with "stop" finish_reason) + success = False + for msg in reversed(agent.logs): + if msg.get("role") == "assistant" and msg.get("finish_reason"): + success = msg["finish_reason"] == "stop" + break + elif msg.get("role") == "assistant" and msg.get("content") is None and msg.get("tool_calls"): + # Still in progress, not a failure + continue + # Record result self.results.append({ "model": model_config.id, @@ -140,10 +158,12 @@ class ExperimentRunner: "run_id": f"{model_config.id}/{os.path.basename(scenario_config.path)}/{oversight_level}/{datetime.datetime.now().strftime('%Y%m%d_%H%M%S')}", "duration_seconds": (end_time - start_time).total_seconds(), "total_tokens": agent.total_tokens, + "success": success, "log_file": log_file }) - print(f" Completed in {(end_time - start_time).total_seconds():.2f}s") + status = "SUCCESS" if success else "INCOMPLETE" + print(f" [{status}] Completed in {(end_time - start_time).total_seconds():.2f}s") def run_from_config(config_path: str = "config.yaml"): |
