Files
ObsiGate/backend/agent/loop.py
T
bruno 4c4e415975
CI / lint (push) Successful in 58s
CI / security (push) Successful in 40s
CI / test (push) Successful in 1m15s
CI / build (push) Successful in 37s
CI / e2e (push) Successful in 10m15s
Desktop Build / build-windows (push) Canceled after 0s
Desktop Build / build-linux (push) Canceled after 0s
feat(ai): services partages A2 + SSE streaming B4 + confirmations UI B5 (#79)
2026-09-11 17:06:40 -04:00

253 lines
8.4 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
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_schemas,
)
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
# Stopping reasons
STOP_DONE = "done"
STOP_MAX_ITERATIONS = "max_iterations"
STOP_CONFIRMATION_REQUIRED = "confirmation_required"
@dataclass
class ToolCallRecord:
"""Audit-friendly record of one executed tool call."""
name: str
arguments: dict[str, Any]
ok: bool
result: Any
@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)
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 _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 a previously-paused mutating tool call and feed its result back.
The pending payload is the ``error`` object emitted by a ``confirmation``
event. 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)
name = error.get("tool")
arguments = error.get("arguments") or {}
call_id = error.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)
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,
on_tool_call: Callable[[ToolCallRecord], 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.
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)
convo = [dict(m) for m in (resume_messages if resume_messages is not None else messages)]
executed: list[ToolCallRecord] = []
if confirm_pending:
_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,
iterations=iteration,
stopped=STOP_DONE,
)
convo.append(_assistant_tool_message(response.content, response.tool_calls))
for call in response.tool_calls:
try:
result = call_tool(call.name, ctx, call.arguments)
payload = result.data
ok = True
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
return AgentResult(
content=response.content or "",
messages=convo,
tool_calls=executed,
iterations=iteration,
stopped=STOP_CONFIRMATION_REQUIRED,
pending=pending,
)
except ToolError as e:
payload = e.to_dict()
ok = False
record = ToolCallRecord(name=call.name, arguments=call.arguments, ok=ok, result=payload)
executed.append(record)
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),
})
logger.warning(f"Agent reached max iterations ({max_iterations})")
return AgentResult(
content="",
messages=convo,
tool_calls=executed,
iterations=max_iterations,
stopped=STOP_MAX_ITERATIONS,
)