oclaw/runtime/chat/agent.py
oliver d245e463c0 Make tool failures actionable with suggestions and scope hints.
Unregistered tools now suggest close matches; MCP insufficient_scope is rewritten with fallbacks so ops agents stop blind-retrying SQL without sql:query.

Co-authored-by: Cursor <cursoragent@cursor.com>
2026-08-10 22:42:37 +08:00

282 lines
11 KiB
Python

from __future__ import annotations
import json
import logging
import os
import time
from concurrent.futures import ThreadPoolExecutor, as_completed
from collections.abc import Callable
from dataclasses import dataclass
from typing import Any, Optional
from runtime.tools.base import ToolRegistry
from svc.persistence.sqlite_store import SqliteStore
from svc.llm.chat_models import (
ChatModel,
LLMResponse,
LLMToolCall,
OpenAIChatModel,
RuleBasedChatModel,
StaticTextChatModel,
_normalize_image_b64_payload,
build_default_model,
gemini_openai_compat_client,
)
from runtime.prompt_templates import render_prompt_for_lang
from runtime.tools.tool_validation import format_invalid_arguments_error, validate_tool_arguments
logger = logging.getLogger(__name__)
SESSION_TITLE_MAX_LEN = 120
AGENT_CONTEXT_MESSAGES = 80
DEFAULT_SYSTEM_PROMPTS: dict[str, str] = {
"zh": render_prompt_for_lang("runtime/default_system", "zh", strict=True),
"en": render_prompt_for_lang("runtime/default_system", "en", strict=True),
}
class GenerationInterrupted(Exception):
"""用户请求中止当前生成过程。"""
@dataclass(frozen=True)
class AgentConfig:
max_messages: int = AGENT_CONTEXT_MESSAGES
max_tool_rounds: int = 100
max_tool_workers: int = 8
class Agent:
def __init__(
self,
store: SqliteStore,
tools: ToolRegistry,
model: Optional[ChatModel] = None,
config: Optional[AgentConfig] = None,
system_prompt: str | None = None,
lang: str = "zh",
llm_profile_mode: str | None = None,
):
self.store = store
self.tools = tools
self.model = model or build_default_model()
self.config = config or AgentConfig()
self.lang = (lang or "zh").strip().lower()
self._system_prompt_base = (system_prompt or DEFAULT_SYSTEM_PROMPTS.get(self.lang, DEFAULT_SYSTEM_PROMPTS["zh"])).strip()
self.llm_profile_mode = ((llm_profile_mode or "").strip().lower() or None)
self._last_turn_outcome: Any | None = None
def _native_tools_sent_by_api(self) -> bool:
"""当前模型这一侧是否会把 tools 放进请求(与 ``llm.OpenAIChatModel._skip_tools`` 对齐)。"""
m = self.model
if isinstance(m, (RuleBasedChatModel, StaticTextChatModel)):
return False
if isinstance(m, OpenAIChatModel):
return not bool(m._skip_tools)
return False
def _compose_system_prompt(self) -> str:
"""系统正文。工具 schema 始终通过原生 tools 字段下发,不再拼接到 prompt。"""
return self._system_prompt_base
def _format_ollama_failure_banner(self, exc: BaseException) -> str:
# Backward compat wrapper; implementation lives in `src.chat.agent_errors`.
from runtime.chat.agent_errors import format_ollama_failure_banner
return format_ollama_failure_banner(lang=self.lang, exc=exc)
def _format_openai_transport_error(self, exc: BaseException) -> str:
# Backward compat wrapper; implementation lives in `src.chat.agent_errors`.
from runtime.chat.agent_errors import format_openai_transport_error
return format_openai_transport_error(lang=self.lang, exc=exc)
def _invoke_tool(self, tc: LLMToolCall) -> tuple[dict[str, Any], int]:
t0 = time.perf_counter()
tool = self.tools.get(tc.name)
if not tool:
from runtime.tools.tool_error_hints import format_unregistered_tool_error
available = [t.name for t in self.tools.list()] if hasattr(self.tools, "list") else []
return format_unregistered_tool_error(
str(tc.name or ""),
available,
lang=str(self.lang or "zh"),
), int((time.perf_counter() - t0) * 1000)
ok, v_err = validate_tool_arguments(tool.parameters, tc.arguments)
if not ok:
return format_invalid_arguments_error(
tool.parameters or {},
str(v_err or "invalid"),
lang=str(self.lang or "zh"),
), int((time.perf_counter() - t0) * 1000)
try:
result = tool.handler(tc.arguments)
return result, int((time.perf_counter() - t0) * 1000)
except Exception as e:
if self.lang.startswith("en"):
err = {"ok": False, "error": f"Tool execution error: {type(e).__name__}: {e}"}
else:
err = {"ok": False, "error": f"工具执行异常: {type(e).__name__}: {e}"}
return err, int((time.perf_counter() - t0) * 1000)
def _emit_progress(self, on_progress: Optional[Callable[[str], None]], en: str, zh: str) -> None:
if on_progress:
on_progress(en if self.lang.startswith("en") else zh)
@staticmethod
def _attachments_from_tool_result(result: Any) -> list[dict[str, Any]]:
"""Extract image/relay references from tool results for rendering."""
if not isinstance(result, dict):
return []
out: list[dict[str, Any]] = []
aid = str(result.get("attachment_id") or "").strip()
if aid:
out.append(
{
"type": "image_ref",
"attachment_id": aid,
"name": str(result.get("name") or "generated-image"),
"mime": str(result.get("mime") or "image/png"),
"bytes": result.get("bytes"),
"width": result.get("width"),
"height": result.get("height"),
}
)
refs = result.get("attachments")
if isinstance(refs, list):
for r in refs:
if not isinstance(r, dict):
continue
# Relay pointer payload (new protocol).
p_uri = str(r.get("pointer_uri") or "").strip()
if p_uri:
out.append(
{
"type": "relay_pointer",
"pointer_uri": p_uri,
"rel_path": str(r.get("rel_path") or ""),
"mime": str(r.get("mime_type") or r.get("mime") or ""),
"bytes": r.get("bytes"),
"sha256": str(r.get("sha256") or ""),
"name": str(r.get("name") or ""),
}
)
continue
r_aid = str(r.get("attachment_id") or "").strip()
if not r_aid:
continue
out.append(
{
"type": "image_ref",
"attachment_id": r_aid,
"name": str(r.get("name") or "generated-image"),
"mime": str(r.get("mime") or "image/png"),
"bytes": r.get("bytes"),
"width": r.get("width"),
"height": r.get("height"),
}
)
# de-dup by attachment_id
uniq: list[dict[str, Any]] = []
seen: set[str] = set()
for a in out:
k = str(a.get("attachment_id") or a.get("pointer_uri") or "")
if not k or k in seen:
continue
seen.add(k)
uniq.append(a)
return uniq
def run_turn(
self,
session_id: str,
user_text: str,
attachments: list[dict[str, Any]] | None = None,
on_progress: Optional[Callable[[str], None]] = None,
on_token: Optional[Callable[[str], None]] = None,
on_tool_ui: Optional[Callable[[str, dict[str, Any]], None]] = None,
should_stop: Optional[Callable[[], bool]] = None,
*,
workspace_owner_session_id: str | None = None,
path_policy_tenant_id: str | None = None,
path_policy_user_id: str | None = None,
interaction_mode: str | None = None,
selected_specialist: str | None = None,
) -> str:
from runtime.gateway import OclawGateway
from runtime.types import StandardMessage
tenant_id = str(path_policy_tenant_id or "").strip()
user_id = str(path_policy_user_id or "").strip()
if not tenant_id or not user_id:
try:
owner = self.store.get_ui_session_owner(session_id=session_id)
except Exception:
owner = None
if isinstance(owner, dict):
tenant_id = tenant_id or str(owner.get("tenant_id") or "")
user_id = user_id or str(owner.get("user_id") or "")
self._emit_progress(
on_progress,
"Received. Working on your request…",
"已收到,正在处理…",
)
meta: dict[str, Any] = {"tenant_id": tenant_id, "user_id": user_id}
if workspace_owner_session_id:
meta["workspace_owner_session_id"] = str(workspace_owner_session_id).strip()
if str(interaction_mode or "").strip():
meta["interaction_mode"] = str(interaction_mode).strip().lower()
if str(selected_specialist or "").strip():
meta["selected_specialist"] = str(selected_specialist).strip().lower()
msg = StandardMessage(
session_id=session_id,
tenant_id=tenant_id,
user_id=user_id,
role="member",
channel="agent_turn",
text=str(user_text or ""),
attachments=list(attachments or []),
metadata=meta,
)
gw = OclawGateway(store=self.store)
try:
res = gw.handle_turn(
msg=msg,
lang=self.lang,
executor=self,
on_token=on_token,
on_progress=on_progress,
on_tool_ui=on_tool_ui,
should_stop=should_stop,
)
except RuntimeError as e:
low = str(e).lower()
if "interrupted" in low and "user" in low:
raise GenerationInterrupted(str(e)) from e
raise
self._last_turn_outcome = getattr(self, "_last_turn_outcome", None)
return str(res.reply_text or "")
def _build_llm_messages(self, session_id: str) -> list[dict[str, Any]]:
from runtime.chat.agent_messages import build_llm_messages
msgs = self.store.get_messages(session_id=session_id, limit=self.config.max_messages)
trunc_raw = str(self.store.get_setting("AIA_TOOL_CONTEXT_TRUNCATE_ENABLED") or "").strip().lower()
tool_context_truncate_enabled = trunc_raw not in ("0", "false", "no", "off")
return build_llm_messages(
store_messages=msgs,
system_prompt=self._compose_system_prompt(),
model=self.model,
lang=self.lang,
tool_context_truncate_enabled=tool_context_truncate_enabled,
)
__all__ = ["AgentConfig", "DEFAULT_SYSTEM_PROMPTS", "GenerationInterrupted", "Agent"]