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
253 lines
8.4 KiB
Python
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,
|
|
)
|