From 366d7247da9b0e8825cb5661974602144349e6e2 Mon Sep 17 00:00:00 2001 From: Adam Moussa <166072409+amoussa1229@users.noreply.github.com> Date: Fri, 15 May 2026 11:20:28 -0400 Subject: [PATCH] Stabilize router and consolidate agent registry Phase 1 stabilization. Removes the four-copy prompt/agent-description drift surface and the silent router fallback. - models.py: hoist model IDs to module-level constants; add with_retries() helper (2 retries on Anthropic+OpenAI transient errors via with_retry). - agents.py: single AGENTS dict (model_fn, prompt, description) and a make_agent_node() factory that collapses six near-identical node functions. - graph.py: router prompt is generated from AGENTS; router_node uses with_structured_output(RouteDecision) and returns an explicit "unknown" route instead of the silent "researcher" fallback. New unknown_node wires to END. All LLM invocations go through with_retries. - state.py: add "unknown" to the route Literal. - run.py: --route-only now imports router_node from graph.py, killing the fourth prompt copy. - tests/: pytest golden-set (20 labelled tasks + size guard). Skips cleanly without ANTHROPIC_API_KEY or COMPOSIO_API_KEY. Validated 21/21 passing. --- agents.py | 101 ++++++++++---------- graph.py | 176 ++++++++++++++++++++--------------- models.py | 71 ++++++++++++-- run.py | 21 ++--- state.py | 1 + tests/__init__.py | 0 tests/conftest.py | 7 ++ tests/test_routing_golden.py | 115 +++++++++++++++++++++++ 8 files changed, 344 insertions(+), 148 deletions(-) create mode 100644 tests/__init__.py create mode 100644 tests/conftest.py create mode 100644 tests/test_routing_golden.py diff --git a/agents.py b/agents.py index 8b5b8d9..2756ede 100644 --- a/agents.py +++ b/agents.py @@ -7,6 +7,7 @@ from models import ( get_cross_reviewer, get_scanner, get_fast_coder, + with_retries, ) IMPLEMENTER_PROMPT = """You are an implementation agent. You write clean, production-ready code. @@ -42,55 +43,59 @@ Implement exactly what is asked. No extras, no refactoring beyond scope. Return complete, working code.""" -def implementer_node(state: OrchestratorState) -> dict: - llm = get_implementer() - response = llm.invoke([ - SystemMessage(content=IMPLEMENTER_PROMPT), - HumanMessage(content=state["task"]), - ]) - return {"result": response.content, "messages": [response]} +# Single source of truth for simple-pattern agents (system + human → result). +# Connector is handled separately in graph.py because it binds tools. +AGENTS = { + "implementer": { + "model_fn": get_implementer, + "prompt": IMPLEMENTER_PROMPT, + "description": "Write new code, add features, fix bugs. Use for any coding task with a clear spec.", + }, + "reviewer": { + "model_fn": get_reviewer, + "prompt": REVIEWER_PROMPT, + "description": "Review code changes (diffs, PRs) for correctness, security, maintainability. Uses Claude.", + }, + "researcher": { + "model_fn": get_researcher, + "prompt": RESEARCHER_PROMPT, + "description": "Look up documentation, API references, technical questions. Fast and cheap.", + }, + "cross_reviewer": { + "model_fn": get_cross_reviewer, + "prompt": CROSS_REVIEWER_PROMPT, + "description": ( + "Independent code review using a different AI model (GPT). Use when you want a " + "second opinion that catches different blind spots than Claude." + ), + }, + "scanner": { + "model_fn": get_scanner, + "prompt": SCANNER_PROMPT, + "description": ( + "Analyze large codebases for patterns, consistency, structural issues. " + "Uses Gemini's large context window." + ), + }, + "fast_coder": { + "model_fn": get_fast_coder, + "prompt": FAST_CODER_PROMPT, + "description": "Quick, bounded coding for crystal-clear specs. Uses DeepSeek. Best for small, well-defined tasks.", + }, +} -def reviewer_node(state: OrchestratorState) -> dict: - llm = get_reviewer() - response = llm.invoke([ - SystemMessage(content=REVIEWER_PROMPT), - HumanMessage(content=state["task"]), - ]) - return {"result": response.content, "messages": [response]} +def make_agent_node(label: str): + cfg = AGENTS[label] + def node(state: OrchestratorState) -> dict: + llm = with_retries(cfg["model_fn"]()) + response = llm.invoke( + [ + SystemMessage(content=cfg["prompt"]), + HumanMessage(content=state["task"]), + ] + ) + return {"result": response.content, "messages": [response]} -def researcher_node(state: OrchestratorState) -> dict: - llm = get_researcher() - response = llm.invoke([ - SystemMessage(content=RESEARCHER_PROMPT), - HumanMessage(content=state["task"]), - ]) - return {"result": response.content, "messages": [response]} - - -def cross_reviewer_node(state: OrchestratorState) -> dict: - llm = get_cross_reviewer() - response = llm.invoke([ - SystemMessage(content=CROSS_REVIEWER_PROMPT), - HumanMessage(content=state["task"]), - ]) - return {"result": response.content, "messages": [response]} - - -def scanner_node(state: OrchestratorState) -> dict: - llm = get_scanner() - response = llm.invoke([ - SystemMessage(content=SCANNER_PROMPT), - HumanMessage(content=state["task"]), - ]) - return {"result": response.content, "messages": [response]} - - -def fast_coder_node(state: OrchestratorState) -> dict: - llm = get_fast_coder() - response = llm.invoke([ - SystemMessage(content=FAST_CODER_PROMPT), - HumanMessage(content=state["task"]), - ]) - return {"result": response.content, "messages": [response]} + return node diff --git a/graph.py b/graph.py index 835338a..1cff246 100644 --- a/graph.py +++ b/graph.py @@ -1,49 +1,56 @@ +from typing import Literal + from dotenv import load_dotenv - -load_dotenv(".env") - -from langchain_core.messages import SystemMessage, HumanMessage +from langchain_core.messages import AIMessage, HumanMessage, SystemMessage from langgraph.graph import StateGraph, START, END from langgraph.prebuilt import ToolNode +from pydantic import BaseModel, Field from state import OrchestratorState -from models import get_orchestrator -from agents import ( - implementer_node, - reviewer_node, - researcher_node, - cross_reviewer_node, - scanner_node, - fast_coder_node, -) +from models import get_orchestrator, with_retries +from agents import AGENTS, make_agent_node from tools import get_composio_tools -ROUTER_PROMPT = """You are a task router for Sea Haven Industries. Analyze the incoming task and decide which agent should handle it. - -Available agents: -- implementer: Write new code, add features, fix bugs. Use for any coding task with a clear spec. -- reviewer: Review code changes (diffs, PRs) for correctness, security, maintainability. Uses Claude. -- researcher: Look up documentation, API references, technical questions. Fast and cheap. -- cross_reviewer: Independent code review using a different AI model (GPT). Use when you want a second opinion that catches different blind spots than Claude. -- scanner: Analyze large codebases for patterns, consistency, structural issues. Uses Gemini's large context window. -- fast_coder: Quick, bounded coding for crystal-clear specs. Uses DeepSeek. Best for small, well-defined tasks. -- connector: Interact with external services (Slack, Notion, Google Drive, GitHub) — send messages, read/update pages, find files. -- done: The task is complete or doesn't need agent delegation (e.g., a simple question you can answer directly). - -Respond with ONLY the agent name, nothing else. Pick the single best match.""" +# Load env before any module-level call that reads it (composio init, build_graph). +load_dotenv(".env") -composio_tools = get_composio_tools() +# Route metadata — agents from AGENTS plus the two routes that don't follow the +# simple agent shape. The router can also return "unknown" when nothing fits. +ROUTE_DESCRIPTIONS: dict[str, str] = { + label: cfg["description"] for label, cfg in AGENTS.items() +} +ROUTE_DESCRIPTIONS["connector"] = ( + "Interact with external services (Slack, Notion, Google Drive, GitHub) — " + "send messages, read/update pages, find files." +) +ROUTE_DESCRIPTIONS["done"] = ( + "The task is complete or doesn't need agent delegation " + "(e.g., a simple question you can answer directly)." +) + +VALID_ROUTES: tuple[str, ...] = tuple(ROUTE_DESCRIPTIONS) + ("unknown",) -def router_node(state: OrchestratorState) -> dict: - llm = get_orchestrator() - response = llm.invoke([ - SystemMessage(content=ROUTER_PROMPT), - HumanMessage(content=state["task"]), - ]) - route = response.content.strip().lower() - valid = { +def _build_router_prompt() -> str: + bullets = "\n".join( + f"- {label}: {desc}" for label, desc in ROUTE_DESCRIPTIONS.items() + ) + return ( + "You are a task router for Sea Haven Industries. Analyze the incoming task and " + "decide which agent should handle it.\n\n" + "Available agents:\n" + f"{bullets}\n\n" + 'Pick the single best match. If no option clearly fits, return route="unknown" ' + "rather than guessing. Always include a brief one-line reasoning." + ) + + +ROUTER_PROMPT = _build_router_prompt() + + +class RouteDecision(BaseModel): + route: Literal[ "implementer", "reviewer", "researcher", @@ -52,23 +59,41 @@ def router_node(state: OrchestratorState) -> dict: "fast_coder", "connector", "done", - } - if route not in valid: - route = "researcher" - return {"route": route, "messages": [response]} + "unknown", + ] = Field(description="The agent or terminal route that should handle this task.") + reasoning: str = Field(description="One short sentence explaining the choice.") + + +composio_tools = get_composio_tools() + + +def router_node(state: OrchestratorState) -> dict: + llm = with_retries(get_orchestrator().with_structured_output(RouteDecision)) + decision: RouteDecision = llm.invoke( + [ + SystemMessage(content=ROUTER_PROMPT), + HumanMessage(content=state["task"]), + ] + ) + log_msg = AIMessage( + content=f"router: route={decision.route} reasoning={decision.reasoning}" + ) + return {"route": decision.route, "messages": [log_msg]} def connector_node(state: OrchestratorState) -> dict: - llm = get_orchestrator().bind_tools(composio_tools) - response = llm.invoke([ - SystemMessage( - content=( - "You help interact with external services. Use the available tools to complete the task. " - "Make exactly ONE tool call, then stop. Do not chain multiple calls." - ) - ), - HumanMessage(content=state["task"]), - ]) + llm = with_retries(get_orchestrator().bind_tools(composio_tools)) + response = llm.invoke( + [ + SystemMessage( + content=( + "You help interact with external services. Use the available tools to complete the task. " + "Make exactly ONE tool call, then stop. Do not chain multiple calls." + ) + ), + HumanMessage(content=state["task"]), + ] + ) return {"messages": [response]} @@ -82,6 +107,15 @@ def summarizer_node(state: OrchestratorState) -> dict: return {"result": content} +def unknown_node(state: OrchestratorState) -> dict: + return { + "result": ( + "Router returned 'unknown': no agent clearly fits this task. " + "Rephrase the request or specify an agent explicitly." + ) + } + + def route_task(state: OrchestratorState) -> str: return state["route"] @@ -90,43 +124,29 @@ def build_graph(): graph = StateGraph(OrchestratorState) graph.add_node("router", router_node) - graph.add_node("implementer", implementer_node) - graph.add_node("reviewer", reviewer_node) - graph.add_node("researcher", researcher_node) - graph.add_node("cross_reviewer", cross_reviewer_node) - graph.add_node("scanner", scanner_node) - graph.add_node("fast_coder", fast_coder_node) + for label in AGENTS: + graph.add_node(label, make_agent_node(label)) graph.add_node("connector", connector_node) graph.add_node("tool_executor", ToolNode(composio_tools)) graph.add_node("summarizer", summarizer_node) + graph.add_node("unknown", unknown_node) graph.add_edge(START, "router") - graph.add_conditional_edges( - "router", - route_task, - { - "implementer": "implementer", - "reviewer": "reviewer", - "researcher": "researcher", - "cross_reviewer": "cross_reviewer", - "scanner": "scanner", - "fast_coder": "fast_coder", - "connector": "connector", - "done": END, - }, - ) + conditional_edges = {label: label for label in AGENTS} + conditional_edges["connector"] = "connector" + conditional_edges["done"] = END + conditional_edges["unknown"] = "unknown" - graph.add_edge("implementer", END) - graph.add_edge("reviewer", END) - graph.add_edge("researcher", END) - graph.add_edge("cross_reviewer", END) - graph.add_edge("scanner", END) - graph.add_edge("fast_coder", END) + graph.add_conditional_edges("router", route_task, conditional_edges) + + for label in AGENTS: + graph.add_edge(label, END) graph.add_edge("connector", "tool_executor") graph.add_edge("tool_executor", "summarizer") graph.add_edge("summarizer", END) + graph.add_edge("unknown", END) return graph.compile() @@ -136,7 +156,11 @@ app = build_graph() if __name__ == "__main__": import sys - task = " ".join(sys.argv[1:]) if len(sys.argv) > 1 else "What is the capital of France?" + task = ( + " ".join(sys.argv[1:]) + if len(sys.argv) > 1 + else "What is the capital of France?" + ) result = app.invoke({"task": task, "messages": []}) print(f"\n--- Route: {result['route']} ---") print(result.get("result", "No result")) diff --git a/models.py b/models.py index 1880861..4feb0d4 100644 --- a/models.py +++ b/models.py @@ -1,37 +1,90 @@ import os + from langchain_anthropic import ChatAnthropic -from langchain_openai import ChatOpenAI +from langchain_core.runnables import Runnable from langchain_google_genai import ChatGoogleGenerativeAI +from langchain_openai import ChatOpenAI + +# Model IDs — single source of truth. Bump here when families ship new revs. +CLAUDE_SONNET = "claude-sonnet-4-20250514" +CLAUDE_HAIKU = "claude-haiku-4-5-20251001" +OPENAI_CROSS_REVIEWER = "gpt-4.1" +GEMINI_SCANNER = "gemini-2.5-pro" +DEEPSEEK_FAST_CODER = "deepseek-coder" +DEEPSEEK_BASE_URL = "https://api.deepseek.com/v1" + + +def _collect_retriable_exceptions() -> tuple[type[BaseException], ...]: + excs: list[type[BaseException]] = [] + try: + import anthropic + + excs.extend( + [ + anthropic.APIConnectionError, + anthropic.RateLimitError, + anthropic.InternalServerError, + ] + ) + except ImportError: + pass + try: + import openai + + excs.extend( + [ + openai.APIConnectionError, + openai.RateLimitError, + openai.InternalServerError, + ] + ) + except ImportError: + pass + return tuple(excs) + + +RETRIABLE_EXCEPTIONS = _collect_retriable_exceptions() + + +def with_retries(runnable: Runnable) -> Runnable: + """Wrap an LLM runnable with up to 2 retries on transient provider errors.""" + if not RETRIABLE_EXCEPTIONS: + return runnable + return runnable.with_retry( + retry_if_exception_type=RETRIABLE_EXCEPTIONS, + stop_after_attempt=3, + wait_exponential_jitter=True, + ) def get_orchestrator(): - return ChatAnthropic(model="claude-sonnet-4-20250514", temperature=0) + return ChatAnthropic(model=CLAUDE_SONNET, temperature=0) def get_implementer(): - return ChatAnthropic(model="claude-sonnet-4-20250514", temperature=0) + return ChatAnthropic(model=CLAUDE_SONNET, temperature=0) def get_reviewer(): - return ChatAnthropic(model="claude-sonnet-4-20250514", temperature=0) + return ChatAnthropic(model=CLAUDE_SONNET, temperature=0) def get_researcher(): - return ChatAnthropic(model="claude-haiku-4-5-20251001", temperature=0) + return ChatAnthropic(model=CLAUDE_HAIKU, temperature=0) def get_cross_reviewer(): - return ChatOpenAI(model="gpt-4.1", temperature=0.2) + return ChatOpenAI(model=OPENAI_CROSS_REVIEWER, temperature=0.2) def get_scanner(): - return ChatGoogleGenerativeAI(model="gemini-2.5-pro", temperature=0) + return ChatGoogleGenerativeAI(model=GEMINI_SCANNER, temperature=0) def get_fast_coder(): return ChatOpenAI( - model="deepseek-coder", - base_url="https://api.deepseek.com/v1", + model=DEEPSEEK_FAST_CODER, + base_url=DEEPSEEK_BASE_URL, api_key=os.getenv("DEEPSEEK_API_KEY"), temperature=0, ) diff --git a/run.py b/run.py index 3c2496b..fa34af6 100755 --- a/run.py +++ b/run.py @@ -5,17 +5,16 @@ Usage: python3 run.py "Write a function that validates emails" python3 run.py --route-only "Send a Slack message to #general" """ -import sys -import os -import json -os.chdir(os.path.dirname(os.path.abspath(__file__))) +import os +import sys from dotenv import load_dotenv +os.chdir(os.path.dirname(os.path.abspath(__file__))) load_dotenv(".env") -from graph import app +from graph import app, router_node # noqa: E402 (env must be loaded before graph imports composio) def main(): @@ -28,16 +27,8 @@ def main(): task = " ".join(args) if route_only: - from langchain_core.messages import SystemMessage, HumanMessage - from models import get_orchestrator - - llm = get_orchestrator() - response = llm.invoke([ - SystemMessage(content="You are a task router. Respond with ONLY the agent name.\n" - "Available: implementer, reviewer, researcher, cross_reviewer, scanner, fast_coder, connector, done"), - HumanMessage(content=task), - ]) - print(response.content.strip().lower()) + out = router_node({"task": task, "messages": []}) + print(out["route"]) return result = app.invoke({"task": task, "messages": []}) diff --git a/state.py b/state.py index 5f7fc9f..01251cf 100644 --- a/state.py +++ b/state.py @@ -16,5 +16,6 @@ class OrchestratorState(TypedDict): "fast_coder", "connector", "done", + "unknown", ] result: str diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..8b575ff --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,7 @@ +import os +import sys + +# Add repo root to sys.path so tests can import top-level modules (graph, agents, ...). +ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), "..")) +if ROOT not in sys.path: + sys.path.insert(0, ROOT) diff --git a/tests/test_routing_golden.py b/tests/test_routing_golden.py new file mode 100644 index 0000000..3e318b8 --- /dev/null +++ b/tests/test_routing_golden.py @@ -0,0 +1,115 @@ +"""Golden-set routing test for the structured router. + +20 labelled tasks → expected agent. Skipped if provider keys are missing +(ANTHROPIC for the router LLM, COMPOSIO because importing the graph eagerly +loads tools). Run with: + + pytest tests/test_routing_golden.py -v +""" + +import os + +import pytest +from dotenv import load_dotenv + +load_dotenv(os.path.join(os.path.dirname(__file__), "..", ".env")) + +if not os.getenv("ANTHROPIC_API_KEY"): + pytest.skip( + "ANTHROPIC_API_KEY not set; skipping live router tests.", + allow_module_level=True, + ) +if not os.getenv("COMPOSIO_API_KEY"): + pytest.skip( + "COMPOSIO_API_KEY not set; skipping live router tests.", allow_module_level=True + ) + +from graph import router_node # noqa: E402 + + +GOLDEN_SET: list[tuple[str, str]] = [ + # implementer (3) + ( + "Write a Python function that validates email addresses using a regex", + "implementer", + ), + ( + "Add a new Lambda handler in handlers/notify.py that publishes an SNS message", + "implementer", + ), + ( + "Fix the bug in our auth middleware where expired tokens are accepted as valid", + "implementer", + ), + # reviewer (3) + ( + "Review this pull request diff for correctness, security, and maintainability concerns", + "reviewer", + ), + ( + "Code review the attached commit and categorize each issue as BLOCK, FIX, or NIT", + "reviewer", + ), + ( + "Review the following code changes and tell me what should block merge", + "reviewer", + ), + # researcher (3) + ( + "What is the latest stable version of the langgraph Python package?", + "researcher", + ), + ( + "Look up the AWS Lambda maximum concurrent execution limit in us-east-1", + "researcher", + ), + ("Find documentation on how to configure DynamoDB TTL", "researcher"), + # cross_reviewer (2) + ( + "Get a cross-family second opinion on this diff using GPT to catch what Claude might miss", + "cross_reviewer", + ), + ( + "Run an independent cross-model review on this Lambda handler change", + "cross_reviewer", + ), + # scanner (3) + ( + "Scan the entire monorepo to find inconsistent error handling patterns across 200+ files", + "scanner", + ), + ( + "Analyze the whole codebase for unused imports and dead code using a large-context model", + "scanner", + ), + ("Audit every Lambda handler in this repo for hardcoded secrets", "scanner"), + # fast_coder (3) + ( + "Quick small task using DeepSeek: write a 5-line Python helper that converts kebab-case to snake_case", + "fast_coder", + ), + ( + "Fast bounded coding job: implement a one-function utility that pads strings to a fixed width", + "fast_coder", + ), + ( + "Use the cheap fast coder to write a short Python snippet that parses a CSV row into a dict", + "fast_coder", + ), + # connector (3) + ("Send a Slack message to the #ops channel announcing the deploy", "connector"), + ("Create a Notion page under the Engineering space called 'Q3 plan'", "connector"), + ("Find a file named contracts.pdf in my Google Drive", "connector"), +] + + +def test_golden_set_size(): + assert len(GOLDEN_SET) == 20 + + +@pytest.mark.parametrize("task,expected", GOLDEN_SET) +def test_router_picks_expected_agent(task: str, expected: str): + out = router_node({"task": task, "messages": []}) + assert out["route"] == expected, ( + f"task={task!r} got={out['route']} expected={expected}" + )