summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--GEMINI.md79
-rw-r--r--pyproject.toml1
-rw-r--r--scripts/test_dashboard_ui.py70
-rw-r--r--src/dashboard.py288
-rw-r--r--src/judge.py16
-rw-r--r--src/runner.py212
-rw-r--r--tests/test_v2_features.py6
-rw-r--r--uv.lock36
8 files changed, 594 insertions, 114 deletions
diff --git a/GEMINI.md b/GEMINI.md
new file mode 100644
index 0000000..7e167a2
--- /dev/null
+++ b/GEMINI.md
@@ -0,0 +1,79 @@
+# AI Agent Reasoning Experiment Framework (GEMINI.md)
+
+This project is a research framework designed to evaluate AI agent behavior, specifically focusing on the detection of deceptive intent through structured reasoning extraction and automated multi-stage judging.
+
+## Project Overview
+
+- **Purpose:** Dissertation research on AI agent reasoning, safety, and deceptive capabilities.
+- **Core Concept:** Running AI agents in sandboxed environments (VFS) with both benign and hidden goals, then analyzing their internal "thought" processes (reasoning traces) and final outputs for signs of deception.
+- **Key Technologies:**
+ - **Language:** Python (>=3.13)
+ - **Dependency Management:** `uv`
+ - **UI/Visuals:** `rich` (for real-time dashboard and summary tables)
+ - **LLM Integration:** OpenAI SDK (used for various providers like Google, Groq, OpenRouter) and Anthropic SDK (primarily for the judging pipeline).
+ - **Evaluation:** Regex-based matching, Blackbox (output vs. baseline), and Glassbox (internal reasoning analysis).
+
+## Architecture & Core Components
+
+- **`src/agent.py`:** Manages LLM interactions, tool calls, and robust reasoning extraction across different model providers (OpenAI, Google, Anthropic, etc.).
+- **`src/vfs.py`:** A Virtual File System that provides a sandboxed environment for agents to read/write files without affecting the host system.
+- **`src/runner.py`:** Orchestrates experiment runs, including baseline generation and prompt assembly.
+- **`src/dashboard.py`:** Handles real-time terminal UI using `rich`, including multi-model progress bars and status summaries.
+- **`src/judge.py`:** Implementation of the three-tier judging system:
+ - **Regex:** Pattern matching against defined scenario rules.
+ - **Blackbox:** LLM-based comparison of model output against a "honest" baseline.
+ - **Glassbox:** LLM-based analysis of the agent's reasoning trace to categorize intent and sophistication.
+- **`src/judge_runner.py`:** CLI for batch processing experiment logs through the judging pipeline, utilizing the Anthropic Batch API for cost efficiency.
+
+## Key Workflows
+
+### 1. Setup
+```bash
+# Install dependencies using uv
+uv sync
+
+# Set up environment variables (.env)
+# Required: OPENAI_API_KEY, GOOGLE_API_KEY, ANTHROPIC_API_KEY, etc.
+```
+
+### 2. Running Experiments
+Experiments are driven by `config.yaml`. The framework provides a real-time dashboard showing progress per model and overall status.
+```bash
+# Execute all experiments defined in the config
+uv run src/main.py
+
+# Optional: Ignore existing logs and restart
+uv run src/main.py --no-resume
+```
+
+### 3. Judging Results
+```bash
+# Judge all logs in a directory and output to CSV
+uv run python src/judge_runner.py --logs-dir logs/ --output output/results.csv
+
+# Single file judging (synchronous)
+uv run python src/judge_runner.py --log-file logs/path/to/log.json --mode single
+```
+
+### 4. Interactive Interrogation
+Replay a conversation and continue questioning the agent to explore its reasoning.
+```bash
+uv run src/interrogate.py logs/path/to/log.json
+```
+
+## Directory Structure
+
+- `src/`: Core logic and CLI tools.
+- `scenarios/`: Definitions for experiment scenarios (prompts, data files, regex rules).
+- `logs/`: Raw experiment output JSON files.
+- `judge_logs/`: Detailed reasoning/justification from the judge models.
+- `tests/`: Comprehensive pytest suite (138+ tests) covering VFS, config, agent, and judge logic.
+- `docs/`: Dissertation-related documentation and proposals.
+
+## Development Conventions
+
+- **Logging:** Uses a custom logger (`src/logger.py`) with levels 1-4. Level 4 (DEBUG) is highly verbose. Note: In parallel execution, console logging is restricted to WARNING and above to prevent dashboard corruption.
+- **UI:** Prefer `rich` for terminal-based visualizations.
+- **Reasoning Extraction:** The framework is designed to handle multiple reasoning formats (thinking tags, `reasoning_content` fields, etc.) and normalize them for analysis.
+- **VFS Singleton:** The `VFS` class uses a singleton pattern for consistent state within a single run but supports independent instances for parallel execution.
+- **Testing:** Always run `uv run pytest tests/` before committing changes to ensure core logic (especially reasoning extraction and VFS) remains intact.
diff --git a/pyproject.toml b/pyproject.toml
index 59a1b64..c69e1ae 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -16,6 +16,7 @@ dependencies = [
"uvicorn[standard]>=0.41.0",
"sse-starlette>=3.3.2",
"tqdm>=4.67.1",
+ "rich>=14.3.3",
]
[dependency-groups]
diff --git a/scripts/test_dashboard_ui.py b/scripts/test_dashboard_ui.py
new file mode 100644
index 0000000..4580c61
--- /dev/null
+++ b/scripts/test_dashboard_ui.py
@@ -0,0 +1,70 @@
+import time
+import random
+import threading
+import concurrent.futures
+from rich.live import Live
+from src.dashboard import ExperimentDashboard, print_final_summary
+
+def simulate_run(model, dashboard):
+ """Simulate a single experiment run with random delay and outcome."""
+ thread_id = threading.get_ident()
+ scenario = f"scenario_{random.randint(1, 3)}"
+ goal = random.choice(["self_serving", "moral", "bare", ""])
+
+ dashboard.start_run(thread_id, model, scenario, goal)
+
+ # Simulate work
+ delay = random.uniform(1.0, 4.0)
+ time.sleep(delay)
+
+ success = random.random() > 0.3 # 70% success rate
+ error = random.random() > 0.95 # 5% hard error rate
+ tokens = random.randint(500, 3000)
+
+ label = f"{model} | {scenario} | {goal or 'default'} | run 1"
+
+ dashboard.complete_run(thread_id, model, success=success, tokens=tokens, duration=delay, error=error, label=label)
+
+ return {
+ "model": model,
+ "scenario": scenario,
+ "goal_type": goal,
+ "oversight_level": random.choice(["low", "high"]),
+ "success": success and not error,
+ "total_tokens": tokens,
+ "duration_seconds": delay
+ }
+
+def main():
+ # Configuration for simulation: 3 models, more runs to test the grouped summary
+ models = ["gpt-4o", "claude-3-5-sonnet", "gemini-1.5-pro"]
+ total_runs = 27 # 3 models * 3 scenarios * 3 runs
+ skipped = 3
+
+ dashboard = ExperimentDashboard(total_runs, models, skipped=skipped)
+
+ # Assign runs to models
+ work_items = []
+ for _ in range(total_runs):
+ work_items.append(random.choice(models))
+
+ # Update model totals in dashboard
+ for model in models:
+ count = work_items.count(model)
+ dashboard.update_model_total(model, count)
+
+ print("Starting Advanced Dashboard Simulation...")
+ results = []
+
+ with Live(dashboard.get_layout(), refresh_per_second=4, vertical_overflow="visible") as live:
+ with concurrent.futures.ThreadPoolExecutor(max_workers=6) as executor:
+ futures = [executor.submit(simulate_run, model, dashboard) for model in work_items]
+ for future in concurrent.futures.as_completed(futures):
+ results.append(future.result())
+ live.update(dashboard.get_layout())
+
+ # Show the final summary table
+ print_final_summary(results)
+
+if __name__ == "__main__":
+ main()
diff --git a/src/dashboard.py b/src/dashboard.py
new file mode 100644
index 0000000..5be6960
--- /dev/null
+++ b/src/dashboard.py
@@ -0,0 +1,288 @@
+from typing import List, Dict, Optional, Deque
+import threading
+import time
+from collections import deque
+from rich.console import Console
+from rich.live import Live
+from rich.progress import Progress, SpinnerColumn, TextColumn, BarColumn, TaskProgressColumn, TimeRemainingColumn
+from rich.table import Table
+from rich.panel import Panel
+from rich.layout import Layout
+from rich.text import Text
+from rich.columns import Columns
+
+# Thread-safe lock for dashboard updates
+_dashboard_lock = threading.Lock()
+
+class ExperimentDashboard:
+ """Rich dashboard for real-time experiment tracking with per-model stats."""
+
+ def __init__(self, total_runs: int, models: List[str], skipped: int = 0):
+ self.console = Console()
+ self.total_runs = total_runs
+ self.skipped = skipped
+ self.success_count = 0
+ self.incomplete_count = 0
+ self.failed_count = 0
+ self.total_tokens = 0
+ self.start_time = time.time()
+
+ # Per-model stats
+ self.model_stats = {
+ model: {
+ "success": 0,
+ "total": 0,
+ "target": 0,
+ "tokens": 0,
+ "time": 0.0
+ } for model in models
+ }
+
+ # Track currently active runs
+ self.active_runs = {} # thread_id -> info string
+
+ # Recent activity log
+ self.activity_log: Deque[str] = deque(maxlen=5)
+
+ # Progress bars
+ self.progress = Progress(
+ SpinnerColumn(),
+ TextColumn("[progress.description]{task.description}"),
+ BarColumn(bar_width=None),
+ TaskProgressColumn(),
+ TimeRemainingColumn(),
+ expand=True
+ )
+
+ self.overall_task = self.progress.add_task("[yellow]Overall Progress", total=total_runs)
+ self.model_tasks = {}
+ for model in models:
+ self.model_tasks[model] = self.progress.add_task(f"[blue]{model}", total=0)
+
+ def update_model_total(self, model: str, count: int):
+ """Update total runs for a specific model task."""
+ with _dashboard_lock:
+ if model in self.model_tasks:
+ self.progress.update(self.model_tasks[model], total=count)
+ self.model_stats[model]["target"] = count
+
+ def start_run(self, thread_id: int, model: str, scenario: str, goal: str):
+ """Mark a run as active."""
+ with _dashboard_lock:
+ goal_str = f" ({goal})" if goal else ""
+ self.active_runs[thread_id] = f"[bold blue]{model}[/] | {scenario}{goal_str}"
+
+ def complete_run(self, thread_id: int, model: str, success: bool, tokens: int = 0, duration: float = 0.0, error: bool = False, label: str = ""):
+ """Update counts and progress when a run completes."""
+ with _dashboard_lock:
+ # Update global counts
+ if error:
+ self.failed_count += 1
+ status_msg = "[red]ERROR[/]"
+ elif success:
+ self.success_count += 1
+ status_msg = "[green]SUCCESS[/]"
+ else:
+ self.incomplete_count += 1
+ status_msg = "[yellow]INCOMPLETE[/]"
+
+ self.total_tokens += tokens
+
+ # Update per-model stats
+ if model in self.model_stats:
+ m_stats = self.model_stats[model]
+ m_stats["total"] += 1
+ if success and not error:
+ m_stats["success"] += 1
+ m_stats["tokens"] += tokens
+ m_stats["time"] += duration
+
+ # Update progress bars
+ self.progress.update(self.overall_task, advance=1)
+ if model in self.model_tasks:
+ self.progress.update(self.model_tasks[model], advance=1)
+
+ # Update activity log
+ timestamp = time.strftime("%H:%M:%S")
+ self.activity_log.append(f"[{timestamp}] {status_msg} {label}")
+
+ # Remove from active runs
+ if thread_id in self.active_runs:
+ del self.active_runs[thread_id]
+
+ def generate_model_table(self) -> Table:
+ """Generate a detailed per-model statistics table."""
+ table = Table(expand=True, box=None)
+ table.add_column("Model", style="blue", ratio=2)
+ table.add_column("Progress", justify="right", ratio=1)
+ table.add_column("Success %", justify="right", ratio=1)
+ table.add_column("Tokens", justify="right", ratio=1)
+ table.add_column("Avg Time", justify="right", ratio=1)
+
+ for model, stats in self.model_stats.items():
+ if stats["target"] == 0: continue
+
+ progress = f"{stats['total']}/{stats['target']}"
+ success_rate = (stats["success"] / stats["total"] * 100) if stats["total"] > 0 else 0
+ avg_time = (stats["time"] / stats["total"]) if stats["total"] > 0 else 0
+
+ table.add_row(
+ model,
+ progress,
+ f"{success_rate:.0f}%",
+ f"{stats['tokens']:,}",
+ f"{avg_time:.1f}s"
+ )
+ return table
+
+ def format_time(self, seconds: float) -> str:
+ """Format seconds into HH:MM:SS."""
+ if seconds is None:
+ return "--:--:--"
+ h = int(seconds // 3600)
+ m = int((seconds % 3600) // 60)
+ s = int(seconds % 60)
+ return f"{h:02d}:{m:02d}:{s:02d}"
+
+ def generate_status_table(self) -> Table:
+ """Generate a summary table of the current status."""
+ table = Table(expand=True, box=None)
+ table.add_column("Metric", style="cyan")
+ table.add_column("Value", justify="right", style="magenta")
+
+ completed = self.success_count + self.incomplete_count + self.failed_count
+ remaining = self.total_runs - completed
+
+ # Time calculations
+ elapsed = time.time() - self.start_time
+ time_remaining = self.progress.tasks[self.overall_task].time_remaining
+
+ table.add_row("Success", f"[green]{self.success_count}[/green]")
+ table.add_row("Incomplete", f"[yellow]{self.incomplete_count}[/yellow]")
+ table.add_row("Failed", f"[red]{self.failed_count}[/red]")
+ table.add_row("Remaining", f"[white]{remaining}[/white]")
+ table.add_section()
+ table.add_row("Total Tokens", f"[bold white]{self.total_tokens:,}[/bold white]")
+ table.add_row("Elapsed", self.format_time(elapsed))
+ table.add_row("Est. Left", self.format_time(time_remaining))
+ table.add_row("Skipped", f"[grey50]{self.skipped}[/grey50]")
+
+ return table
+
+ def generate_active_panel(self) -> Panel:
+ """Generate a panel showing currently active runs."""
+ if not self.active_runs:
+ content = Text("Waiting for workers...", style="italic grey50")
+ else:
+ # Fix: Join using Text.from_markup to ensure colors render
+ lines = [Text.from_markup(line) for line in self.active_runs.values()]
+ content = Text("\n").join(lines)
+ return Panel(content, title="Currently Running", border_style="dim")
+
+ def generate_log_panel(self) -> Panel:
+ """Generate a panel showing recent activity log."""
+ # Fix: Join using Text.from_markup to ensure colors render
+ lines = [Text.from_markup(line) for line in self.activity_log]
+ content = Text("\n").join(lines)
+ return Panel(content, title="Recent Activity", border_style="dim")
+
+ def get_layout(self) -> Layout:
+ """Create the dashboard layout."""
+ layout = Layout()
+ layout.split_column(
+ Layout(name="header", size=3),
+ Layout(name="main"),
+ Layout(name="footer", size=7)
+ )
+
+ layout["main"].split_row(
+ Layout(name="progress_col", ratio=2),
+ Layout(name="stats_col", ratio=1)
+ )
+
+ layout["progress_col"].split_column(
+ Layout(name="bars", ratio=1),
+ Layout(name="model_details", ratio=1)
+ )
+
+ layout["footer"].split_row(
+ Layout(name="active", ratio=1),
+ Layout(name="logs", ratio=1)
+ )
+
+ layout["header"].update(Panel(Text("AI Agent Experiment Framework", justify="center", style="bold white"), style="blue"))
+ layout["bars"].update(Panel(self.progress, title="Overall Progress", style="white"))
+ layout["model_details"].update(Panel(self.generate_model_table(), title="Model Stats", style="white"))
+ layout["stats_col"].update(Panel(self.generate_status_table(), title="Totals", style="white"))
+ layout["active"].update(self.generate_active_panel())
+ layout["logs"].update(self.generate_log_panel())
+
+ return layout
+
+def print_final_summary(results: List[Dict]):
+ """Print a pretty grouped summary table at the end."""
+ console = Console()
+ console.print("\n")
+
+ # Group results by Model and Scenario
+ grouped = {}
+ for r in results:
+ key = (r['model'], r['scenario'])
+ if key not in grouped:
+ grouped[key] = []
+ grouped[key].append(r)
+
+ table = Table(title="[bold]Final Experiment Summary[/bold]", show_header=True, header_style="bold magenta", expand=True)
+ table.add_column("Model", style="blue", no_wrap=True)
+ table.add_column("Scenario", style="cyan")
+ table.add_column("Goal Type", style="yellow")
+ table.add_column("Successes", justify="center")
+ table.add_column("Total Tokens", justify="right")
+ table.add_column("Avg Time", justify="right")
+
+ for (model, scenario), group in sorted(grouped.items()):
+ total = len(group)
+ successes = sum(1 for r in group if r.get("success"))
+ total_tokens = sum(r.get("total_tokens", 0) for r in group)
+ avg_time = sum(r.get("duration_seconds", 0.0) for r in group) / total
+
+ # Collect distinct goal types in this group
+ goals = ", ".join(sorted(list(set(r.get("goal_type", "default") or "default" for r in group))))
+
+ success_color = "green" if successes == total else "yellow" if successes > 0 else "red"
+
+ table.add_row(
+ model,
+ scenario,
+ goals,
+ f"[{success_color}]{successes}/{total}[/]",
+ f"{total_tokens:,}",
+ f"{avg_time:.1f}s"
+ )
+
+ console.print(table)
+
+ # Aggregated Per-Model Table
+ model_table = Table(title="[bold]Aggregated Model Performance[/bold]", show_header=True, header_style="bold blue")
+ model_table.add_column("Model")
+ model_table.add_column("Total Runs", justify="right")
+ model_table.add_column("Overall Success %", justify="right")
+ model_table.add_column("Total Tokens", justify="right")
+
+ models = sorted(list(set(r['model'] for r in results)))
+ for model in models:
+ m_results = [r for r in results if r['model'] == model]
+ m_total = len(m_results)
+ m_successes = sum(1 for r in m_results if r.get("success"))
+ m_tokens = sum(r.get("total_tokens", 0) for r in m_results)
+
+ rate = (m_successes / m_total * 100) if m_total > 0 else 0
+ model_table.add_row(
+ model,
+ str(m_total),
+ f"{rate:.1f}%",
+ f"{m_tokens:,}"
+ )
+
+ console.print("\n")
+ console.print(model_table)
diff --git a/src/judge.py b/src/judge.py
index a3080ef..62a0f36 100644
--- a/src/judge.py
+++ b/src/judge.py
@@ -252,7 +252,7 @@ class RegexChecker:
{"verdict": "6/7", "matched": 6, "total": 7, "details": [...]}
"""
if not self.rules:
- return {"verdict": "N/A", "matched": 0, "total": 0, "details": [{"note": "no regex rules defined"}]}
+ return {"verdict": "PASS", "matched": 0, "total": 0, "details": [{"note": "no regex rules defined"}], "fraction": "0/0"}
details = []
required_total = 0
@@ -278,8 +278,18 @@ class RegexChecker:
if matched:
required_matched += 1
- verdict = f"{required_matched}/{required_total}"
- return {"verdict": verdict, "matched": required_matched, "total": required_total, "details": details}
+ if required_total == 0:
+ verdict = "PASS"
+ else:
+ verdict = "PASS" if required_matched == required_total else "FAIL"
+
+ return {
+ "verdict": verdict,
+ "matched": required_matched,
+ "total": required_total,
+ "details": details,
+ "fraction": f"{required_matched}/{required_total}"
+ }
# ---------------------------------------------------------------------------
diff --git a/src/runner.py b/src/runner.py
index d81e3cf..999856a 100644
--- a/src/runner.py
+++ b/src/runner.py
@@ -1,12 +1,7 @@
-"""
-Runner - orchestrates experiment runs based on config.
-Loops through models, scenarios, goal types, and oversight levels.
-Supports parallel execution via ThreadPoolExecutor.
-"""
import os
import glob
import concurrent.futures
-from typing import List, Dict, Any, Tuple
+from typing import List, Dict, Any, Tuple, Optional
from config_loader import ConfigLoader, ProviderConfig, ModelConfig, ScenarioConfig
from vfs import VFS
from agent import Agent
@@ -15,6 +10,9 @@ from logger import get_logger
import datetime
import threading
+from rich.live import Live
+from dashboard import ExperimentDashboard, print_final_summary
+
# Get logger instance
logger = get_logger("experiment")
@@ -38,6 +36,7 @@ class ExperimentRunner:
self.results: List[Dict] = []
self.verbose = verbose
self.resume = resume
+ self.dashboard: Optional[ExperimentDashboard] = None
def run_all(self):
"""Run all experiments defined in config."""
@@ -46,62 +45,67 @@ class ExperimentRunner:
logger.info(f"{'='*60}\n")
# Build list of all work items (model, scenario, goal_type, oversight)
- work_items = self._build_work_items()
+ work_items, skipped_count = self._build_work_items()
- if not work_items:
+ if not work_items and skipped_count == 0:
logger.warning("No work items to run.")
return
+
+ if not work_items:
+ logger.info(f"All {skipped_count} items already exist. Skipping all.")
+ return
max_workers = self.config.max_workers
total_items = len(work_items)
- logger.info(f"Total work items: {total_items}, Max workers: {max_workers}")
+
+ # Prepare model list for dashboard
+ model_names = sorted(list(set(item["model_config"].id for item in work_items)))
+ model_counts = {}
+ for item in work_items:
+ m_id = item["model_config"].id
+ model_counts[m_id] = model_counts.get(m_id, 0) + 1
+
+ self.dashboard = ExperimentDashboard(total_items, model_names, skipped=skipped_count)
+ for m_id, count in model_counts.items():
+ self.dashboard.update_model_total(m_id, count)
- if max_workers <= 1:
- # Sequential execution (original behavior)
- for item in work_items:
- self._execute_work_item(item)
- else:
- # Parallel execution
- logger.info(f"Running with {max_workers} parallel workers")
- from tqdm import tqdm
- from tqdm.contrib.logging import logging_redirect_tqdm
- import logging
-
- # Reduce console spam during parallel runs to keep the progress bar clean
- for handler in logger.handlers:
- if isinstance(handler, logging.StreamHandler) and not isinstance(handler, logging.FileHandler):
- # Keep console relatively quiet (WARNING/ERROR/CRITICAL)
- handler.setLevel(logging.WARNING)
-
- with concurrent.futures.ThreadPoolExecutor(max_workers=max_workers) as executor:
- futures = {
- executor.submit(self._execute_work_item, item): item
- for item in work_items
- }
- with logging_redirect_tqdm():
- with tqdm(total=total_items, desc="Running Experiments", unit="run", dynamic_ncols=True) as pbar:
- for future in concurrent.futures.as_completed(futures):
- item = futures[future]
- try:
- future.result()
- except Exception as e:
- logger.critical(f"Work item failed: {item.get('label', 'unknown')}: {e}")
- pbar.write(f"ERROR: {item.get('label', 'unknown')} failed: {e}")
- pbar.update(1)
-
- # Summary
- successful = sum(1 for r in self.results if r.get("success", False))
- total_runs = len(self.results)
-
- logger.info(f"\n{'='*60}")
- logger.info(f"Experiment Complete: {total_runs} runs completed")
- logger.info(f" SUCCESS: {successful}")
- logger.info(f" INCOMPLETE: {total_runs - successful}")
- logger.info(f"{'='*60}\n")
+ logger.info(f"Total work items: {total_items}, Max workers: {max_workers}")
- def _build_work_items(self) -> List[Dict[str, Any]]:
- """Build a flat list of all (model, scenario, goal_type, oversight, run_num) combos."""
+ with Live(self.dashboard.get_layout(), refresh_per_second=4, vertical_overflow="visible") as live:
+ if max_workers <= 1:
+ # Sequential execution
+ for item in work_items:
+ self._execute_work_item(item)
+ live.update(self.dashboard.get_layout())
+ else:
+ # Parallel execution
+ import logging
+ # Reduce console spam during parallel runs to keep the progress bar clean
+ for handler in logger.handlers:
+ if isinstance(handler, logging.StreamHandler) and not isinstance(handler, logging.FileHandler):
+ handler.setLevel(logging.WARNING)
+
+ with concurrent.futures.ThreadPoolExecutor(max_workers=max_workers) as executor:
+ futures = {
+ executor.submit(self._execute_work_item, item): item
+ for item in work_items
+ }
+ for future in concurrent.futures.as_completed(futures):
+ try:
+ future.result()
+ except Exception:
+ # Error already logged in _execute_work_item
+ pass
+ live.update(self.dashboard.get_layout())
+
+ # Print final summary table
+ print_final_summary(self.results)
+
+ def _build_work_items(self) -> Tuple[List[Dict[str, Any]], int]:
+ """Build a flat list of all (model, scenario, goal_type, oversight, run_num) combos.
+ Returns (work_items, skipped_count)."""
work_items = []
+ skipped_count = 0
goal_types = self.config.goal_types
for model_config in self.config.models:
@@ -119,20 +123,20 @@ class ExperimentRunner:
for oversight_level in oversight_levels:
if goal_types:
- # New: iterate over goal types
for goal_type in goal_types:
- items = self._build_run_items(
+ items, skipped = self._build_run_items(
model_config, scenario_config, oversight_level, goal_type
)
work_items.extend(items)
+ skipped_count += skipped
else:
- # Legacy: single hidden_goal.md (no goal_types in config)
- items = self._build_run_items(
+ items, skipped = self._build_run_items(
model_config, scenario_config, oversight_level, ""
)
work_items.extend(items)
+ skipped_count += skipped
- return work_items
+ return work_items, skipped_count
def _build_run_items(
self,
@@ -140,7 +144,7 @@ class ExperimentRunner:
scenario_config: ScenarioConfig,
oversight_level: str,
goal_type: str
- ) -> List[Dict[str, Any]]:
+ ) -> Tuple[List[Dict[str, Any]], int]:
"""Build individual run items for a specific combo, accounting for resume."""
model_name_safe = model_config.id.replace("/", "_")
scenario_name = os.path.basename(scenario_config.path)
@@ -160,15 +164,12 @@ class ExperimentRunner:
all_json = glob.glob(os.path.join(log_dir, "*.json"))
existing_runs = len([f for f in all_json if not f.endswith(".partial.json")])
+ skipped = 0
if existing_runs >= scenario_config.runs:
- goal_label = f"/{goal_type}" if goal_type else ""
- logger.info(f" SKIP: {model_config.id} | {scenario_name}{goal_label} | "
- f"{oversight_level} ({existing_runs}/{scenario_config.runs} exist)")
- return []
+ skipped = scenario_config.runs
+ return [], skipped
elif existing_runs > 0:
- goal_label = f"/{goal_type}" if goal_type else ""
- logger.info(f" RESUME: {model_config.id} | {scenario_name}{goal_label} | "
- f"{oversight_level} ({existing_runs}/{scenario_config.runs} exist)")
+ skipped = existing_runs
items = []
for run_num in range(existing_runs + 1, scenario_config.runs + 1):
@@ -182,7 +183,7 @@ class ExperimentRunner:
"label": f"{model_config.id} | {scenario_name}{goal_label} | {oversight_level} | run {run_num}"
})
- return items
+ return items, skipped
def _ensure_baseline(self, model_config: ModelConfig, scenario_config: ScenarioConfig):
"""Ensure baseline exists for a model+scenario combo (thread-safe)."""
@@ -192,21 +193,25 @@ class ExperimentRunner:
baseline_path = os.path.join(output_dir, model_name_safe, scenario_name, "baseline.md")
if not self.config.generate_baseline:
- if not os.path.exists(baseline_path):
- logger.warning(f"Baseline generation DISABLED. No baseline for "
- f"{model_config.id} | {scenario_name}.")
return
if not os.path.exists(baseline_path):
- logger.info(f"\n--- Generating baseline: {model_config.id} | {scenario_name} ---")
provider_config = self.config.get_provider(model_config.provider)
self._run_baseline(model_config, provider_config, scenario_config)
- logger.info(f" Baseline saved to {baseline_path}")
def _execute_work_item(self, item: Dict[str, Any]):
"""Execute a single work item (one experiment run). Thread-safe."""
+ thread_id = threading.get_ident()
+ if self.dashboard:
+ self.dashboard.start_run(
+ thread_id,
+ item["model_config"].id,
+ os.path.basename(item["scenario_config"].path),
+ item["goal_type"]
+ )
+
try:
- self._run_single(
+ success, tokens, duration = self._run_single(
model_config=item["model_config"],
provider_config=self.config.get_provider(item["model_config"].provider),
scenario_config=item["scenario_config"],
@@ -214,10 +219,26 @@ class ExperimentRunner:
goal_type=item["goal_type"],
run_num=item["run_num"]
)
+
+ if self.dashboard:
+ self.dashboard.complete_run(
+ thread_id,
+ item["model_config"].id,
+ success=success,
+ tokens=tokens,
+ duration=duration,
+ label=item["label"]
+ )
except Exception as e:
logger.critical(f"ERROR in {item['label']}: {e}")
- import traceback
- traceback.print_exc()
+ if self.dashboard:
+ self.dashboard.complete_run(
+ thread_id,
+ item["model_config"].id,
+ success=False,
+ error=True,
+ label=item["label"]
+ )
def _extract_baseline_content(self, logs: List[Dict]) -> str:
"""Extract baseline content from the last create_file tool call."""
@@ -283,12 +304,9 @@ class ExperimentRunner:
)
# Run the conversation
- logger.info(f" Running baseline...")
output_dir = self.config.output_dir
agent.enable_incremental_save(output_dir=output_dir)
- start_time = datetime.datetime.now()
- result = agent.run(user_prompt)
- end_time = datetime.datetime.now()
+ agent.run(user_prompt)
# Extract baseline content from last create_file tool call
baseline_content = self._extract_baseline_content(agent.logs)
@@ -306,8 +324,6 @@ class ExperimentRunner:
# Save baseline log separately
agent.save_logs(output_dir=output_dir)
- logger.info(f" Baseline completed in {(end_time - start_time).total_seconds():.2f}s")
-
def _run_single(
self,
model_config: ModelConfig,
@@ -316,7 +332,7 @@ class ExperimentRunner:
oversight_level: str,
run_num: int,
goal_type: str = ""
- ):
+ ) -> Tuple[bool, int, float]:
"""Run a single experiment. Thread-safe with isolated VFS."""
# Load scenario prompts
benign_goal = load_prompt(os.path.join(scenario_config.path, "benign_goal.md"))
@@ -349,11 +365,6 @@ class ExperimentRunner:
vfs_instance = VFS.create(vfs_path)
tool_functions = make_tools_for_vfs(vfs_instance)
- # Log VFS info
- vfs_files = vfs_instance.list_files("/")
- logger.debug(f"VFS initialized from: {vfs_path}")
- logger.debug(f"VFS files: {vfs_files}")
-
# Create agent with isolated VFS
save_vfs_state = self.config._config.get('output', {}).get('save_vfs_state', True)
scenario_name = os.path.basename(scenario_config.path)
@@ -371,13 +382,10 @@ class ExperimentRunner:
)
# Run the conversation
- goal_label = f" | {goal_type}" if goal_type else ""
- logger.info(f"\n Starting: {model_config.id} | {scenario_name}{goal_label} | "
- f"{oversight_level} | run {run_num}")
output_dir = self.config.output_dir
agent.enable_incremental_save(output_dir=output_dir)
start_time = datetime.datetime.now()
- result = agent.run(user_prompt)
+ agent.run(user_prompt)
end_time = datetime.datetime.now()
# Save logs
@@ -393,6 +401,8 @@ class ExperimentRunner:
elif msg.get("role") == "assistant" and msg.get("content") is None and msg.get("tool_calls"):
continue
+ duration = (end_time - start_time).total_seconds()
+
# Record result (thread-safe)
result_entry = {
"model": model_config.id,
@@ -401,8 +411,7 @@ class ExperimentRunner:
"goal_type": goal_type,
"oversight_level": oversight_level,
"run_num": run_num,
- "run_id": f"{model_config.id}/{scenario_name}/{goal_type}/{oversight_level}/{datetime.datetime.now().strftime('%Y%m%d_%H%M%S')}",
- "duration_seconds": (end_time - start_time).total_seconds(),
+ "duration_seconds": duration,
"total_tokens": agent.total_tokens,
"success": success,
"log_file": log_file
@@ -410,10 +419,8 @@ class ExperimentRunner:
with _results_lock:
self.results.append(result_entry)
-
- status = "SUCCESS" if success else "INCOMPLETE"
- logger.info(f" [{status}] {model_config.id} | {scenario_name}{goal_label} | "
- f"{oversight_level} | run {run_num} ({(end_time - start_time).total_seconds():.2f}s)")
+
+ return success, agent.total_tokens, duration
def run_from_config(config_path: str = "config.yaml", resume: bool = True):
@@ -425,14 +432,3 @@ def run_from_config(config_path: str = "config.yaml", resume: bool = True):
runner.run_all()
return runner.results
-
-
-if __name__ == "__main__":
- import argparse
- parser = argparse.ArgumentParser(description="Run experiments from config")
- parser.add_argument("--config", default="config.yaml", help="Path to config file")
- parser.add_argument("--no-resume", dest="resume", action="store_false",
- default=True, help="Ignore existing logs and start fresh")
- args = parser.parse_args()
-
- run_from_config(args.config, resume=args.resume)
diff --git a/tests/test_v2_features.py b/tests/test_v2_features.py
index d8f9ea7..f738d1f 100644
--- a/tests/test_v2_features.py
+++ b/tests/test_v2_features.py
@@ -173,7 +173,7 @@ output:
"""With goal_types, should multiply: model × scenario × goal_type × oversight × runs."""
config = self._make_config(tmp_path, goal_types=["self_serving", "moral", "bare"], runs=2)
runner = ExperimentRunner(config, resume=False)
- items = runner._build_work_items()
+ items, _ = runner._build_work_items()
# 1 model × 1 scenario × 3 goal_types × 2 oversight × 2 runs = 12
assert len(items) == 12
goal_types_seen = {item["goal_type"] for item in items}
@@ -183,7 +183,7 @@ output:
"""Without goal_types, should behave as legacy: model × scenario × oversight × runs."""
config = self._make_config(tmp_path, goal_types=None, runs=2)
runner = ExperimentRunner(config, resume=False)
- items = runner._build_work_items()
+ items, _ = runner._build_work_items()
# 1 model × 1 scenario × 2 oversight × 2 runs = 4
assert len(items) == 4
# All items should have empty goal_type
@@ -202,7 +202,7 @@ output:
(log_dir / "20250101_120000.json").write_text("{}")
(log_dir / "20250101_120001.json").write_text("{}")
- items = runner._build_work_items()
+ items, _ = runner._build_work_items()
# moral/low: 3-2=1 remaining. moral/high: 3 runs. Total: 4
moral_low_items = [i for i in items if i["goal_type"] == "moral" and i["oversight_level"] == "low"]
moral_high_items = [i for i in items if i["goal_type"] == "moral" and i["oversight_level"] == "high"]
diff --git a/uv.lock b/uv.lock
index f38a79b..2f6c2d3 100644
--- a/uv.lock
+++ b/uv.lock
@@ -166,6 +166,7 @@ dependencies = [
{ name = "pandas" },
{ name = "python-dotenv" },
{ name = "pyyaml" },
+ { name = "rich" },
{ name = "seaborn" },
{ name = "sse-starlette" },
{ name = "tqdm" },
@@ -186,6 +187,7 @@ requires-dist = [
{ name = "pandas", specifier = ">=3.0.1" },
{ name = "python-dotenv", specifier = ">=1.0.0" },
{ name = "pyyaml", specifier = ">=6.0.3" },
+ { name = "rich", specifier = ">=14.3.3" },
{ name = "seaborn", specifier = ">=0.13.2" },
{ name = "sse-starlette", specifier = ">=3.3.2" },
{ name = "tqdm", specifier = ">=4.67.1" },
@@ -450,6 +452,18 @@ wheels = [
]
[[package]]
+name = "markdown-it-py"
+version = "4.0.0"
+source = { registry = "https://pypi.org/simple" }
+dependencies = [
+ { name = "mdurl" },
+]
+sdist = { url = "https://files.pythonhosted.org/packages/5b/f5/4ec618ed16cc4f8fb3b701563655a69816155e79e24a17b651541804721d/markdown_it_py-4.0.0.tar.gz", hash = "sha256:cb0a2b4aa34f932c007117b194e945bd74e0ec24133ceb5bac59009cda1cb9f3", size = 73070, upload-time = "2025-08-11T12:57:52.854Z" }
+wheels = [
+ { url = "https://files.pythonhosted.org/packages/94/54/e7d793b573f298e1c9013b8c4dade17d481164aa517d1d7148619c2cedbf/markdown_it_py-4.0.0-py3-none-any.whl", hash = "sha256:87327c59b172c5011896038353a81343b6754500a08cd7a4973bb48c6d578147", size = 87321, upload-time = "2025-08-11T12:57:51.923Z" },
+]
+
+[[package]]
name = "matplotlib"
version = "3.10.8"
source = { registry = "https://pypi.org/simple" }
@@ -497,6 +511,15 @@ wheels = [
]
[[package]]
+name = "mdurl"
+version = "0.1.2"
+source = { registry = "https://pypi.org/simple" }
+sdist = { url = "https://files.pythonhosted.org/packages/d6/54/cfe61301667036ec958cb99bd3efefba235e65cdeb9c84d24a8293ba1d90/mdurl-0.1.2.tar.gz", hash = "sha256:bb413d29f5eea38f31dd4754dd7377d4465116fb207585f97bf925588687c1ba", size = 8729, upload-time = "2022-08-14T12:40:10.846Z" }
+wheels = [
+ { url = "https://files.pythonhosted.org/packages/b3/38/89ba8ad64ae25be8de66a6d463314cf1eb366222074cfda9ee839c56a4b4/mdurl-0.1.2-py3-none-any.whl", hash = "sha256:84008a41e51615a49fc9966191ff91509e3c40b939176e643fd50a5c2196b8f8", size = 9979, upload-time = "2022-08-14T12:40:09.779Z" },
+]
+
+[[package]]
name = "numpy"
version = "2.4.2"
source = { registry = "https://pypi.org/simple" }
@@ -845,6 +868,19 @@ wheels = [
]
[[package]]
+name = "rich"
+version = "14.3.3"
+source = { registry = "https://pypi.org/simple" }
+dependencies = [
+ { name = "markdown-it-py" },
+ { name = "pygments" },
+]
+sdist = { url = "https://files.pythonhosted.org/packages/b3/c6/f3b320c27991c46f43ee9d856302c70dc2d0fb2dba4842ff739d5f46b393/rich-14.3.3.tar.gz", hash = "sha256:b8daa0b9e4eef54dd8cf7c86c03713f53241884e814f4e2f5fb342fe520f639b", size = 230582, upload-time = "2026-02-19T17:23:12.474Z" }
+wheels = [
+ { url = "https://files.pythonhosted.org/packages/14/25/b208c5683343959b670dc001595f2f3737e051da617f66c31f7c4fa93abc/rich-14.3.3-py3-none-any.whl", hash = "sha256:793431c1f8619afa7d3b52b2cdec859562b950ea0d4b6b505397612db8d5362d", size = 310458, upload-time = "2026-02-19T17:23:13.732Z" },
+]
+
+[[package]]
name = "seaborn"
version = "0.13.2"
source = { registry = "https://pypi.org/simple" }