fix: issues with local tool calling on non openai endpoint

This commit is contained in:
Alex committed 2026-06-03 22:50:03 +01:00
1 parent 8ec9474cc6
commit fbd37e627c
6 files changed
+258 -18

No files matched your search

+27 -18
View File
@@ -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,
@@ -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:
+6
View File
@@ -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"]
+89
View File
@@ -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
+88
View File
@@ -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:
+32
View File
@@ -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 = {