456 lines
17 KiB
Python
456 lines
17 KiB
Python
"""In-app agent loop — multi-step tool calling.
|
|
|
|
The loop drives an LLM that may request tool calls, executes them through the
|
|
shared tool layer (``backend.tools``), feeds the results back, and repeats
|
|
until the model produces a final answer or the iteration budget is exhausted.
|
|
|
|
The LLM is injected as an async callable so the loop is fully testable without
|
|
network access::
|
|
|
|
async def fake_llm(messages, tools):
|
|
return LLMResponse(content="done")
|
|
|
|
result = await run_agent(messages, ctx=ctx, llm=fake_llm)
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import logging
|
|
import os
|
|
from collections.abc import Callable
|
|
from dataclasses import dataclass, field
|
|
from typing import Any
|
|
|
|
from backend.tools.api import (
|
|
ToolConfirmationRequired,
|
|
ToolContext,
|
|
ToolError,
|
|
ToolScope,
|
|
call_tool,
|
|
get_tool,
|
|
get_tool_schemas,
|
|
)
|
|
from backend.tools.labels import thought_step_label, tool_step_label
|
|
|
|
logger = logging.getLogger("obsigate.agent.loop")
|
|
|
|
DEFAULT_MAX_ITERATIONS = 10
|
|
# Cap the size of a tool result fed back to the model (chars).
|
|
MAX_TOOL_RESULT_CHARS = 100_000
|
|
# Quota: maximum tool calls executed per agent run (``BOOKSLM_MAX_TOOL_CALLS``).
|
|
DEFAULT_MAX_TOOL_CALLS = int(os.environ.get("BOOKSLM_MAX_TOOL_CALLS", "25"))
|
|
|
|
# Sent as a last user turn when the loop stopped before the model produced an
|
|
# answer (iteration/quota budget exhausted while it was still calling tools).
|
|
_FINALIZE_INSTRUCTION = (
|
|
"N'appelle plus aucun outil. Réponds maintenant directement à l'utilisateur, "
|
|
"en français, à partir des informations déjà recueillies ci-dessus. "
|
|
"Structure la réponse en Markdown, cite les liens sources utiles, et si les "
|
|
"informations sont insuffisantes, dis-le explicitement."
|
|
)
|
|
|
|
# Stopping reasons
|
|
STOP_DONE = "done"
|
|
STOP_MAX_ITERATIONS = "max_iterations"
|
|
STOP_CONFIRMATION_REQUIRED = "confirmation_required"
|
|
STOP_QUOTA_EXCEEDED = "quota_exceeded"
|
|
|
|
|
|
@dataclass
|
|
class ToolCallRecord:
|
|
"""Audit-friendly record of one executed tool call."""
|
|
|
|
name: str
|
|
arguments: dict[str, Any]
|
|
ok: bool
|
|
result: Any
|
|
# Human-readable « step » label for the Notion-style UI
|
|
# ({key, params} — see backend.tools.labels).
|
|
step: dict[str, Any] = field(default_factory=dict)
|
|
|
|
|
|
@dataclass
|
|
class AgentResult:
|
|
"""Outcome of an agent run."""
|
|
|
|
content: str = ""
|
|
messages: list[dict[str, Any]] = field(default_factory=list)
|
|
tool_calls: list[ToolCallRecord] = field(default_factory=list)
|
|
# Ordered Notion-style step descriptors ({key, params}); tool steps and
|
|
# intermediate reasoning notes interleaved by execution order.
|
|
steps: list[dict[str, Any]] = field(default_factory=list)
|
|
iterations: int = 0
|
|
stopped: str = STOP_DONE
|
|
pending: dict[str, Any] | None = None
|
|
|
|
|
|
async def provider_llm(messages: list[dict[str, Any]], tools: list[dict[str, Any]] | None):
|
|
"""Default LLM adapter backed by ``backend.ai_chat.chat_completion``."""
|
|
from backend.ai_chat import chat_completion
|
|
|
|
return await chat_completion(messages, tools=tools)
|
|
|
|
|
|
def _truncate(payload: Any) -> Any:
|
|
"""Truncate an oversized tool result before feeding it back to the model."""
|
|
serialized = json.dumps(payload, ensure_ascii=False, default=str)
|
|
if len(serialized) <= MAX_TOOL_RESULT_CHARS:
|
|
return payload
|
|
return {"truncated": True, "content": serialized[:MAX_TOOL_RESULT_CHARS]}
|
|
|
|
|
|
def _assistant_tool_message(content: str | None, tool_calls: list[Any]) -> dict[str, Any]:
|
|
"""Build the OpenAI-style assistant message carrying tool calls."""
|
|
return {
|
|
"role": "assistant",
|
|
"content": content or "",
|
|
"tool_calls": [
|
|
{
|
|
"id": call.id,
|
|
"type": "function",
|
|
"function": {
|
|
"name": call.name,
|
|
"arguments": json.dumps(call.arguments, ensure_ascii=False, default=str),
|
|
},
|
|
}
|
|
for call in tool_calls
|
|
],
|
|
}
|
|
|
|
|
|
def _deferred_tool_message(call: Any, reason: str | None = None) -> dict[str, Any]:
|
|
"""Answer a tool call that was not reached because the run stopped early.
|
|
|
|
A single LLM response may carry several tool calls; when the run stops
|
|
before reaching some of them (tool-call quota), the assistant message still
|
|
lists *all* of them, so every ``tool_call_id`` must get a tool result
|
|
before the next LLM call (the OpenAI tool protocol rejects dangling ids).
|
|
The calls that were not reached get a synthetic ``deferred`` result.
|
|
|
|
Note: mutating calls that pause the run for confirmation are no longer
|
|
deferred — they are batched and applied together on resume (BUG-075); this
|
|
helper remains for budget stops (BUG-050/BUG-052).
|
|
"""
|
|
return {
|
|
"role": "tool",
|
|
"tool_call_id": call.id,
|
|
"name": call.name,
|
|
"content": json.dumps({
|
|
"status": "deferred",
|
|
"reason": reason or (
|
|
"Not executed: the run stopped before reaching this tool call. "
|
|
"Re-issue this call if it is still needed."
|
|
),
|
|
}, ensure_ascii=False),
|
|
}
|
|
|
|
|
|
def _action_descriptor(call: Any) -> dict[str, Any]:
|
|
"""Describe one paused mutating tool call for the confirmation payload.
|
|
|
|
A single LLM response may request several mutations (create a folder and
|
|
the files inside it…). They are batched into one confirmation so the user
|
|
approves the whole plan in one click (BUG-075). ``step`` reuses the
|
|
Notion-style label, so the confirmation card reads like the steps block.
|
|
"""
|
|
return {
|
|
"id": call.id,
|
|
"tool": call.name,
|
|
"arguments": call.arguments,
|
|
"step": tool_step_label(call.name, call.arguments),
|
|
}
|
|
|
|
|
|
def _fallback_summary(executed: list[ToolCallRecord]) -> str:
|
|
"""Deterministic non-empty answer built from the gathered tool results.
|
|
|
|
Used only if the final synthesis call fails or returns nothing, so a turn
|
|
never ends on an empty message (BUG-052).
|
|
"""
|
|
lines: list[str] = []
|
|
for record in executed:
|
|
data = record.result
|
|
if not isinstance(data, dict):
|
|
continue
|
|
for item in (data.get("results") or [])[:5]:
|
|
if not isinstance(item, dict):
|
|
continue
|
|
title = item.get("title") or item.get("url") or ""
|
|
url = item.get("url") or ""
|
|
lines.append(f"- [{title}]({url})" if url else f"- {title}")
|
|
if data.get("url") and data.get("text"):
|
|
title = data.get("title") or data["url"]
|
|
lines.append(f"- [{title}]({data['url']})")
|
|
if not lines:
|
|
return "Je n'ai pas pu produire de réponse à partir des résultats obtenus."
|
|
unique = list(dict.fromkeys(lines))
|
|
return "Voici les sources pertinentes trouvées :\n" + "\n".join(unique)
|
|
|
|
|
|
async def _finalize_answer(
|
|
llm: Callable[..., Any],
|
|
convo: list[dict[str, Any]],
|
|
executed: list[ToolCallRecord],
|
|
steps: list[dict[str, Any]],
|
|
iterations: int,
|
|
stopped: str,
|
|
) -> AgentResult:
|
|
"""Guarantee a textual answer when the loop stopped before producing one.
|
|
|
|
Web research often exhausts the iteration budget while the model is still
|
|
calling tools; returning ``content=""`` left the conversation with steps and
|
|
sources but no answer. One final tool-less call asks the model to synthesize
|
|
the gathered results, and a deterministic source list is used as a last
|
|
resort (BUG-052).
|
|
"""
|
|
content = ""
|
|
if executed:
|
|
try:
|
|
response = await llm(
|
|
[*convo, {"role": "user", "content": _FINALIZE_INSTRUCTION}], []
|
|
)
|
|
content = (response.content or "").strip()
|
|
except Exception as e:
|
|
logger.warning(f"Agent final synthesis failed: {e}")
|
|
if not content:
|
|
content = _fallback_summary(executed)
|
|
return AgentResult(
|
|
content=content,
|
|
messages=convo,
|
|
tool_calls=executed,
|
|
steps=steps,
|
|
iterations=iterations,
|
|
stopped=stopped,
|
|
)
|
|
|
|
|
|
def _execute_confirmed(
|
|
ctx: ToolContext,
|
|
confirm_pending: dict[str, Any],
|
|
convo: list[dict[str, Any]],
|
|
executed: list[ToolCallRecord],
|
|
on_tool_call: Callable[[ToolCallRecord], None] | None,
|
|
) -> None:
|
|
"""Apply previously-paused mutating tool calls and feed their results back.
|
|
|
|
The pending payload is the ``error`` object emitted by a ``confirmation``
|
|
event, optionally carrying an ``actions`` list with every mutating call of
|
|
the LLM turn (BUG-075). Each action is applied with a one-shot confirmation
|
|
and its ``tool_call_id`` answered, keeping the conversation valid for the
|
|
resumed turn. The assistant tool-call message is expected to already be in
|
|
``convo`` (it is part of the snapshot returned with the confirmation).
|
|
"""
|
|
from backend.ai_chat import ToolCall
|
|
|
|
error = confirm_pending.get("error", confirm_pending) or {}
|
|
actions = confirm_pending.get("actions")
|
|
if not isinstance(actions, list) or not actions:
|
|
# Legacy single-action payload (no ``actions`` list).
|
|
actions = [{
|
|
"id": error.get("id") or "call_pending",
|
|
"tool": error.get("tool"),
|
|
"arguments": error.get("arguments") or {},
|
|
}]
|
|
|
|
for action in actions:
|
|
name = action.get("tool")
|
|
arguments = action.get("arguments") or {}
|
|
call_id = action.get("id") or "call_pending"
|
|
|
|
if not name:
|
|
raise ToolError("Malformed confirmation payload", code="invalid_confirmation")
|
|
|
|
# Make sure the assistant tool-call message is present in the snapshot.
|
|
if not any(
|
|
m.get("role") == "assistant" and any(
|
|
tc.get("id") == call_id for tc in (m.get("tool_calls") or [])
|
|
)
|
|
for m in convo
|
|
):
|
|
convo.append(_assistant_tool_message(None, [ToolCall(id=call_id, name=name, arguments=arguments)]))
|
|
|
|
try:
|
|
result = call_tool(name, ctx, arguments, confirm=True)
|
|
payload = result.data
|
|
ok = True
|
|
except ToolError as e:
|
|
payload = e.to_dict()
|
|
ok = False
|
|
|
|
record = ToolCallRecord(
|
|
name=name, arguments=arguments, ok=ok, result=payload,
|
|
step=tool_step_label(name, arguments),
|
|
)
|
|
executed.append(record)
|
|
if on_tool_call is not None:
|
|
on_tool_call(record)
|
|
|
|
convo.append({
|
|
"role": "tool",
|
|
"tool_call_id": call_id,
|
|
"name": name,
|
|
"content": json.dumps(_truncate(payload), ensure_ascii=False, default=str),
|
|
})
|
|
|
|
|
|
async def run_agent(
|
|
messages: list[dict[str, Any]],
|
|
*,
|
|
ctx: ToolContext,
|
|
llm: Callable[..., Any] | None = None,
|
|
tools: list[dict[str, Any]] | None = None,
|
|
max_iterations: int = DEFAULT_MAX_ITERATIONS,
|
|
max_tool_calls: int | None = None,
|
|
on_tool_call: Callable[[ToolCallRecord], None] | None = None,
|
|
on_thought: Callable[[dict[str, Any]], None] | None = None,
|
|
resume_messages: list[dict[str, Any]] | None = None,
|
|
confirm_pending: dict[str, Any] | None = None,
|
|
) -> AgentResult:
|
|
"""Run the tool-calling loop until completion.
|
|
|
|
Args:
|
|
messages: Initial conversation (OpenAI-style), typically a system
|
|
message followed by the conversation history and the user message.
|
|
ctx: Tool execution context (identity, mode, confirmation state).
|
|
llm: Async callable ``(messages, tools) -> LLMResponse``. Defaults to
|
|
the real provider adapter.
|
|
tools: Tool schemas to expose. ``None`` exposes all in-app tools;
|
|
pass ``[]`` to disable tool calling (plain chat).
|
|
max_iterations: Hard cap on LLM round-trips.
|
|
max_tool_calls: Hard cap on the total number of executed tool calls
|
|
(quota, defaults to ``BOOKSLM_MAX_TOOL_CALLS``).
|
|
on_tool_call: Optional callback invoked after each executed tool call.
|
|
resume_messages: Conversation snapshot from a paused run (returned with
|
|
a ``confirmation`` event). When set, the loop resumes from it.
|
|
confirm_pending: Pending mutating tool call to apply before resuming
|
|
(two-step propose/apply).
|
|
|
|
Returns:
|
|
An :class:`AgentResult`. ``stopped`` is ``done``, ``max_iterations`` or
|
|
``confirmation_required`` (in which case ``pending`` holds the payload
|
|
to confirm, for the two-step propose/apply flow).
|
|
"""
|
|
llm = llm or provider_llm
|
|
if tools is None:
|
|
tools = get_tool_schemas(scope=ToolScope.IN_APP)
|
|
quota = DEFAULT_MAX_TOOL_CALLS if max_tool_calls is None else max_tool_calls
|
|
steps: list[dict[str, Any]] = []
|
|
|
|
def _emit_note(text: str) -> None:
|
|
"""Record an intermediate reasoning note as a visible step."""
|
|
note = thought_step_label(text)
|
|
if note["params"]["value"]:
|
|
steps.append(note)
|
|
if on_thought is not None:
|
|
on_thought(note)
|
|
|
|
convo = [dict(m) for m in (resume_messages if resume_messages is not None else messages)]
|
|
executed: list[ToolCallRecord] = []
|
|
|
|
def _run_call(call: Any) -> None:
|
|
"""Execute one tool call, record it and answer its ``tool_call_id``.
|
|
|
|
``ToolConfirmationRequired`` propagates to the caller so the loop can
|
|
pause and batch the mutating calls of the turn (BUG-075).
|
|
"""
|
|
try:
|
|
result = call_tool(call.name, ctx, call.arguments)
|
|
payload: Any = result.data
|
|
ok = True
|
|
except ToolConfirmationRequired:
|
|
raise
|
|
except ToolError as e:
|
|
payload = e.to_dict()
|
|
ok = False
|
|
record = ToolCallRecord(
|
|
name=call.name, arguments=call.arguments, ok=ok, result=payload,
|
|
step=tool_step_label(call.name, call.arguments),
|
|
)
|
|
executed.append(record)
|
|
steps.append(record.step)
|
|
if on_tool_call is not None:
|
|
on_tool_call(record)
|
|
convo.append({
|
|
"role": "tool",
|
|
"tool_call_id": call.id,
|
|
"name": call.name,
|
|
"content": json.dumps(_truncate(payload), ensure_ascii=False, default=str),
|
|
})
|
|
|
|
if confirm_pending:
|
|
if quota is not None and len(executed) >= quota:
|
|
return AgentResult(
|
|
content="",
|
|
messages=convo,
|
|
tool_calls=executed,
|
|
steps=steps,
|
|
iterations=0,
|
|
stopped=STOP_QUOTA_EXCEEDED,
|
|
)
|
|
_execute_confirmed(ctx, confirm_pending, convo, executed, on_tool_call)
|
|
|
|
for iteration in range(1, max_iterations + 1):
|
|
response = await llm(convo, tools)
|
|
|
|
if not response.has_tool_calls:
|
|
return AgentResult(
|
|
content=response.content or "",
|
|
messages=convo,
|
|
tool_calls=executed,
|
|
steps=steps,
|
|
iterations=iteration,
|
|
stopped=STOP_DONE,
|
|
)
|
|
|
|
# Intermediate reasoning shown alongside tool calls → a "thought" step.
|
|
_emit_note(response.content or "")
|
|
convo.append(_assistant_tool_message(response.content, response.tool_calls))
|
|
|
|
for index, call in enumerate(response.tool_calls):
|
|
if quota is not None and len(executed) >= quota:
|
|
logger.warning(f"Agent reached the tool-call quota ({quota})")
|
|
# Keep the conversation valid for the synthesis call: the
|
|
# assistant message announced every tool call of the batch.
|
|
for skipped in response.tool_calls[index:]:
|
|
convo.append(_deferred_tool_message(
|
|
skipped, "Not executed: the tool-call quota was reached."
|
|
))
|
|
return await _finalize_answer(
|
|
llm, convo, executed, steps, iteration, STOP_QUOTA_EXCEEDED
|
|
)
|
|
try:
|
|
_run_call(call)
|
|
except ToolConfirmationRequired as e:
|
|
logger.info(f"Agent paused: confirmation required for '{call.name}'")
|
|
pending = e.to_dict()
|
|
# Include the tool-call id so the client can echo it back.
|
|
pending["error"]["id"] = call.id
|
|
# BUG-075: batch every mutating call of this LLM turn so the
|
|
# user approves the whole plan at once (one resume applies them
|
|
# all) instead of approving one action after another. Read-only
|
|
# calls of the batch run immediately and answer their
|
|
# ``tool_call_id`` so the resumed turn stays valid.
|
|
actions = [_action_descriptor(call)]
|
|
for after in response.tool_calls[index + 1:]:
|
|
spec = get_tool(after.name)
|
|
if spec is not None and spec.requires_confirmation:
|
|
actions.append(_action_descriptor(after))
|
|
else:
|
|
_run_call(after)
|
|
pending["actions"] = actions
|
|
return AgentResult(
|
|
content=response.content or "",
|
|
messages=convo,
|
|
tool_calls=executed,
|
|
steps=steps,
|
|
iterations=iteration,
|
|
stopped=STOP_CONFIRMATION_REQUIRED,
|
|
pending=pending,
|
|
)
|
|
|
|
logger.warning(f"Agent reached max iterations ({max_iterations})")
|
|
return await _finalize_answer(
|
|
llm, convo, executed, steps, max_iterations, STOP_MAX_ITERATIONS
|
|
)
|