Add retrieval-augmented routing tests, fix lazy loading and open items #2-3
- Add 6 retrieval-augmented routing tests (3 live retrieval, 3 off-topic fake memories) to unblock Phase 5 - Defer Composio tool loading and graph construction to first use so expired or missing keys don't crash imports - Atomic cache write in retriever via temp file (open item #2) - Log rotation in weekly_summary.py, pruning JSONL >90 days (open item #3)
This commit is contained in:
parent
4d4381b5eb
commit
c23ea679e4
6 changed files with 115 additions and 11 deletions
|
|
@ -41,7 +41,7 @@ System updates:
|
|||
|
||||
### 1. Retrieval-augmented routing test — **DO BEFORE PHASE 5**
|
||||
|
||||
**Status:** open, ~30 min effort.
|
||||
**Status:** DONE (2026-05-26). Added 6 tests: 3 with live retrieval, 3 with off-topic fake memories. Also fixed Composio lazy loading so tests don't crash on expired keys.
|
||||
|
||||
**What's missing.** The golden-set test (`tests/test_routing_golden.py`) calls `router_node` with no `retrieved` key, so all 21 cases exercise the router on a bare prompt without memory injection. We currently have zero test signal on whether retrieval flips a routing decision.
|
||||
|
||||
|
|
@ -84,7 +84,7 @@ def test_router_unchanged_with_offtopic_memory():
|
|||
|
||||
### 2. Non-atomic cache write — **LOW**
|
||||
|
||||
**Status:** open, ~5 min effort.
|
||||
**Status:** DONE (2026-05-26). Atomic write via temp file in `retriever._save_cache`.
|
||||
|
||||
**What's wrong.** `retriever._save_cache` calls `CACHE_FILE.write_text(...)`, which truncates the file to zero before writing. A crash, Ctrl-C, or kill during the ~50ms write window leaves `.cache/embeddings.json` empty or half-written. Next run, `_load_cache` hits `JSONDecodeError`, returns `{}`, and the retriever re-embeds all 87 memories (~10s + ~$0.001).
|
||||
|
||||
|
|
@ -108,7 +108,7 @@ def _save_cache(cache: dict) -> None:
|
|||
|
||||
### 3. Telemetry log rotation — **LOW**
|
||||
|
||||
**Status:** open, ~5 min added to the digest `/schedule`.
|
||||
**Status:** DONE (2026-05-26). `prune_old_logs()` added to `weekly_summary.py`; deletes JSONL files >90 days old, runs automatically after each digest.
|
||||
|
||||
**What's wrong.** `~/.claude/logs/orchestrator/YYYY-MM-DD.jsonl` files accumulate forever. No deletion, no compression. At current volume (20 runs/day) the directory grows ~2 MB/year — disk is not the issue, but file count climbs and the digest only reads the last 7 days, so older logs are noise.
|
||||
|
||||
|
|
|
|||
25
graph.py
25
graph.py
|
|
@ -65,7 +65,14 @@ class RouteDecision(BaseModel):
|
|||
reasoning: str = Field(description="One short sentence explaining the choice.")
|
||||
|
||||
|
||||
composio_tools = get_composio_tools()
|
||||
_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
|
||||
|
||||
|
||||
def retriever_node(state: OrchestratorState) -> dict:
|
||||
|
|
@ -105,7 +112,7 @@ def router_node(state: OrchestratorState) -> dict:
|
|||
|
||||
|
||||
def connector_node(state: OrchestratorState) -> dict:
|
||||
llm = with_retries(get_orchestrator().bind_tools(composio_tools))
|
||||
llm = with_retries(get_orchestrator().bind_tools(_get_tools()))
|
||||
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."
|
||||
|
|
@ -152,7 +159,7 @@ def build_graph():
|
|||
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("tool_executor", ToolNode(_get_tools()))
|
||||
graph.add_node("summarizer", summarizer_node)
|
||||
graph.add_node("unknown", unknown_node)
|
||||
|
||||
|
|
@ -177,7 +184,15 @@ def build_graph():
|
|||
return graph.compile()
|
||||
|
||||
|
||||
app = build_graph()
|
||||
_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
|
||||
|
|
@ -187,6 +202,6 @@ if __name__ == "__main__":
|
|||
if len(sys.argv) > 1
|
||||
else "What is the capital of France?"
|
||||
)
|
||||
result = app.invoke({"task": task, "messages": []})
|
||||
result = get_app().invoke({"task": task, "messages": []})
|
||||
print(f"\n--- Route: {result['route']} ---")
|
||||
print(result.get("result", "No result"))
|
||||
|
|
|
|||
|
|
@ -72,7 +72,9 @@ def _load_cache() -> dict:
|
|||
|
||||
def _save_cache(cache: dict) -> None:
|
||||
CACHE_DIR.mkdir(exist_ok=True)
|
||||
CACHE_FILE.write_text(json.dumps(cache))
|
||||
tmp = CACHE_FILE.with_suffix(".json.tmp")
|
||||
tmp.write_text(json.dumps(cache))
|
||||
tmp.replace(CACHE_FILE)
|
||||
|
||||
|
||||
def get_or_build_embeddings(memories: list[Memory]) -> dict[str, list[float]]:
|
||||
|
|
|
|||
4
run.py
4
run.py
|
|
@ -15,7 +15,7 @@ from dotenv import load_dotenv
|
|||
os.chdir(os.path.dirname(os.path.abspath(__file__)))
|
||||
load_dotenv(".env")
|
||||
|
||||
from graph import app, retriever_node, router_node # noqa: E402 (env must be loaded before graph imports composio)
|
||||
from graph import get_app, retriever_node, router_node # noqa: E402 (env must be loaded before graph imports composio)
|
||||
from telemetry import build_record, log_run # noqa: E402
|
||||
|
||||
|
||||
|
|
@ -47,7 +47,7 @@ def main():
|
|||
result: dict | None = None
|
||||
error: str | None = None
|
||||
try:
|
||||
result = app.invoke({"task": task, "messages": []})
|
||||
result = get_app().invoke({"task": task, "messages": []})
|
||||
except Exception as exc:
|
||||
error = f"{type(exc).__name__}: {exc}"
|
||||
finished = time.monotonic()
|
||||
|
|
|
|||
|
|
@ -119,5 +119,30 @@ def summarize(records: list[dict], window_days: int = WINDOW_DAYS) -> str:
|
|||
return "\n".join(lines)
|
||||
|
||||
|
||||
RETENTION_DAYS = 90
|
||||
|
||||
|
||||
def prune_old_logs(retention_days: int = RETENTION_DAYS) -> int:
|
||||
"""Delete JSONL logs older than retention_days. Returns count deleted."""
|
||||
if not LOG_DIR.is_dir():
|
||||
return 0
|
||||
cutoff = datetime.datetime.now(datetime.UTC).date() - datetime.timedelta(
|
||||
days=retention_days
|
||||
)
|
||||
deleted = 0
|
||||
for path in LOG_DIR.glob("*.jsonl"):
|
||||
try:
|
||||
file_date = datetime.date.fromisoformat(path.stem)
|
||||
except ValueError:
|
||||
continue
|
||||
if file_date < cutoff:
|
||||
path.unlink()
|
||||
deleted += 1
|
||||
return deleted
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
print(summarize(load_recent_records()))
|
||||
pruned = prune_old_logs()
|
||||
if pruned:
|
||||
print(f"\n_Pruned {pruned} log file(s) older than {RETENTION_DAYS} days._")
|
||||
|
|
|
|||
|
|
@ -22,6 +22,13 @@ requires_keys = pytest.mark.skipif(
|
|||
reason=f"Requires {', '.join(_REQUIRED_KEYS)}; missing: {', '.join(_missing)}",
|
||||
)
|
||||
|
||||
_RETRIEVAL_KEYS = ("ANTHROPIC_API_KEY", "COMPOSIO_API_KEY", "OPENAI_API_KEY")
|
||||
_missing_retrieval = [k for k in _RETRIEVAL_KEYS if not os.getenv(k)]
|
||||
requires_keys_with_retrieval = pytest.mark.skipif(
|
||||
bool(_missing_retrieval),
|
||||
reason=f"Requires {', '.join(_RETRIEVAL_KEYS)}; missing: {', '.join(_missing_retrieval)}",
|
||||
)
|
||||
|
||||
|
||||
GOLDEN_SET: list[tuple[str, str]] = [
|
||||
# implementer (3)
|
||||
|
|
@ -114,3 +121,58 @@ def test_router_picks_expected_agent(task: str, expected: str):
|
|||
assert out["route"] == expected, (
|
||||
f"task={task!r} got={out['route']} expected={expected}"
|
||||
)
|
||||
|
||||
|
||||
# --- Retrieval-augmented routing tests (Phase 5 gate) ---
|
||||
|
||||
RETRIEVAL_CASES: list[tuple[str, str]] = [
|
||||
("Send a Slack message to the #ops channel about the deploy", "connector"),
|
||||
(
|
||||
"Review this pull request diff for correctness and security",
|
||||
"reviewer",
|
||||
),
|
||||
(
|
||||
"Scan the entire repo for hardcoded AWS credentials across all files",
|
||||
"scanner",
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
@requires_keys_with_retrieval
|
||||
@pytest.mark.parametrize("task,expected", RETRIEVAL_CASES)
|
||||
def test_router_stable_with_live_retrieval(task: str, expected: str):
|
||||
"""Route stays correct when real retrieved memories are injected."""
|
||||
from graph import router_node
|
||||
from retriever import retrieve
|
||||
|
||||
retrieved = retrieve(task, k=3)
|
||||
out = router_node({"task": task, "messages": [], "retrieved": retrieved})
|
||||
assert out["route"] == expected, (
|
||||
f"task={task!r} got={out['route']} expected={expected} "
|
||||
f"retrieved={[m['name'] for m in retrieved]}"
|
||||
)
|
||||
|
||||
|
||||
@requires_keys
|
||||
@pytest.mark.parametrize("task,expected", RETRIEVAL_CASES)
|
||||
def test_router_stable_with_offtopic_memory(task: str, expected: str):
|
||||
"""Off-topic retrieved memories don't mislead the router."""
|
||||
from graph import router_node
|
||||
|
||||
fake_retrieved = [
|
||||
{
|
||||
"name": "feedback_emulator_coordinates",
|
||||
"score": 0.1,
|
||||
"content": "ADB input uses landscape coords; getevent uses portrait.",
|
||||
},
|
||||
{
|
||||
"name": "project_unrelated_game",
|
||||
"score": 0.05,
|
||||
"content": "Unity project for a 2D platformer. Uses C# and .NET 8.",
|
||||
},
|
||||
]
|
||||
out = router_node({"task": task, "messages": [], "retrieved": fake_retrieved})
|
||||
assert out["route"] == expected, (
|
||||
f"task={task!r} got={out['route']} expected={expected} "
|
||||
f"with off-topic memories injected"
|
||||
)
|
||||
|
|
|
|||
Reference in a new issue