This repository has been archived on 2026-08-04. You can view files and clone it, but cannot push or open issues or pull requests.
orchestrator/graph.py

208 lines
6.6 KiB
Python
Raw Normal View History

from typing import Literal
from dotenv import load_dotenv
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, with_retries
from agents import AGENTS, make_agent_node
2026-05-15 11:36:03 -04:00
from retriever import format_memories_for_prompt, retrieve
from tools import get_composio_tools
# Load env before any module-level call that reads it (composio init, build_graph).
load_dotenv(".env")
# 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 _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",
"cross_reviewer",
"scanner",
"fast_coder",
"connector",
"done",
"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_cache: list | None = None
def _get_tools():
global _composio_tools_cache
if _composio_tools_cache is None:
_composio_tools_cache = get_composio_tools()
return _composio_tools_cache
2026-05-15 11:36:03 -04:00
def retriever_node(state: OrchestratorState) -> dict:
"""Retrieve top-3 relevant memories for the task. Failures are non-fatal —
if retrieval errors out, the run continues with no memory context."""
try:
retrieved = retrieve(state["task"], k=3)
except Exception as exc:
log = AIMessage(
content=f"retriever: error ({type(exc).__name__}: {exc}); continuing without memory"
)
return {"retrieved": [], "messages": [log]}
names = ", ".join(m["name"] for m in retrieved) or "none"
log = AIMessage(content=f"retriever: {names}")
return {"retrieved": retrieved, "messages": [log]}
def _router_system_prompt(retrieved: list[dict] | None) -> str:
memory_block = format_memories_for_prompt(retrieved or [])
if not memory_block:
return ROUTER_PROMPT
return f"{ROUTER_PROMPT}\n\n{memory_block}"
def router_node(state: OrchestratorState) -> dict:
llm = with_retries(get_orchestrator().with_structured_output(RouteDecision))
decision: RouteDecision = llm.invoke(
[
2026-05-15 11:36:03 -04:00
SystemMessage(content=_router_system_prompt(state.get("retrieved"))),
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 = with_retries(get_orchestrator().bind_tools(_get_tools()))
2026-05-15 11:36:03 -04:00
base_prompt = (
"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."
)
memory_block = format_memories_for_prompt(state.get("retrieved") or [])
system_content = f"{base_prompt}\n\n{memory_block}" if memory_block else base_prompt
response = llm.invoke(
[
2026-05-15 11:36:03 -04:00
SystemMessage(content=system_content),
HumanMessage(content=state["task"]),
]
)
return {"messages": [response]}
def summarizer_node(state: OrchestratorState) -> dict:
last_msg = state["messages"][-1]
content = last_msg.content if hasattr(last_msg, "content") else str(last_msg)
if isinstance(content, list):
content = "\n".join(str(c) for c in content)
if len(content) > 2000:
content = content[:2000] + "...(truncated)"
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"]
def build_graph():
graph = StateGraph(OrchestratorState)
2026-05-15 11:36:03 -04:00
graph.add_node("retriever", retriever_node)
graph.add_node("router", router_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(_get_tools()))
graph.add_node("summarizer", summarizer_node)
graph.add_node("unknown", unknown_node)
2026-05-15 11:36:03 -04:00
graph.add_edge(START, "retriever")
graph.add_edge("retriever", "router")
conditional_edges = {label: label for label in AGENTS}
conditional_edges["connector"] = "connector"
conditional_edges["done"] = END
conditional_edges["unknown"] = "unknown"
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()
_app_cache = None
def get_app():
global _app_cache
if _app_cache is None:
_app_cache = build_graph()
return _app_cache
if __name__ == "__main__":
import sys
task = (
" ".join(sys.argv[1:])
if len(sys.argv) > 1
else "What is the capital of France?"
)
result = get_app().invoke({"task": task, "messages": []})
print(f"\n--- Route: {result['route']} ---")
print(result.get("result", "No result"))