diff --git a/s16code/routes.py b/s16code/routes.py index ee09f81..5fd4988 100644 --- a/s16code/routes.py +++ b/s16code/routes.py @@ -163,15 +163,33 @@ def _waiting_channel_approval(request: Request, body: ChannelMessageBody): return None +_NEGATIVE_APPROVAL_CHOICES = frozenset({"no", "n", "reject", "deny", "cancel"}) + + +def _channel_approval_choice(text: str, wait: dict[str, Any]) -> str | None: + """Return the matched listed choice, or None if this message is not an answer.""" + answer = text.casefold().strip() + if not answer: + return None + choices = [str(choice).casefold().strip() for choice in (wait.get("choices") or []) if str(choice).strip()] + if not choices: + return answer + return answer if answer in choices else None + + async def _resume_channel_approval(request: Request, body: ChannelMessageBody): pending = _waiting_channel_approval(request, body) if pending is None: return None run_id, wait = pending + matched = _channel_approval_choice(body.text or "", wait) + if matched is None: + return None completion = request.app.state.runtime.graph.complete_waiting( wait["handle"], "approval.received", {"response": body.text or "", "channel": body.channel, "channel_user_id": body.channel_user_id, "thread_id": body.thread_id}, + success=matched not in _NEGATIVE_APPROVAL_CHOICES, ) if completion is None: return None diff --git a/tests/test_channel_connections.py b/tests/test_channel_connections.py index 4cb238b..5ec29d2 100644 --- a/tests/test_channel_connections.py +++ b/tests/test_channel_connections.py @@ -122,7 +122,7 @@ async def gateway(request: httpx.Request) -> httpx.Response: await http.aclose() -def test_a_reply_in_the_same_channel_thread_resumes_a_waiting_approval(app_client, monkeypatch): +def _park_yes_no_approval(app_client, monkeypatch, origin_id: str): monkeypatch.setenv("S16_CHANNEL_BRIDGE_TOKEN", "shared") monkeypatch.setenv("S16_CHANNEL_ALLOWED_SIDE_EFFECTS", "request_approval") app_client.app.state.runtime.memory.embedder = DeterministicEmbedder(128) @@ -150,9 +150,16 @@ async def fake_gateway(_app, prompt: str, system: str): monkeypatch.setattr(agent_route, "gateway_text_llm", fake_gateway) headers = {"Authorization": "Bearer shared"} first = app_client.post("/v1/agent/channel-messages", headers=headers, - json=_channel_message(metadata={"message_id": "approval-1"})) + json=_channel_message(metadata={"message_id": origin_id})) assert first.status_code == 200 assert first.json()["text"] == "Approval needed: Send the final report? Choices: yes, no" + events = app_client.get("/v1/agent/events").json()["events"] + run_id = events[-1]["decisions"][0]["run_id"] + return headers, run_id + + +def test_a_reply_in_the_same_channel_thread_resumes_a_waiting_approval(app_client, monkeypatch): + headers, _run_id = _park_yes_no_approval(app_client, monkeypatch, "approval-1") second = app_client.post("/v1/agent/channel-messages", headers=headers, json=_channel_message(text="yes", metadata={"message_id": "approval-2"})) @@ -165,6 +172,29 @@ async def fake_gateway(_app, prompt: str, system: str): assert any(event["kind"] == "external_event_received" for event in journal) +def test_a_no_does_not_succeed_a_waiting_channel_approval(app_client, monkeypatch): + headers, run_id = _park_yes_no_approval(app_client, monkeypatch, "approval-no-1") + second = app_client.post("/v1/agent/channel-messages", headers=headers, + json=_channel_message(text="no", metadata={"message_id": "approval-no-2"})) + assert second.status_code == 200 + journal = app_client.get(f"/v1/agent/runs/{run_id}").json() + assert journal["nodes"]["approve"]["state"] == "failed" + + +def test_unrelated_text_does_not_consume_a_waiting_channel_approval(app_client, monkeypatch): + headers, run_id = _park_yes_no_approval(app_client, monkeypatch, "approval-other-1") + second = app_client.post( + "/v1/agent/channel-messages", headers=headers, + json=_channel_message(text="What's on my calendar tomorrow?", + metadata={"message_id": "approval-other-2"}), + ) + assert second.status_code == 200 + journal = app_client.get(f"/v1/agent/runs/{run_id}").json() + assert journal["nodes"]["approve"]["state"] == "waiting" + events = app_client.get("/v1/agent/events").json()["events"] + assert events[-1]["decisions"][0]["subscription_id"] != "channel-approval" + + def test_job_callback_resumes_and_pushes_final_answer_to_originating_channel(app_client, monkeypatch): monkeypatch.setenv("S16_CHANNEL_BRIDGE_TOKEN", "shared") app_client.app.state.runtime.memory.embedder = DeterministicEmbedder(128)