diff --git a/application/api/answer/routes/base.py b/application/api/answer/routes/base.py index 32852edf..2f95d10c 100644 --- a/application/api/answer/routes/base.py +++ b/application/api/answer/routes/base.py @@ -588,24 +588,34 @@ class BaseAnswerResource: exc_info=True, ) - # Notify the user out-of-band so they can navigate - # back to the conversation and decide on the - # pending tool calls. Gated on ``state_saved``: a - # missing pending_tool_state row would 404 the - # resume endpoint, so an unfulfillable notification - # is worse than no notification. + # Notify the user out-of-band so they can navigate back and + # resolve the pause. Only ``awaiting_approval`` pauses need a + # human; ``requires_client_execution`` pauses are resolved by + # the client, so notifying for those is non-actionable noise. + # Also gated on ``state_saved``: a missing pending_tool_state + # row would 404 the resume endpoint. user_id_for_event = ( decoded_token.get("sub") if decoded_token else None ) - if state_saved and user_id_for_event and conversation_id: - pending_calls = continuation.get( - "pending_tool_calls", [] - ) if continuation else [] - # Trim each pending tool call to its identifying - # metadata so a tool with a multi-MB argument - # doesn't blow out the per-event payload size - # cap. The resume page fetches full args from - # ``pending_tool_state`` regardless. + approval_calls = [ + tc + for tc in ( + continuation.get("pending_tool_calls", []) + if continuation + else [] + ) + if isinstance(tc, dict) + and tc.get("pause_type") == "awaiting_approval" + ] + if ( + state_saved + and user_id_for_event + and conversation_id + and approval_calls + ): + # Trim each pending tool call to its identifying metadata + # so a multi-MB argument can't blow out the per-event + # payload cap. Full args come from pending_tool_state. pending_summaries = [ { k: tc.get(k) @@ -615,10 +625,9 @@ class BaseAnswerResource: "action_name", "name", ) - if isinstance(tc, dict) and tc.get(k) is not None + if tc.get(k) is not None } - for tc in (pending_calls or []) - if isinstance(tc, dict) + for tc in approval_calls ] publish_user_event( user_id_for_event, diff --git a/application/api/answer/services/stream_processor.py b/application/api/answer/services/stream_processor.py index fd212205..e923b26a 100644 --- a/application/api/answer/services/stream_processor.py +++ b/application/api/answer/services/stream_processor.py @@ -1004,6 +1004,22 @@ class StreamProcessor: from application.llm.handlers.handler_creator import LLMHandlerCreator from application.llm.llm_creator import LLMCreator + # api_key-in-body auth carries no JWT, so initial_user_id is None — but + # the state was saved under the agent owner. Resolve the owner so the + # lookup / mark_resuming / delete_state key on the same id. (No-op for + # v1, which already passes an owner-scoped decoded_token.) + if self.initial_user_id is None and self.data.get("api_key"): + with db_readonly() as conn: + agent_doc = AgentsRepository(conn).find_by_key(self.data["api_key"]) + owner = ( + (agent_doc.get("user_id") or agent_doc.get("user")) + if agent_doc + else None + ) + if owner: + self.initial_user_id = owner + self.decoded_token = {"sub": owner} + cont_service = ContinuationService() state = cont_service.load_state(conversation_id, self.initial_user_id) if not state: diff --git a/application/api/v1/translator.py b/application/api/v1/translator.py index 5330a029..a4d48f0c 100644 --- a/application/api/v1/translator.py +++ b/application/api/v1/translator.py @@ -146,6 +146,12 @@ def translate_request( "tool_actions": tool_actions, "api_key": api_key, } + # A continuation only exists if turn 1 was saved, so default to True — + # otherwise the final turn and its WAL row are never persisted. An + # explicit override is honoured if the client resends it. + result["save_conversation"] = bool( + data.get("docsgpt", {}).get("save_conversation", True) + ) # Carry tools forward for next iteration if data.get("tools"): result["client_tools"] = data["tools"] diff --git a/tests/api/answer/routes/test_base.py b/tests/api/answer/routes/test_base.py index 7c132858..08bcf2ea 100644 --- a/tests/api/answer/routes/test_base.py +++ b/tests/api/answer/routes/test_base.py @@ -308,6 +308,95 @@ class TestCompleteStreamMethod: assert seen_conv_id_on_gen["value"] == fresh_conv_id assert tool_executor.conversation_id == fresh_conv_id + def _run_paused(self, resource, pending_calls): + """Drive complete_stream into its paused branch with the given pending + tool calls, mocking out the WAL row and continuation save.""" + agent = MagicMock() + agent.gen.return_value = iter( + [{"type": "tool_calls_pending", "data": {"pending_tool_calls": pending_calls}}] + ) + agent._pending_continuation = { + "messages": [], + "pending_tool_calls": pending_calls, + "tools_dict": {}, + } + # Make the WAL reservation no-op so we stay off the journal/DB. + resource.conversation_service = MagicMock() + resource.conversation_service.save_user_question.side_effect = Exception("skip") + list( + resource.complete_stream( + question="Do the test", + agent=agent, + conversation_id="conv-1", + user_api_key=None, + decoded_token={"sub": "user123"}, + should_save_conversation=True, + ) + ) + + def test_paused_skips_notification_for_client_execution( + self, mock_mongo_db, flask_app + ): + """A pure ``requires_client_execution`` pause must NOT publish a + ``tool.approval.required`` event — the client resolves it, so the + notification would be non-actionable noise.""" + from application.api.answer.routes import base as base_mod + + with flask_app.app_context(), patch.object( + base_mod, "publish_user_event" + ) as published, patch.object( + base_mod, "ContinuationService", lambda: MagicMock() + ): + self._run_paused( + base_mod.BaseAnswerResource(), + [ + { + "call_id": "c1", + "name": "create_file", + "tool_name": "create_file", + "action_name": "create_file", + "pause_type": "requires_client_execution", + } + ], + ) + published.assert_not_called() + + def test_paused_publishes_notification_only_for_awaiting_approval( + self, mock_mongo_db, flask_app + ): + """A pause with an ``awaiting_approval`` call publishes once, and the + payload surfaces only the approval call (not the client-side one).""" + from application.api.answer.routes import base as base_mod + + with flask_app.app_context(), patch.object( + base_mod, "publish_user_event" + ) as published, patch.object( + base_mod, "ContinuationService", lambda: MagicMock() + ): + self._run_paused( + base_mod.BaseAnswerResource(), + [ + { + "call_id": "a1", + "name": "delete_thing", + "tool_name": "api_tool", + "action_name": "delete_thing", + "pause_type": "awaiting_approval", + }, + { + "call_id": "c1", + "name": "create_file", + "tool_name": "create_file", + "action_name": "create_file", + "pause_type": "requires_client_execution", + }, + ], + ) + published.assert_called_once() + args, _ = published.call_args + assert args[1] == "tool.approval.required" + summaries = args[2]["pending_tool_calls"] + assert [s["call_id"] for s in summaries] == ["a1"] @pytest.mark.unit diff --git a/tests/test_continuation.py b/tests/test_continuation.py index d7304fc3..d068e545 100644 --- a/tests/test_continuation.py +++ b/tests/test_continuation.py @@ -983,6 +983,94 @@ class TestResumeMarkResuming: assert sp.reserved_message_id == reserved_id + def test_resume_resolves_owner_from_api_key_when_no_jwt(self, monkeypatch): + """api_key-authenticated resumes (no JWT) must resolve the agent owner + before loading state. + + On the native /stream and /api/answer routes the agent key lives in the + request body, so ``request.decoded_token`` — and hence + ``initial_user_id`` — is None. The pending state was saved under the + owner's id during the first turn, so the resume has to resolve the owner + here or the lookup misses and the run 400s with "No pending tool state + found for this conversation". + """ + from contextlib import contextmanager + + from application.api.answer.services import ( + continuation_service as cont_mod, + ) + from application.api.answer.services import stream_processor as sp_mod + from application.llm import llm_creator as llm_creator_mod + from application.llm.handlers import handler_creator as handler_mod + + cont_service = MagicMock() + cont_service.load_state.return_value = { + "messages": [], + "pending_tool_calls": [], + "tools_dict": {}, + "tool_schemas": [], + "agent_config": { + "model_id": "m1", + "model_user_id": None, + "llm_name": "openai", + "api_key": "k", + "user_api_key": None, + "agent_id": None, + "agent_type": "ClassicAgent", + "prompt": "", + "json_schema": None, + "retriever_config": None, + }, + "client_tools": None, + } + cont_service.mark_resuming.return_value = True + monkeypatch.setattr(cont_mod, "ContinuationService", lambda: cont_service) + monkeypatch.setattr( + llm_creator_mod.LLMCreator, "create_llm", lambda *a, **kw: MagicMock(), + ) + monkeypatch.setattr( + handler_mod.LLMHandlerCreator, "create_handler", + lambda *a, **kw: MagicMock(), + ) + from application.agents import agent_creator as ac_mod + from application.agents import tool_executor as te_mod + + monkeypatch.setattr( + te_mod, "ToolExecutor", lambda **kw: MagicMock(client_tools=None) + ) + monkeypatch.setattr( + ac_mod.AgentCreator, "create_agent", lambda *a, **kw: MagicMock() + ) + + # The body api_key resolves to its owning user. + fake_repo = MagicMock() + fake_repo.find_by_key.return_value = {"user_id": "owner-1"} + + @contextmanager + def _fake_db_readonly(): + yield MagicMock() + + monkeypatch.setattr(sp_mod, "db_readonly", _fake_db_readonly) + monkeypatch.setattr(sp_mod, "AgentsRepository", lambda conn: fake_repo) + + conv_id = "00000000-0000-0000-0000-000000000009" + sp = sp_mod.StreamProcessor.__new__(sp_mod.StreamProcessor) + sp.data = {"api_key": "agent-key-1"} + sp.decoded_token = None + sp.initial_user_id = None + sp.conversation_id = conv_id + sp.agent_config = {} + sp.reserved_message_id = None + + sp.resume_from_tool_actions(tool_actions=[], conversation_id=conv_id) + + fake_repo.find_by_key.assert_called_once_with("agent-key-1") + # The lookup + claim now key on the owner id, not None. + cont_service.load_state.assert_called_once_with(conv_id, "owner-1") + cont_service.mark_resuming.assert_called_once_with(conv_id, "owner-1") + assert sp.initial_user_id == "owner-1" + assert sp.decoded_token == {"sub": "owner-1"} + @pytest.mark.unit class TestContinuationServiceMarkResuming: diff --git a/tests/test_v1_translator.py b/tests/test_v1_translator.py index 91157e59..e8de8d8f 100644 --- a/tests/test_v1_translator.py +++ b/tests/test_v1_translator.py @@ -245,6 +245,38 @@ class TestTranslateRequest: assert len(result["tool_actions"]) == 1 assert result["tool_actions"][0]["call_id"] == "c1" + def test_continuation_persists_by_default(self): + """A continuation implies the first turn was saved, so the resumed turn + must persist too (otherwise the final answer + WAL row are lost).""" + data = { + "messages": [ + {"role": "user", "content": "Search for X"}, + { + "role": "assistant", + "tool_calls": [{"id": "c1", "type": "function", "function": {"name": "search", "arguments": "{}"}}], + }, + {"role": "tool", "tool_call_id": "c1", "content": "done"}, + ], + } + result = translate_request(data, "key") + assert result["save_conversation"] is True + + def test_continuation_honours_explicit_save_conversation_override(self): + """An explicit docsgpt.save_conversation=false on the continuation wins.""" + data = { + "docsgpt": {"save_conversation": False}, + "messages": [ + {"role": "user", "content": "Search for X"}, + { + "role": "assistant", + "tool_calls": [{"id": "c1", "type": "function", "function": {"name": "search", "arguments": "{}"}}], + }, + {"role": "tool", "tool_call_id": "c1", "content": "done"}, + ], + } + result = translate_request(data, "key") + assert result["save_conversation"] is False + def test_continuation_with_top_level_conversation_id(self): """Standard clients send conversation_id at request level, not in messages.""" data = {