Spaces:
Sleeping
Sleeping
| #!/usr/bin/env python3 | |
| """ | |
| Random Agent — sanity test for the WorkSim environment loop. | |
| Calls reset(), then randomly selects tools and calls step() until | |
| the episode ends. Used to verify: | |
| - No crashes on random inputs | |
| - World state evolves correctly | |
| - Termination conditions trigger properly | |
| - Audit logging works | |
| Usage: | |
| python -m scripts.run_random_agent | |
| python -m scripts.run_random_agent --episodes 5 --max-steps 30 --seed 42 | |
| """ | |
| import argparse | |
| import json | |
| import random | |
| import sys | |
| import os | |
| sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..")) | |
| from env.gymnasium_env import WorkSimEnv | |
| from env.tools import get_available_tools | |
| # Random argument generators per tool | |
| def _random_query(): | |
| queries = [ | |
| "budget", "deadline", "vendor", "pricing", "schedule", "meeting", | |
| "report", "Acme", "design", "review", "approved", "urgent", | |
| "quarterly", "contract", "milestone", "requirements", | |
| ] | |
| return random.choice(queries) | |
| def _random_args_for_tool(tool_name: str, obs: dict) -> dict: | |
| """Generate plausible random arguments for a tool.""" | |
| # Get asset IDs from observation | |
| asset_ids = list(obs.get("assets", {}).keys()) | |
| email_ids = [aid for aid, a in obs.get("assets", {}).items() if a.get("asset_type") == "email_thread"] | |
| chat_ids = [aid for aid, a in obs.get("assets", {}).items() if a.get("asset_type") == "chat_thread"] | |
| doc_ids = [aid for aid, a in obs.get("assets", {}).items() if a.get("asset_type") in ("document_asset", "meeting_artifact")] | |
| sheet_ids = [aid for aid, a in obs.get("assets", {}).items() if a.get("asset_type") == "spreadsheet_asset"] | |
| people = [p.get("name", "Unknown") for p in obs.get("org_directory", {}).values()] | |
| if tool_name == "mail.search": | |
| return {"query": _random_query()} | |
| elif tool_name == "mail.open_thread": | |
| return {"thread_id": random.choice(email_ids) if email_ids else "nonexistent"} | |
| elif tool_name == "mail.open_message": | |
| return {"message_id": f"msg_{random.randint(1, 100)}"} | |
| elif tool_name == "mail.list_inbox": | |
| return {"page": 1} | |
| elif tool_name == "mail.draft_reply": | |
| return {"thread_id": random.choice(email_ids) if email_ids else "x", "content": "Thank you for the update."} | |
| elif tool_name == "mail.send_reply": | |
| drafts = obs.get("drafts", {}) | |
| return {"draft_id": random.choice(list(drafts.keys())) if drafts else "draft_none"} | |
| elif tool_name == "chat.list_channels": | |
| return {} | |
| elif tool_name == "chat.search": | |
| return {"query": _random_query()} | |
| elif tool_name == "chat.open_thread": | |
| return {"thread_id": random.choice(chat_ids) if chat_ids else "nonexistent"} | |
| elif tool_name == "chat.open_channel": | |
| return {"channel_id": random.choice(chat_ids) if chat_ids else "x", "window": 10} | |
| elif tool_name == "chat.post_message": | |
| return {"channel_id": random.choice(chat_ids) if chat_ids else "x", "content": "Checking in on this."} | |
| elif tool_name == "drive.list_files": | |
| return {} | |
| elif tool_name == "drive.search": | |
| return {"query": _random_query()} | |
| elif tool_name == "drive.open_file": | |
| return {"file_id": random.choice(doc_ids) if doc_ids else "nonexistent"} | |
| elif tool_name == "drive.compare_versions": | |
| return {"file_id": random.choice(doc_ids) if doc_ids else "x"} | |
| elif tool_name == "sheet.open": | |
| return {"file_id": random.choice(sheet_ids) if sheet_ids else "x"} | |
| elif tool_name == "sheet.read_range": | |
| return {"file_id": random.choice(sheet_ids) if sheet_ids else "x", "tab_name": "Budget", "range_ref": "A1:E5"} | |
| elif tool_name == "sheet.write_cell": | |
| return {"file_id": random.choice(sheet_ids) if sheet_ids else "x", "tab_name": "Budget", "cell_ref": "C7", "value": random.randint(100, 9999)} | |
| elif tool_name == "sheet.write_range": | |
| return {"file_id": random.choice(sheet_ids) if sheet_ids else "x", "tab_name": "Budget", "values": {"C7": 1000, "C8": 2000}} | |
| elif tool_name == "sheet.get_formula": | |
| return {"file_id": random.choice(sheet_ids) if sheet_ids else "x", "tab_name": "Budget", "cell_ref": "D10"} | |
| elif tool_name == "calendar.view": | |
| return {"participants": random.sample(people, min(2, len(people))) if people else []} | |
| elif tool_name == "calendar.check_conflicts": | |
| return {"participants": random.sample(people, min(2, len(people))) if people else [], "proposed_time": "2024-03-12T10:00:00"} | |
| elif tool_name == "calendar.propose_time": | |
| return {"participants": random.sample(people, min(2, len(people))) if people else [], "constraints": {"duration_minutes": 60}} | |
| elif tool_name == "calendar.create_hold": | |
| return {"title": "Random Meeting", "participants": people[:2] if people else ["A"], "time_slot": {"start": "2024-03-12T14:00:00", "end": "2024-03-12T15:00:00"}} | |
| elif tool_name == "notes.write": | |
| return {"content": f"Step {random.randint(1,50)} notes: investigating {_random_query()}"} | |
| elif tool_name == "memo.create": | |
| return {"title": "Draft Memo", "content": "This is a draft analysis based on findings so far."} | |
| elif tool_name == "memo.submit": | |
| # Only submit if a memo was created | |
| deliverables = obs.get("drafts", {}) | |
| return {"memo_id": "memo_none"} # Will usually fail gracefully | |
| elif tool_name == "task.mark_done": | |
| objectives = obs.get("objectives", []) | |
| return {"task_id": objectives[0]["id"] if objectives else "unknown"} | |
| elif tool_name == "search.global": | |
| return {"query": _random_query()} | |
| elif tool_name == "entity.resolve": | |
| return {"name_or_alias": random.choice(people) if people else "unknown"} | |
| elif tool_name == "workspace.status": | |
| return {} | |
| return {} | |
| def run_random_episode( | |
| project_type: str, difficulty: int, seed: int, max_steps: int, verbose: bool = True | |
| ) -> dict: | |
| """Run one random episode and return statistics.""" | |
| env = WorkSimEnv(project_type=project_type, difficulty_level=difficulty, max_steps=max_steps) | |
| obs, info = env.reset(seed=seed) | |
| rng = random.Random(seed) | |
| tools = get_available_tools() | |
| # Weight exploration tools higher than submission | |
| safe_tools = [t for t in tools if t not in ("memo.submit",)] | |
| stats = { | |
| "world_id": info.get("world_id", ""), | |
| "project_type": project_type, | |
| "difficulty": difficulty, | |
| "seed": seed, | |
| "steps": 0, | |
| "total_reward": 0.0, | |
| "tool_calls": {}, | |
| "errors": 0, | |
| "terminated": False, | |
| "truncated": False, | |
| "termination_reason": "", | |
| } | |
| if verbose: | |
| print(f"\n{'='*60}") | |
| print(f"Random Agent — {info.get('world_id', 'N/A')}") | |
| print(f" Project: {project_type} | Difficulty: {difficulty} | Seed: {seed}") | |
| print(f" Task: {obs.get('task_goal', '')[:70]}...") | |
| print(f"{'='*60}") | |
| while not env.done: | |
| # Pick a random tool | |
| tool = rng.choice(safe_tools) | |
| # After step 15, occasionally try memo.create + submit to test termination | |
| if env.step_count > max_steps * 0.7 and rng.random() < 0.1: | |
| tool = "memo.create" | |
| args = _random_args_for_tool(tool, obs) | |
| action = {"tool_name": tool, "arguments": args} | |
| obs, reward, terminated, truncated, step_info = env.step(action) | |
| stats["steps"] += 1 | |
| stats["total_reward"] += reward | |
| stats["tool_calls"][tool] = stats["tool_calls"].get(tool, 0) + 1 | |
| tool_result = step_info.get("tool_result", {}) | |
| if tool_result.get("status") == "error": | |
| stats["errors"] += 1 | |
| if verbose and stats["steps"] % 5 == 0: | |
| print(f" Step {stats['steps']}: {tool} → {tool_result.get('status', '?')} (reward: {reward:+.3f}, total: {env.episode_reward:.3f})") | |
| # If we just created a memo, submit it on the next round | |
| if tool == "memo.create" and tool_result.get("status") == "success": | |
| memo_id = tool_result.get("result", {}).get("memo_id", "") | |
| if memo_id: | |
| submit_action = {"tool_name": "memo.submit", "arguments": {"memo_id": memo_id}} | |
| obs, reward, terminated, truncated, step_info = env.step(submit_action) | |
| stats["steps"] += 1 | |
| stats["total_reward"] += reward | |
| stats["tool_calls"]["memo.submit"] = stats["tool_calls"].get("memo.submit", 0) + 1 | |
| stats["terminated"] = env._world.terminated | |
| stats["truncated"] = env._world.truncated | |
| stats["termination_reason"] = env.get_full_state().get("completion_state", {}).get("termination_reason", "unknown") | |
| if verbose: | |
| print(f"\n{'─'*60}") | |
| print(f"Episode ended: {'TERMINATED' if stats['terminated'] else 'TRUNCATED'}") | |
| print(f" Reason: {stats['termination_reason']}") | |
| print(f" Steps: {stats['steps']}") | |
| print(f" Total reward: {stats['total_reward']:.3f}") | |
| print(f" Errors: {stats['errors']}") | |
| print(f" Tool distribution: {json.dumps(stats['tool_calls'], indent=2)}") | |
| env.render() | |
| return stats | |
| def main(): | |
| parser = argparse.ArgumentParser(description="Random Agent Sanity Test") | |
| parser.add_argument("--episodes", type=int, default=3, help="Number of episodes") | |
| parser.add_argument("--max-steps", type=int, default=30, help="Max steps per episode") | |
| parser.add_argument("--seed", type=int, default=42, help="Base random seed") | |
| parser.add_argument("--project", default=None, help="Project type (default: all)") | |
| parser.add_argument("--difficulty", type=int, default=2, help="Difficulty level") | |
| parser.add_argument("--quiet", action="store_true", help="Less output") | |
| args = parser.parse_args() | |
| project_types = [args.project] if args.project else ["client_brief", "calendar_conflict"] | |
| all_stats = [] | |
| for i in range(args.episodes): | |
| pt = project_types[i % len(project_types)] | |
| seed = args.seed + i * 1000 | |
| stats = run_random_episode( | |
| project_type=pt, | |
| difficulty=args.difficulty, | |
| seed=seed, | |
| max_steps=args.max_steps, | |
| verbose=not args.quiet, | |
| ) | |
| all_stats.append(stats) | |
| # Summary | |
| print(f"\n{'='*60}") | |
| print(f"SUMMARY — {len(all_stats)} episode(s)") | |
| print(f"{'='*60}") | |
| total_steps = sum(s["steps"] for s in all_stats) | |
| total_errors = sum(s["errors"] for s in all_stats) | |
| total_reward = sum(s["total_reward"] for s in all_stats) | |
| terminated = sum(1 for s in all_stats if s["terminated"]) | |
| truncated = sum(1 for s in all_stats if s["truncated"]) | |
| print(f" Total steps: {total_steps}") | |
| print(f" Total errors: {total_errors} ({total_errors/max(total_steps,1)*100:.1f}%)") | |
| print(f" Avg reward: {total_reward/len(all_stats):.3f}") | |
| print(f" Terminated: {terminated} | Truncated: {truncated}") | |
| # Aggregate tool usage | |
| all_tools = {} | |
| for s in all_stats: | |
| for t, c in s["tool_calls"].items(): | |
| all_tools[t] = all_tools.get(t, 0) + c | |
| print(f"\n Tool usage (total):") | |
| for t in sorted(all_tools, key=all_tools.get, reverse=True): | |
| print(f" {t}: {all_tools[t]}") | |
| print(f"\n✅ Random agent sanity test complete.") | |
| if __name__ == "__main__": | |
| main() | |