mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-05 04:13:25 +00:00
fix: issues with local tool calling on non openai endpoint
This commit is contained in:
1 parent
8ec9474cc6
commit
fbd37e627c
6 files changed
+258
-18
No files matched your search
@@ -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:
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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 = {
|
||||
|
||||
Reference in new issue
Block a user