mirror of
https://github.com/hansjone/oclaw.git
synced 2026-10-09 05:50:44 +08:00
重构仓库目录为统一的 runtime 分层并清理历史 openclaw 残留。
本次迁移将网关/通道/工具/技能/脚本与协议资源集中到新结构,统一路径常量与脚本转发机制,减少顶层噪音并保证运行与测试行为一致。 Made-with: Cursor
This commit is contained in:
parent
ba3836f00f
commit
4a23b715a2
498 changed files with 2760 additions and 2200 deletions
1
runtime/chat/__init__.py
Normal file
1
runtime/chat/__init__.py
Normal file
|
|
@ -0,0 +1 @@
|
|||
# oclaw.chat package
|
||||
278
runtime/chat/agent.py
Normal file
278
runtime/chat/agent.py
Normal file
|
|
@ -0,0 +1,278 @@
|
|||
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 oclaw.runtime.tools.base import ToolRegistry
|
||||
from oclaw.platform.persistence.sqlite_store import SqliteStore
|
||||
from oclaw.platform.llm.chat_models import (
|
||||
ChatModel,
|
||||
LLMResponse,
|
||||
LLMToolCall,
|
||||
OpenAIChatModel,
|
||||
RuleBasedChatModel,
|
||||
StaticTextChatModel,
|
||||
_normalize_image_b64_payload,
|
||||
build_default_model,
|
||||
gemini_openai_compat_client,
|
||||
)
|
||||
from oclaw.prompts.loader import render_prompt_for_lang
|
||||
from oclaw.runtime.tools.tool_validation import 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 = 8
|
||||
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 oclaw.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 oclaw.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:
|
||||
msg = f"Unregistered tool: {tc.name}" if self.lang.startswith("en") else f"未注册的工具: {tc.name}"
|
||||
return {"ok": False, "error": msg}, int((time.perf_counter() - t0) * 1000)
|
||||
|
||||
ok, v_err = validate_tool_arguments(tool.parameters, tc.arguments)
|
||||
if not ok:
|
||||
msg = f"Invalid arguments: {v_err}" if self.lang.startswith("en") else f"参数不合法: {v_err}"
|
||||
return {"ok": False, "error": msg}, 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 oclaw.runtime.gateway import OclawGateway
|
||||
from oclaw.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 "")
|
||||
|
||||
session = self.store.get_session(session_id)
|
||||
if session and session.title in ("新会话", "New Chat"):
|
||||
title = user_text.strip().replace("\n", " ")
|
||||
if not title and attachments:
|
||||
title = str(attachments[0].get("name") or "New Chat")
|
||||
if title:
|
||||
self.store.rename_session(session_id, title[:SESSION_TITLE_MAX_LEN])
|
||||
|
||||
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 oclaw.runtime.chat.agent_messages import build_llm_messages
|
||||
|
||||
msgs = self.store.get_messages(session_id=session_id, limit=self.config.max_messages)
|
||||
return build_llm_messages(
|
||||
store_messages=msgs,
|
||||
system_prompt=self._compose_system_prompt(),
|
||||
model=self.model,
|
||||
lang=self.lang,
|
||||
)
|
||||
|
||||
|
||||
__all__ = ["AgentConfig", "DEFAULT_SYSTEM_PROMPTS", "GenerationInterrupted", "Agent"]
|
||||
69
runtime/chat/agent_errors.py
Normal file
69
runtime/chat/agent_errors.py
Normal file
|
|
@ -0,0 +1,69 @@
|
|||
from __future__ import annotations
|
||||
|
||||
"""Agent 错误处理模块。
|
||||
|
||||
把 `Agent` 内的错误格式化逻辑下沉到此处,方便 manager/specialist 复用。
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
from oclaw.prompts import render_prompt
|
||||
|
||||
|
||||
def format_ollama_failure_banner(*, lang: str, exc: BaseException) -> str:
|
||||
prompt_id = "fallback/ollama_failure.en.md" if (lang or "zh").startswith("en") else "fallback/ollama_failure.zh.md"
|
||||
return render_prompt(
|
||||
prompt_id,
|
||||
variables={"error_type": type(exc).__name__, "error_message": str(exc)},
|
||||
strict=True,
|
||||
)
|
||||
|
||||
|
||||
def format_openai_transport_error(*, lang: str, exc: BaseException) -> str:
|
||||
blob = str(exc).lower()
|
||||
oversized_tool = ("30000" in blob or "input length" in blob) and ("range" in blob or "length" in blob)
|
||||
gemini_sig = "thought_signature" in blob
|
||||
if (lang or "zh").startswith("en"):
|
||||
if oversized_tool:
|
||||
return render_prompt(
|
||||
"fallback/openai_transport_oversized.en.md",
|
||||
variables={"error_type": type(exc).__name__, "error_message": str(exc)},
|
||||
strict=True,
|
||||
)
|
||||
tail = (
|
||||
"\n\n_Gemini 3 with tools: the API requires echoing `thought_signature` from each tool use in chat history. "
|
||||
"If this persists, update the app or use a model/SDK path that preserves provider-specific tool fields._"
|
||||
if gemini_sig
|
||||
else ""
|
||||
)
|
||||
return render_prompt(
|
||||
"fallback/openai_transport_error.en.md",
|
||||
variables={"error_type": type(exc).__name__, "error_message": str(exc), "extra_tail": tail},
|
||||
strict=True,
|
||||
)
|
||||
if oversized_tool:
|
||||
return render_prompt(
|
||||
"fallback/openai_transport_oversized.zh.md",
|
||||
variables={"error_type": type(exc).__name__, "error_message": str(exc)},
|
||||
strict=True,
|
||||
)
|
||||
tail = (
|
||||
"\n\n(**Gemini 3 + 工具调用**:接口要求把模型返回的 **thought_signature** 随该次 `tool_calls` 一并写回对话历史;"
|
||||
"首轮能跑工具、第二轮 400 多为丢失该字段。若已更新本应用仍报错,请确认代理/OpenAI 兼容层是否透传该字段。)"
|
||||
if gemini_sig
|
||||
else ""
|
||||
)
|
||||
return render_prompt(
|
||||
"fallback/openai_transport_error.zh.md",
|
||||
variables={"error_type": type(exc).__name__, "error_message": str(exc), "extra_tail": tail},
|
||||
strict=True,
|
||||
)
|
||||
|
||||
|
||||
def safe_str(e: Any) -> str:
|
||||
try:
|
||||
return str(e)
|
||||
except Exception:
|
||||
return repr(e)
|
||||
|
||||
|
||||
__all__ = ["format_ollama_failure_banner", "format_openai_transport_error", "safe_str"]
|
||||
494
runtime/chat/agent_messages.py
Normal file
494
runtime/chat/agent_messages.py
Normal file
|
|
@ -0,0 +1,494 @@
|
|||
from __future__ import annotations
|
||||
|
||||
"""Agent 消息构建模块。
|
||||
|
||||
把 `Agent._build_llm_messages` 的职责下沉到此处,便于:
|
||||
- Manager 决策/Final merge 复用同一套“消息规范化与附件注入”规则
|
||||
- 后续 Workspace/RAG/Trace 插入上下文时有单一入口
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
from typing import Any
|
||||
|
||||
from oclaw.platform.llm.chat_models import _normalize_image_b64_payload, gemini_openai_compat_client, ChatModel
|
||||
from oclaw.runtime.chat.tool_runtime import tool_llm_message_max_chars, truncate_tool_result_for_llm_messages
|
||||
from oclaw.prompts import render_prompt
|
||||
from oclaw.platform.files.attachment_assets import attachment_id_to_data_url
|
||||
from oclaw.runtime.relay_pointer import parse_pointer_uri
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
_THINK_BLOCK_RE = re.compile(r"<think>\s*(.*?)\s*</think>\s*", flags=re.IGNORECASE | re.DOTALL)
|
||||
|
||||
|
||||
def _replay_recent_tool_rounds() -> int:
|
||||
raw = str(os.getenv("AIA_REPLAY_TOOL_FULL_ROUNDS") or "").strip()
|
||||
if raw.isdigit():
|
||||
return max(0, min(int(raw), 12))
|
||||
return 3
|
||||
|
||||
|
||||
def _allow_reasoning_signature_replay(model: ChatModel) -> bool:
|
||||
# - auto (default): only providers that require signature continuity (Gemini paths).
|
||||
# - on: always include signature metadata on assistant tool_calls.
|
||||
# - off: never include signature metadata.
|
||||
policy = str(os.getenv("AIA_REPLAY_REASONING_SIGNATURE_POLICY") or "auto").strip().lower()
|
||||
if policy in ("0", "off", "false", "no"):
|
||||
return False
|
||||
if policy in ("1", "on", "true", "yes"):
|
||||
return True
|
||||
if gemini_openai_compat_client(model):
|
||||
return True
|
||||
return model.__class__.__name__ == "GoogleGeminiChatModel"
|
||||
|
||||
|
||||
def _strip_reasoning_blocks(text: str) -> str:
|
||||
return _THINK_BLOCK_RE.sub("", str(text or "")).strip()
|
||||
|
||||
|
||||
def _parse_tool_calls(raw_tc: Any) -> list[dict[str, Any]]:
|
||||
if not raw_tc:
|
||||
return []
|
||||
try:
|
||||
data = json.loads(raw_tc) if isinstance(raw_tc, str) else raw_tc
|
||||
except Exception:
|
||||
return []
|
||||
if not isinstance(data, list):
|
||||
return []
|
||||
return [x for x in data if isinstance(x, dict)]
|
||||
|
||||
|
||||
def _tool_call_id_from_tool_row(raw_tc: Any) -> str:
|
||||
if not raw_tc:
|
||||
return ""
|
||||
try:
|
||||
meta = json.loads(raw_tc) if isinstance(raw_tc, str) else raw_tc
|
||||
except Exception:
|
||||
return ""
|
||||
if not isinstance(meta, dict):
|
||||
return ""
|
||||
return str(meta.get("tool_call_id") or "").strip()
|
||||
|
||||
|
||||
def _collect_historical_tool_call_ids(store_messages: list[Any], *, full_rounds: int) -> set[str]:
|
||||
if full_rounds < 0:
|
||||
full_rounds = 0
|
||||
full_ids: set[str] = set()
|
||||
rounds = 0
|
||||
for m in reversed(store_messages or []):
|
||||
role = str(getattr(m, "role", "") or "")
|
||||
if role != "assistant":
|
||||
continue
|
||||
tcs = _parse_tool_calls(getattr(m, "tool_calls", None))
|
||||
tc_ids = [str(tc.get("id") or "").strip() for tc in tcs if str(tc.get("id") or "").strip()]
|
||||
if not tc_ids:
|
||||
continue
|
||||
rounds += 1
|
||||
if rounds <= full_rounds:
|
||||
full_ids.update(tc_ids)
|
||||
historical_ids: set[str] = set()
|
||||
for m in store_messages or []:
|
||||
if str(getattr(m, "role", "") or "") != "tool":
|
||||
continue
|
||||
tcid = _tool_call_id_from_tool_row(getattr(m, "tool_calls", None))
|
||||
if tcid and tcid not in full_ids:
|
||||
historical_ids.add(tcid)
|
||||
return historical_ids
|
||||
|
||||
|
||||
def _summarize_historical_tool_content(raw: str, *, cap: int) -> str:
|
||||
s = str(raw or "").strip()
|
||||
if not s:
|
||||
return json.dumps({"ok": None, "summary": "", "_history_summarized": True}, ensure_ascii=False)
|
||||
out: dict[str, Any] = {"_history_summarized": True}
|
||||
try:
|
||||
obj = json.loads(s)
|
||||
except Exception:
|
||||
preview = s[: max(1, cap - 120)] + ("\n...<truncated>" if len(s) > cap else "")
|
||||
out["summary"] = preview
|
||||
return json.dumps(out, ensure_ascii=False)
|
||||
if not isinstance(obj, dict):
|
||||
out["summary"] = s[: max(1, cap - 120)] + ("\n...<truncated>" if len(s) > cap else "")
|
||||
return json.dumps(out, ensure_ascii=False)
|
||||
out["ok"] = obj.get("ok")
|
||||
for key in ("error_code", "error", "hint"):
|
||||
v = str(obj.get(key) or "").strip()
|
||||
if v:
|
||||
out[key] = v
|
||||
if "result" in obj:
|
||||
r = obj.get("result")
|
||||
if isinstance(r, dict):
|
||||
out["result_keys"] = sorted(list(r.keys()))[:20]
|
||||
preview = s[: max(1, cap - 260)] + ("\n...<truncated>" if len(s) > cap else "")
|
||||
out["preview"] = preview
|
||||
return json.dumps(out, ensure_ascii=False)
|
||||
|
||||
|
||||
def _summarize_unpaired_tool_content(raw: str, *, cap: int) -> str:
|
||||
"""Best-effort summarize tool JSON for model-friendly context."""
|
||||
s = str(raw or "").strip()
|
||||
if not s:
|
||||
return ""
|
||||
if cap > 0 and len(s) > cap:
|
||||
s = s[: max(1, cap - 80)] + "\n...<truncated>"
|
||||
try:
|
||||
obj = json.loads(s)
|
||||
except Exception:
|
||||
return s
|
||||
if not isinstance(obj, dict):
|
||||
return s
|
||||
lines: list[str] = []
|
||||
ok = obj.get("ok")
|
||||
if ok is not None:
|
||||
lines.append(f"ok={bool(ok)}")
|
||||
ec = str(obj.get("error_code") or "").strip()
|
||||
if ec:
|
||||
lines.append(f"error_code={ec}")
|
||||
err = str(obj.get("error") or "").strip()
|
||||
if err:
|
||||
lines.append(f"error={err}")
|
||||
hint = str(obj.get("hint") or "").strip()
|
||||
if hint:
|
||||
lines.append(f"hint={hint}")
|
||||
# Extract MCP-style text blocks when present.
|
||||
try:
|
||||
nested = obj.get("result")
|
||||
content = None
|
||||
if isinstance(nested, dict):
|
||||
content = nested.get("content")
|
||||
if isinstance(content, list):
|
||||
texts = []
|
||||
for b in content:
|
||||
if isinstance(b, dict) and str(b.get("type") or "").strip().lower() == "text":
|
||||
t = str(b.get("text") or "").strip()
|
||||
if t:
|
||||
texts.append(t)
|
||||
if texts:
|
||||
lines.append("content_text=" + " | ".join(texts)[: min(800, cap)])
|
||||
except Exception:
|
||||
pass
|
||||
head = " ".join(lines).strip()
|
||||
if head:
|
||||
return head + "\n" + s
|
||||
return s
|
||||
|
||||
|
||||
def build_llm_messages(
|
||||
*,
|
||||
store_messages: list[Any],
|
||||
system_prompt: str,
|
||||
model: ChatModel,
|
||||
lang: str,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""把 DB 中的消息序列转换为 LLM messages。"""
|
||||
out: list[dict[str, Any]] = [{"role": "system", "content": (system_prompt or "").strip()}]
|
||||
allow_signature_replay = _allow_reasoning_signature_replay(model)
|
||||
historical_tool_ids = _collect_historical_tool_call_ids(
|
||||
store_messages=store_messages, full_rounds=_replay_recent_tool_rounds()
|
||||
)
|
||||
# Some OpenAI-compatible gateways error if a tool message references a tool_call_id
|
||||
# that is not present in the assistant tool_calls within the same request context.
|
||||
# This can happen when context windows are trimmed and the assistant tool_calls row is dropped.
|
||||
valid_tool_call_ids: set[str] = set()
|
||||
for m in store_messages:
|
||||
role = str(getattr(m, "role", "") or "")
|
||||
event_type = str(getattr(m, "event_type", "") or "").strip().lower()
|
||||
if event_type == "reasoning":
|
||||
continue
|
||||
|
||||
if role == "user":
|
||||
content_list: list[dict[str, Any]] = []
|
||||
text = getattr(m, "content", None)
|
||||
if text:
|
||||
content_list.append({"type": "text", "text": str(text)})
|
||||
|
||||
attachments = []
|
||||
raw_att = getattr(m, "attachments", None)
|
||||
if raw_att:
|
||||
try:
|
||||
attachments = json.loads(raw_att) if isinstance(raw_att, str) else raw_att
|
||||
except Exception:
|
||||
attachments = []
|
||||
|
||||
for att in attachments or []:
|
||||
if not isinstance(att, dict):
|
||||
continue
|
||||
att_type = att.get("type")
|
||||
if att_type in ("image", "input_image"):
|
||||
b64 = _normalize_image_b64_payload(att.get("image_base64") or att.get("data"))
|
||||
if not b64:
|
||||
continue
|
||||
content_list.append(
|
||||
{
|
||||
"type": "input_image",
|
||||
"image_base64": b64,
|
||||
"mime": att.get("mime") or "image/jpeg",
|
||||
}
|
||||
)
|
||||
elif att_type == "image_ref":
|
||||
# Prefer actual image bytes so multi-agent/image specialist can truly "see" history images.
|
||||
name = str(att.get("name") or "image")
|
||||
mime = str(att.get("mime") or "image/jpeg")
|
||||
aid = str(att.get("attachment_id") or "")
|
||||
data_url = attachment_id_to_data_url(aid, mime=mime) if aid else ""
|
||||
if data_url:
|
||||
if ";base64," in data_url:
|
||||
b64 = data_url.split(";base64,", 1)[1]
|
||||
content_list.append(
|
||||
{
|
||||
"type": "input_image",
|
||||
"image_base64": b64,
|
||||
"mime": mime,
|
||||
}
|
||||
)
|
||||
continue
|
||||
w = att.get("width")
|
||||
h = att.get("height")
|
||||
sz = att.get("bytes")
|
||||
meta_line = f"- name={name} mime={mime} id={aid}"
|
||||
if w and h:
|
||||
meta_line += f" size={w}x{h}"
|
||||
if sz:
|
||||
meta_line += f" bytes={sz}"
|
||||
content_list.append(
|
||||
{
|
||||
"type": "text",
|
||||
"text": render_prompt(
|
||||
"tools/image_attachment_meta.md",
|
||||
variables={"meta_line": meta_line},
|
||||
strict=True,
|
||||
),
|
||||
}
|
||||
)
|
||||
elif att_type == "text":
|
||||
name = att.get("name", "file")
|
||||
text_content = att.get("content", "")
|
||||
content_list.append(
|
||||
{
|
||||
"type": "text",
|
||||
"text": render_prompt(
|
||||
"tools/text_attachment_wrap.md",
|
||||
variables={"name": str(name), "content": str(text_content)},
|
||||
strict=True,
|
||||
),
|
||||
}
|
||||
)
|
||||
elif att_type == "tabular_ref":
|
||||
name = str(att.get("name") or "table")
|
||||
table_id = str(att.get("table_id") or "")
|
||||
rows = int(att.get("rows") or 0)
|
||||
cols = int(att.get("cols") or 0)
|
||||
aid = str(att.get("attachment_id") or "")
|
||||
sheets = att.get("sheets") if isinstance(att.get("sheets"), list) else []
|
||||
sheet_hint = ""
|
||||
if sheets:
|
||||
names = [str((x or {}).get("sheet_name") or "") for x in sheets if isinstance(x, dict)]
|
||||
names = [x for x in names if x]
|
||||
if names:
|
||||
sheet_hint = f"\n- sheets: {', '.join(names[:8])}"
|
||||
content_list.append(
|
||||
{
|
||||
"type": "text",
|
||||
"text": (
|
||||
f"[LargeTableAttachment]\n"
|
||||
f"- name: {name}\n"
|
||||
f"- table_id: {table_id}\n"
|
||||
f"- attachment_id: {aid}\n"
|
||||
f"- rows: {rows}\n"
|
||||
f"- cols: {cols}\n"
|
||||
f"{sheet_hint}\n"
|
||||
f"- tools: query_tabular_attachment, run_tabular_sql, analyze_tabular_attachment_full_scan"
|
||||
),
|
||||
}
|
||||
)
|
||||
elif att_type == "relay_pointer":
|
||||
p_uri = str(att.get("pointer_uri") or "").strip()
|
||||
if not p_uri:
|
||||
continue
|
||||
mime = str(att.get("mime") or att.get("mime_type") or "").strip()
|
||||
aid = str(att.get("attachment_id") or "").strip()
|
||||
if (not aid) and p_uri:
|
||||
try:
|
||||
_scope, _fid = parse_pointer_uri(p_uri)
|
||||
aid = str(_fid or "").strip()
|
||||
except Exception:
|
||||
aid = ""
|
||||
if aid and mime.startswith("image/"):
|
||||
data_url = attachment_id_to_data_url(aid, mime=mime)
|
||||
if data_url and ";base64," in data_url:
|
||||
b64 = data_url.split(";base64,", 1)[1]
|
||||
content_list.append(
|
||||
{
|
||||
"type": "input_image",
|
||||
"image_base64": b64,
|
||||
"mime": mime or "image/jpeg",
|
||||
}
|
||||
)
|
||||
rel_path = str(att.get("rel_path") or "").strip()
|
||||
sz = att.get("bytes")
|
||||
sha = str(att.get("sha256") or "").strip()
|
||||
pointer_line = f"- pointer_uri={p_uri}"
|
||||
if rel_path:
|
||||
pointer_line += f" rel_path={rel_path}"
|
||||
if mime:
|
||||
pointer_line += f" mime={mime}"
|
||||
if sz:
|
||||
pointer_line += f" bytes={sz}"
|
||||
if sha:
|
||||
pointer_line += f" sha256={sha}"
|
||||
content_list.append({"type": "text", "text": pointer_line})
|
||||
|
||||
if not content_list:
|
||||
placeholder = "(No text content)" if str(lang or "").startswith("en") else "(无文本内容)"
|
||||
content_list.append({"type": "text", "text": placeholder})
|
||||
|
||||
if len(content_list) == 1 and content_list[0].get("type") == "text":
|
||||
out.append({"role": "user", "content": content_list[0]["text"]})
|
||||
else:
|
||||
out.append({"role": "user", "content": content_list})
|
||||
continue
|
||||
|
||||
if role == "assistant":
|
||||
tool_calls = None
|
||||
raw_tc = getattr(m, "tool_calls", None)
|
||||
if raw_tc:
|
||||
try:
|
||||
tool_calls = json.loads(raw_tc) if isinstance(raw_tc, str) else raw_tc
|
||||
except Exception:
|
||||
tool_calls = None
|
||||
|
||||
if tool_calls and isinstance(tool_calls, list):
|
||||
api_tool_calls = []
|
||||
gemini_fc = gemini_openai_compat_client(model)
|
||||
for idx, tc in enumerate(tool_calls):
|
||||
if not isinstance(tc, dict) or not tc.get("id") or not tc.get("name"):
|
||||
continue
|
||||
try:
|
||||
valid_tool_call_ids.add(str(tc.get("id") or ""))
|
||||
except Exception:
|
||||
pass
|
||||
entry: dict[str, Any] = {
|
||||
"id": tc.get("id"),
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": tc.get("name"),
|
||||
"arguments": json.dumps(tc.get("arguments", {}), ensure_ascii=False),
|
||||
},
|
||||
}
|
||||
raw_sig = tc.get("thought_signature")
|
||||
if allow_signature_replay and gemini_fc:
|
||||
if isinstance(raw_sig, str):
|
||||
sig = raw_sig
|
||||
elif idx == 0:
|
||||
sig = "skip_thought_signature_validator"
|
||||
else:
|
||||
sig = ""
|
||||
entry["extra_content"] = {"google": {"thought_signature": sig}}
|
||||
elif allow_signature_replay and isinstance(raw_sig, str):
|
||||
entry["extra_content"] = {"google": {"thought_signature": raw_sig}}
|
||||
api_tool_calls.append(entry)
|
||||
if api_tool_calls:
|
||||
out.append(
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": _strip_reasoning_blocks(getattr(m, "content", "") or ""),
|
||||
"tool_calls": api_tool_calls,
|
||||
}
|
||||
)
|
||||
else:
|
||||
out.append({"role": "assistant", "content": _strip_reasoning_blocks(getattr(m, "content", "") or "")})
|
||||
else:
|
||||
out.append({"role": "assistant", "content": _strip_reasoning_blocks(getattr(m, "content", "") or "")})
|
||||
continue
|
||||
|
||||
if role == "tool":
|
||||
tool_call_id = None
|
||||
raw_tc = getattr(m, "tool_calls", None)
|
||||
if raw_tc:
|
||||
try:
|
||||
meta = json.loads(raw_tc) if isinstance(raw_tc, str) else raw_tc
|
||||
if isinstance(meta, dict):
|
||||
tool_call_id = meta.get("tool_call_id")
|
||||
except Exception:
|
||||
tool_call_id = None
|
||||
if tool_call_id is not None:
|
||||
try:
|
||||
tool_call_id = str(tool_call_id).strip()
|
||||
except Exception:
|
||||
tool_call_id = ""
|
||||
if tool_call_id:
|
||||
# Guard against dangling tool_call_id (assistant tool_calls missing from this trimmed context window).
|
||||
if str(tool_call_id) not in valid_tool_call_ids:
|
||||
# Preserve tool evidence, but downgrade to plain assistant text when pairing is broken.
|
||||
# Some OpenAI-compatible gateways reject a role=tool message if tool_call_id cannot be paired
|
||||
# to an assistant.tool_calls.id within the same request context.
|
||||
r0 = getattr(m, "content", "") or ""
|
||||
cap0 = tool_llm_message_max_chars()
|
||||
pretty = _summarize_unpaired_tool_content(r0, cap=cap0)
|
||||
out.append(
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": render_prompt(
|
||||
"tools/tool_result_unpaired.md",
|
||||
variables={"tag": "tool_use_result:unpaired", "payload": pretty},
|
||||
strict=True,
|
||||
),
|
||||
}
|
||||
)
|
||||
continue
|
||||
raw_tc_content = getattr(m, "content", "") or ""
|
||||
tool_content_out = raw_tc_content
|
||||
cap = tool_llm_message_max_chars()
|
||||
if str(tool_call_id) in historical_tool_ids:
|
||||
summary_cap = 1800
|
||||
if cap > 0:
|
||||
summary_cap = max(600, min(2400, cap // 3))
|
||||
tool_content_out = _summarize_historical_tool_content(raw_tc_content, cap=summary_cap)
|
||||
elif cap > 0 and len(raw_tc_content) > cap:
|
||||
try:
|
||||
parsed = json.loads(raw_tc_content)
|
||||
if isinstance(parsed, dict):
|
||||
tool_content_out = json.dumps(
|
||||
truncate_tool_result_for_llm_messages(parsed), ensure_ascii=False, default=str
|
||||
)
|
||||
else:
|
||||
tool_content_out = raw_tc_content[: max(1, cap - 80)] + "\n...<truncated>"
|
||||
except Exception:
|
||||
tool_content_out = raw_tc_content[: max(1, cap - 80)] + "\n...<truncated>"
|
||||
tool_row: dict[str, Any] = {
|
||||
"role": "tool",
|
||||
"tool_call_id": tool_call_id,
|
||||
"content": tool_content_out,
|
||||
}
|
||||
# Some OpenAI-compatible gateways expect `call_id` instead of `tool_call_id`.
|
||||
# Sending both (non-empty) keeps compatibility; servers should ignore unknown fields.
|
||||
tool_row["call_id"] = tool_call_id
|
||||
try:
|
||||
meta2 = json.loads(raw_tc) if isinstance(raw_tc, str) else raw_tc
|
||||
except Exception:
|
||||
meta2 = None
|
||||
if isinstance(meta2, dict) and meta2.get("name"):
|
||||
tool_row["name"] = str(meta2["name"])
|
||||
out.append(tool_row)
|
||||
else:
|
||||
r = getattr(m, "content", "") or ""
|
||||
cap2 = tool_llm_message_max_chars()
|
||||
pretty2 = _summarize_unpaired_tool_content(r, cap=cap2)
|
||||
out.append(
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": render_prompt(
|
||||
"tools/tool_result_unpaired.md",
|
||||
variables={"tag": "tool_use_result:no_id", "payload": pretty2},
|
||||
strict=True,
|
||||
),
|
||||
}
|
||||
)
|
||||
continue
|
||||
|
||||
return out
|
||||
|
||||
|
||||
__all__ = ["build_llm_messages"]
|
||||
630
runtime/chat/tool_runtime.py
Normal file
630
runtime/chat/tool_runtime.py
Normal file
|
|
@ -0,0 +1,630 @@
|
|||
from __future__ import annotations
|
||||
|
||||
"""Agent 工具执行模块。
|
||||
|
||||
本模块把“工具执行(校验/并发/落库/回写)”从 `Agent.run_turn` 中下沉出来,
|
||||
以便被单 Agent 与编排器(manager/specialist)复用。
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
import os
|
||||
import threading
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from concurrent.futures import TimeoutError as FuturesTimeoutError
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Callable, Optional
|
||||
|
||||
from oclaw.platform.persistence.sqlite_store import SqliteStore
|
||||
from oclaw.runtime.tools.base import ToolRegistry
|
||||
from oclaw.platform.llm.chat_models import LLMToolCall
|
||||
from oclaw.runtime.tools.tool_validation import validate_tool_arguments
|
||||
from oclaw.runtime.tools.experts.workspace.workspace_base import workspace_path_access_scope
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_TOOL_ERROR_MAP = {
|
||||
"tool_timeout_or_failed": "tool_timeout_or_failed",
|
||||
}
|
||||
_SQL_REPLAY_COMPACT_TOOL_NAMES = {
|
||||
"query_tabular_attachment",
|
||||
"run_tabular_sql",
|
||||
"analyze_tabular_attachment_full_scan",
|
||||
}
|
||||
|
||||
|
||||
def normalize_tool_result(result: Any) -> dict[str, Any]:
|
||||
if isinstance(result, dict):
|
||||
out = dict(result)
|
||||
if "ok" in out:
|
||||
out["ok"] = bool(out.get("ok"))
|
||||
else:
|
||||
# Backward compatibility: many lightweight tools return payload-only dicts.
|
||||
# Treat those as success unless they explicitly carry error semantics.
|
||||
has_error = bool(str(out.get("error_code") or "").strip() or str(out.get("error") or "").strip())
|
||||
out["ok"] = not has_error
|
||||
else:
|
||||
out = {"ok": False, "error": "tool_result_not_dict", "data": result}
|
||||
if not out["ok"]:
|
||||
raw_ec = str(out.get("error_code") or "").strip()
|
||||
raw_err = str(out.get("error") or "").strip()
|
||||
if not raw_ec:
|
||||
out["error_code"] = _TOOL_ERROR_MAP.get(raw_err, "tool_failed")
|
||||
return out
|
||||
|
||||
|
||||
def tool_llm_message_max_chars() -> int:
|
||||
raw = str(os.getenv("AIA_TOOL_LLM_MESSAGE_MAX_CHARS") or "").strip()
|
||||
if raw.isdigit():
|
||||
n = int(raw)
|
||||
if n == 0:
|
||||
return 0
|
||||
return max(4096, min(n, 500_000))
|
||||
return 0
|
||||
|
||||
|
||||
def tool_history_summary_after_calls() -> int:
|
||||
raw = str(os.getenv("AIA_TOOL_HISTORY_SUMMARY_AFTER_CALLS") or "").strip()
|
||||
if raw.isdigit():
|
||||
return max(0, min(int(raw), 200))
|
||||
# Default: when the same tool is called >= 3 times in one turn, keep history compact.
|
||||
return 3
|
||||
|
||||
|
||||
def _json_blob_size(obj: Any) -> int:
|
||||
try:
|
||||
return len(json.dumps(obj, ensure_ascii=False, default=str))
|
||||
except Exception:
|
||||
return len(repr(obj))
|
||||
|
||||
|
||||
def _estimate_observed_rows(result: dict[str, Any]) -> int:
|
||||
if not isinstance(result, dict):
|
||||
return 0
|
||||
try:
|
||||
rr = result.get("rows_returned")
|
||||
if isinstance(rr, (int, float)):
|
||||
return max(0, int(rr))
|
||||
except Exception:
|
||||
pass
|
||||
rows = result.get("rows")
|
||||
if isinstance(rows, list):
|
||||
return max(0, len(rows))
|
||||
nested = result.get("result")
|
||||
if isinstance(nested, dict):
|
||||
nrows = nested.get("rows")
|
||||
if isinstance(nrows, list):
|
||||
return max(0, len(nrows))
|
||||
return 0
|
||||
|
||||
|
||||
def _deep_truncate_for_llm(obj: Any, *, max_str: int, max_list: int) -> Any:
|
||||
if isinstance(obj, dict):
|
||||
return {str(k): _deep_truncate_for_llm(v, max_str=max_str, max_list=max_list) for k, v in obj.items()}
|
||||
if isinstance(obj, list):
|
||||
items = obj
|
||||
omitted = 0
|
||||
if len(items) > max_list:
|
||||
omitted = len(items) - max_list
|
||||
items = items[:max_list]
|
||||
out: list[Any] = [_deep_truncate_for_llm(x, max_str=max_str, max_list=max_list) for x in items]
|
||||
if omitted:
|
||||
out.append(f"…({omitted} more list items omitted)")
|
||||
return out
|
||||
if isinstance(obj, str) and len(obj) > max_str:
|
||||
return obj[:max_str] + "\n...<truncated>"
|
||||
return obj
|
||||
|
||||
|
||||
def partition_tool_use_batches(
|
||||
tool_uses: list[LLMToolCall],
|
||||
registry: ToolRegistry,
|
||||
) -> list[list[LLMToolCall]]:
|
||||
"""Split tool uses into ordered batches (cc-mini ``Engine.submit`` scheduling).
|
||||
|
||||
Consecutive tools whose ``ToolSpec.is_read_only()`` is true are merged into one batch
|
||||
and may run in parallel when the batch length is greater than one. Any other tool
|
||||
starts a new batch (typically length 1), which runs sequentially relative to other
|
||||
batches and uses a single worker within the batch.
|
||||
"""
|
||||
batches: list[tuple[bool, list[LLMToolCall]]] = []
|
||||
for tc in tool_uses:
|
||||
spec = registry.get(tc.name)
|
||||
is_concurrent = bool(spec and spec.is_read_only())
|
||||
if batches and batches[-1][0] == is_concurrent and is_concurrent:
|
||||
batches[-1][1].append(tc)
|
||||
else:
|
||||
batches.append((is_concurrent, [tc]))
|
||||
return [chunk for _, chunk in batches]
|
||||
|
||||
|
||||
def truncate_tool_result_for_llm_messages(result: dict[str, Any], *, max_chars: int | None = None) -> dict[str, Any]:
|
||||
"""Return a copy safe to put in ``role=tool`` ``content`` so the next LLM request stays under provider limits."""
|
||||
cap = tool_llm_message_max_chars() if max_chars is None else max(0, min(int(max_chars), 500_000))
|
||||
if cap == 0:
|
||||
return result if isinstance(result, dict) else {"ok": False, "error": "tool_result_not_dict", "data": result}
|
||||
if not isinstance(result, dict):
|
||||
return {"ok": False, "error": "tool_result_not_dict", "payload_type": type(result).__name__}
|
||||
if _json_blob_size(result) <= cap:
|
||||
return result
|
||||
orig_files_n = len(result["files"]) if isinstance(result.get("files"), list) else 0
|
||||
pairs = (
|
||||
(12_000, 800),
|
||||
(8000, 500),
|
||||
(4000, 300),
|
||||
(2000, 200),
|
||||
(1200, 120),
|
||||
(800, 80),
|
||||
(500, 50),
|
||||
(400, 40),
|
||||
)
|
||||
for max_str, max_list in pairs:
|
||||
slim = _deep_truncate_for_llm(result, max_str=max_str, max_list=max_list)
|
||||
if not isinstance(slim, dict):
|
||||
slim = {"ok": bool(result.get("ok")), "payload": slim}
|
||||
if _json_blob_size(slim) <= cap:
|
||||
slim = dict(slim)
|
||||
slim["_truncated_for_llm"] = True
|
||||
if orig_files_n and isinstance(slim.get("files"), list):
|
||||
kept = sum(1 for x in slim["files"] if isinstance(x, str))
|
||||
if kept < orig_files_n:
|
||||
slim["files_total"] = orig_files_n
|
||||
slim["files_omitted"] = orig_files_n - kept
|
||||
return slim
|
||||
return {
|
||||
"ok": bool(result.get("ok")),
|
||||
"_truncated_for_llm": True,
|
||||
"hint": (
|
||||
"Tool output exceeded model message size limits. "
|
||||
"Narrow the glob, lower max_results, or list a subdirectory. / "
|
||||
"工具输出超过模型单条消息限制,请缩小列举范围或降低 max_results。"
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ToolExecutionConfig:
|
||||
max_workers: int = 8
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ToolExecutionContext:
|
||||
store: SqliteStore
|
||||
tools: ToolRegistry
|
||||
session_id: str
|
||||
lang: str = "zh"
|
||||
user_text: str = ""
|
||||
specialist: str = ""
|
||||
task_kind: str = ""
|
||||
policy_engine: Any | None = None
|
||||
trace_id: str | None = None
|
||||
parent_span_id: str | None = None
|
||||
#: When ``session_id`` is a specialist temp chat row (no ``ui_session_owner``), use the user's UI session for ``extra_roots`` / allowlist.
|
||||
workspace_owner_session_id: str | None = None
|
||||
#: If ``get_ui_session_owner`` fails, load allowlist for this (tenant, user) from the HTTP/gateway request (``metadata``).
|
||||
path_policy_tenant_id: str | None = None
|
||||
path_policy_user_id: str | None = None
|
||||
turn_uuid: str | None = None
|
||||
|
||||
|
||||
class ToolExecutor:
|
||||
"""执行一组 tool uses,并把结果写回 store。"""
|
||||
|
||||
def __init__(self, *, config: ToolExecutionConfig | None = None):
|
||||
self.config = config or ToolExecutionConfig()
|
||||
|
||||
def _execute_tool(self, ctx: ToolExecutionContext, tc: LLMToolCall) -> tuple[dict[str, Any], int]:
|
||||
t0 = time.perf_counter()
|
||||
|
||||
tool = ctx.tools.get(tc.name)
|
||||
if not tool:
|
||||
msg = f"Unregistered tool: {tc.name}" if ctx.lang.startswith("en") else f"未注册的工具: {tc.name}"
|
||||
return {"ok": False, "error_code": "tool_not_registered", "error": msg}, int((time.perf_counter() - t0) * 1000)
|
||||
|
||||
ok, v_err = validate_tool_arguments(tool.parameters, tc.arguments)
|
||||
if not ok:
|
||||
msg = f"Invalid arguments: {v_err}" if ctx.lang.startswith("en") else f"参数不合法: {v_err}"
|
||||
return {"ok": False, "error_code": "tool_invalid_arguments", "error": msg}, int((time.perf_counter() - t0) * 1000)
|
||||
|
||||
try:
|
||||
timeout_s = getattr(tool, "timeout_s", None)
|
||||
# Default timeout for plugin tools if not specified.
|
||||
if timeout_s is None and "plugin" in getattr(tool, "tags", frozenset()):
|
||||
timeout_s = 30.0
|
||||
|
||||
def _call() -> Any:
|
||||
with workspace_path_access_scope(
|
||||
ctx.store,
|
||||
ctx.session_id,
|
||||
owner_fallback_session_id=ctx.workspace_owner_session_id,
|
||||
allowlist_tenant_id=ctx.path_policy_tenant_id,
|
||||
allowlist_user_id=ctx.path_policy_user_id,
|
||||
):
|
||||
return tool.handler(tc.arguments)
|
||||
|
||||
if isinstance(timeout_s, (int, float)) and float(timeout_s) > 0:
|
||||
ex = ThreadPoolExecutor(max_workers=1)
|
||||
fut = ex.submit(_call)
|
||||
try:
|
||||
result = fut.result(timeout=float(timeout_s))
|
||||
except FuturesTimeoutError as e:
|
||||
try:
|
||||
fut.cancel()
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
ex.shutdown(wait=False, cancel_futures=True)
|
||||
except Exception:
|
||||
ex.shutdown(wait=False)
|
||||
return {"ok": False, "error_code": "tool_timeout_or_failed", "error": "tool_timeout_or_failed", "detail": f"{type(e).__name__}: {e}"}, int(
|
||||
(time.perf_counter() - t0) * 1000
|
||||
)
|
||||
except Exception as e:
|
||||
try:
|
||||
ex.shutdown(wait=False, cancel_futures=True)
|
||||
except Exception:
|
||||
ex.shutdown(wait=False)
|
||||
return {"ok": False, "error_code": "tool_timeout_or_failed", "error": "tool_timeout_or_failed", "detail": f"{type(e).__name__}: {e}"}, int(
|
||||
(time.perf_counter() - t0) * 1000
|
||||
)
|
||||
else:
|
||||
try:
|
||||
ex.shutdown(wait=False, cancel_futures=True)
|
||||
except Exception:
|
||||
ex.shutdown(wait=False)
|
||||
else:
|
||||
result = _call()
|
||||
return normalize_tool_result(result), int((time.perf_counter() - t0) * 1000)
|
||||
except Exception as e:
|
||||
if ctx.lang.startswith("en"):
|
||||
err = {"ok": False, "error_code": "tool_execution_error", "error": f"Tool execution error: {type(e).__name__}: {e}"}
|
||||
else:
|
||||
err = {"ok": False, "error_code": "tool_execution_error", "error": f"工具执行异常: {type(e).__name__}: {e}"}
|
||||
return normalize_tool_result(err), int((time.perf_counter() - t0) * 1000)
|
||||
|
||||
@staticmethod
|
||||
def _json_dumps_safe(obj: Any) -> str:
|
||||
try:
|
||||
return json.dumps(obj, ensure_ascii=False, default=str)
|
||||
except (TypeError, ValueError):
|
||||
return json.dumps({"ok": False, "error": "tool result is not JSON-serializable"}, ensure_ascii=False)
|
||||
|
||||
def execute_tool_uses(
|
||||
self,
|
||||
*,
|
||||
ctx: ToolExecutionContext,
|
||||
assistant_msg_id: int,
|
||||
tool_uses: list[LLMToolCall],
|
||||
on_tool_ui: Optional[Callable[[str, dict[str, Any]], None]] = None,
|
||||
should_stop: Optional[Callable[[], bool]] = None,
|
||||
signature_budget: int = 2,
|
||||
) -> tuple[list[dict[str, Any]], dict[str, tuple[dict[str, Any], int]]]:
|
||||
"""执行并回写 tool messages。
|
||||
|
||||
Returns:
|
||||
- tool_messages: 用于写入对话 history 的 `role=tool` 消息 payload 列表(与 tool_uses 顺序一致)
|
||||
- results_by_id: tool_call_id -> (result_dict, duration_ms)
|
||||
"""
|
||||
|
||||
def _check_stop() -> None:
|
||||
if should_stop and should_stop():
|
||||
raise RuntimeError("generation interrupted by user")
|
||||
|
||||
def _trace(event_type: str, payload: dict[str, Any]) -> None:
|
||||
if not ctx.trace_id:
|
||||
return
|
||||
try:
|
||||
from oclaw.runtime.orchestration.trace import new_span_id
|
||||
|
||||
ctx.store.add_trace_event(
|
||||
session_id=ctx.session_id,
|
||||
trace_id=str(ctx.trace_id),
|
||||
span_id=new_span_id(),
|
||||
parent_span_id=ctx.parent_span_id,
|
||||
event_type=str(event_type),
|
||||
payload=dict(payload or {}),
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def _load_turn_tool_stats() -> tuple[dict[str, int], dict[str, int]]:
|
||||
counts: dict[str, int] = {}
|
||||
observed_rows: dict[str, int] = {}
|
||||
if not str(ctx.turn_uuid or "").strip():
|
||||
return counts, observed_rows
|
||||
try:
|
||||
rows = ctx.store.get_messages(session_id=ctx.session_id, limit=500)
|
||||
except Exception:
|
||||
return counts, observed_rows
|
||||
for m in rows or []:
|
||||
if str(getattr(m, "role", "") or "") != "tool":
|
||||
continue
|
||||
if str(getattr(m, "turn_uuid", "") or "") != str(ctx.turn_uuid or ""):
|
||||
continue
|
||||
raw_tc = getattr(m, "tool_calls", None)
|
||||
name = ""
|
||||
if isinstance(raw_tc, str):
|
||||
try:
|
||||
parsed = json.loads(raw_tc)
|
||||
except Exception:
|
||||
parsed = None
|
||||
else:
|
||||
parsed = raw_tc
|
||||
if isinstance(parsed, dict):
|
||||
name = str(parsed.get("name") or "").strip()
|
||||
if not name:
|
||||
continue
|
||||
counts[name] = int(counts.get(name, 0)) + 1
|
||||
try:
|
||||
raw_content = str(getattr(m, "content", "") or "")
|
||||
payload = json.loads(raw_content) if raw_content else {}
|
||||
except Exception:
|
||||
payload = {}
|
||||
if isinstance(payload, dict):
|
||||
observed_rows[name] = int(observed_rows.get(name, 0)) + int(
|
||||
_estimate_observed_rows(payload)
|
||||
or payload.get("_tool_observed_rows_this_call")
|
||||
or 0
|
||||
)
|
||||
return counts, observed_rows
|
||||
|
||||
def _compact_tool_result_for_history(
|
||||
*,
|
||||
tool_name: str,
|
||||
result: dict[str, Any],
|
||||
call_index: int,
|
||||
threshold: int,
|
||||
observed_rows_this_call: int,
|
||||
observed_rows_cumulative_in_turn: int,
|
||||
) -> dict[str, Any]:
|
||||
out: dict[str, Any] = {
|
||||
"ok": bool(result.get("ok")),
|
||||
"_history_compacted": True,
|
||||
"_history_compact_reason": "repeated_tool_calls_in_turn",
|
||||
"tool_name": str(tool_name or ""),
|
||||
"call_index_in_turn_for_tool": int(call_index),
|
||||
"compact_threshold": int(threshold),
|
||||
"_tool_observed_rows_this_call": int(observed_rows_this_call),
|
||||
"_tool_observed_rows_cumulative_in_turn": int(observed_rows_cumulative_in_turn),
|
||||
"result_keys": sorted(list(result.keys()))[:30],
|
||||
"result_bytes": int(_json_blob_size(result)),
|
||||
"hint": (
|
||||
"Repeated tool calls in this turn were compacted in chat history to avoid context bloat. "
|
||||
"Full payload remains in tool logs."
|
||||
),
|
||||
"audit_note": (
|
||||
"History is compacted by system optimization. If more detail is needed, continue querying "
|
||||
"with the same SQL/tool parameters from this turn."
|
||||
),
|
||||
}
|
||||
for key in ("error_code", "error", "rows_returned", "limit", "table_id", "engine"):
|
||||
if key in result:
|
||||
out[key] = result.get(key)
|
||||
for key in ("input_sql", "executed_sql"):
|
||||
v = str(result.get(key) or "").strip()
|
||||
if v:
|
||||
out[key] = v[:1200]
|
||||
guard = result.get("sql_guard")
|
||||
if isinstance(guard, dict):
|
||||
out["sql_guard"] = {
|
||||
"readonly_enforced": bool(guard.get("readonly_enforced")),
|
||||
"auto_limit_applied": bool(guard.get("auto_limit_applied")),
|
||||
"result_row_cap": int(guard.get("result_row_cap") or 0),
|
||||
}
|
||||
return out
|
||||
|
||||
_check_stop()
|
||||
if not tool_uses:
|
||||
return [], {}
|
||||
history_summary_threshold = int(tool_history_summary_after_calls())
|
||||
turn_tool_name_counts, turn_tool_observed_rows = _load_turn_tool_stats()
|
||||
local_turn_tool_name_counts: dict[str, int] = {}
|
||||
local_turn_tool_observed_rows: dict[str, int] = {}
|
||||
local_turn_written_tool_msgs: dict[str, list[dict[str, Any]]] = {}
|
||||
|
||||
results_by_id: dict[str, tuple[dict[str, Any], int]] = {}
|
||||
runnable_tool_uses: list[LLMToolCall] = []
|
||||
sig_seen: dict[str, int] = {}
|
||||
budget = max(1, min(int(signature_budget or 2), 8))
|
||||
for tc in tool_uses:
|
||||
sig = f"{tc.name}:{self._json_dumps_safe(dict(tc.arguments or {}))}"
|
||||
count = int(sig_seen.get(sig, 0))
|
||||
if count >= budget:
|
||||
results_by_id[tc.id] = (
|
||||
{
|
||||
"ok": False,
|
||||
"error_code": "tool_loop_guard",
|
||||
"error": f"tool loop guard triggered for signature: {tc.name}",
|
||||
},
|
||||
0,
|
||||
)
|
||||
_trace(
|
||||
"tool_loop_guard",
|
||||
{
|
||||
"tool_name": tc.name,
|
||||
"signature": sig[:300],
|
||||
"budget": budget,
|
||||
},
|
||||
)
|
||||
continue
|
||||
sig_seen[sig] = count + 1
|
||||
runnable_tool_uses.append(tc)
|
||||
|
||||
for batch in partition_tool_use_batches(runnable_tool_uses, ctx.tools):
|
||||
_check_stop()
|
||||
_trace(
|
||||
"tool_batch_started",
|
||||
{
|
||||
"batch_size": len(batch),
|
||||
"tool_names": [str(getattr(x, "name", "") or "") for x in batch],
|
||||
},
|
||||
)
|
||||
if len(batch) > 1:
|
||||
workers = min(int(self.config.max_workers), len(batch))
|
||||
with ThreadPoolExecutor(max_workers=workers) as ex:
|
||||
fut_to_tc = {ex.submit(self._execute_tool, ctx, tc): tc for tc in batch}
|
||||
for fut in as_completed(fut_to_tc):
|
||||
tc = fut_to_tc[fut]
|
||||
results_by_id[tc.id] = fut.result()
|
||||
else:
|
||||
for tc in batch:
|
||||
results_by_id[tc.id] = self._execute_tool(ctx, tc)
|
||||
_trace(
|
||||
"tool_batch_finished",
|
||||
{
|
||||
"batch_size": len(batch),
|
||||
"tool_names": [str(getattr(x, "name", "") or "") for x in batch],
|
||||
},
|
||||
)
|
||||
|
||||
tool_messages: list[dict[str, Any]] = []
|
||||
for tc in tool_uses:
|
||||
_check_stop()
|
||||
_trace(
|
||||
"tool_called",
|
||||
{
|
||||
"tool_name": tc.name,
|
||||
"arguments": tc.arguments,
|
||||
"arguments_bytes": _json_blob_size(tc.arguments),
|
||||
},
|
||||
)
|
||||
result, duration_ms = results_by_id[tc.id]
|
||||
result = normalize_tool_result(result)
|
||||
logger.info(
|
||||
"tool_runtime tool session=%s name=%s duration_ms=%d ok=%s",
|
||||
ctx.session_id[:12],
|
||||
tc.name,
|
||||
duration_ms,
|
||||
result.get("ok") if isinstance(result, dict) else None,
|
||||
)
|
||||
t_db1 = time.perf_counter()
|
||||
ctx.store.add_tool_log(
|
||||
session_id=ctx.session_id,
|
||||
tool_name=tc.name,
|
||||
args=tc.arguments,
|
||||
result=result,
|
||||
specialist=ctx.specialist,
|
||||
duration_ms=duration_ms,
|
||||
)
|
||||
tool_log_write_ms = int((time.perf_counter() - t_db1) * 1000)
|
||||
# Full payload stays in tool_log; chat history must stay under provider per-message limits.
|
||||
t_trunc = time.perf_counter()
|
||||
observed_rows_this_call = int(_estimate_observed_rows(result))
|
||||
result_for_llm = truncate_tool_result_for_llm_messages(result)
|
||||
should_compact_history = tc.name in _SQL_REPLAY_COMPACT_TOOL_NAMES
|
||||
if history_summary_threshold > 0 and should_compact_history:
|
||||
prior = int(turn_tool_name_counts.get(tc.name, 0))
|
||||
current = int(local_turn_tool_name_counts.get(tc.name, 0))
|
||||
call_index = prior + current + 1
|
||||
prior_rows = int(turn_tool_observed_rows.get(tc.name, 0))
|
||||
current_rows = int(local_turn_tool_observed_rows.get(tc.name, 0))
|
||||
observed_rows_cumulative_in_turn = prior_rows + current_rows + observed_rows_this_call
|
||||
if call_index >= history_summary_threshold:
|
||||
result_for_llm = _compact_tool_result_for_history(
|
||||
tool_name=tc.name,
|
||||
result=result,
|
||||
call_index=call_index,
|
||||
threshold=history_summary_threshold,
|
||||
observed_rows_this_call=observed_rows_this_call,
|
||||
observed_rows_cumulative_in_turn=observed_rows_cumulative_in_turn,
|
||||
)
|
||||
local_turn_tool_name_counts[tc.name] = current + 1
|
||||
local_turn_tool_observed_rows[tc.name] = current_rows + observed_rows_this_call
|
||||
trunc_ms = int((time.perf_counter() - t_trunc) * 1000)
|
||||
tool_content = self._json_dumps_safe(result_for_llm)
|
||||
t_db2 = time.perf_counter()
|
||||
msg_row = ctx.store.add_message(
|
||||
session_id=ctx.session_id,
|
||||
role="tool",
|
||||
content=tool_content,
|
||||
tool_calls={"tool_call_id": tc.id, "name": tc.name, "assistant_message_id": assistant_msg_id},
|
||||
turn_uuid=ctx.turn_uuid,
|
||||
event_type="tool_result",
|
||||
event_payload={"tool_name": tc.name, "observed_rows": int(observed_rows_this_call)},
|
||||
)
|
||||
tool_msg_write_ms = int((time.perf_counter() - t_db2) * 1000)
|
||||
tool_messages.append({"role": "tool", "tool_call_id": tc.id, "content": tool_content, "name": tc.name})
|
||||
tool_messages_idx = len(tool_messages) - 1
|
||||
call_index_for_tool = int(turn_tool_name_counts.get(tc.name, 0)) + int(local_turn_tool_name_counts.get(tc.name, 0))
|
||||
local_turn_written_tool_msgs.setdefault(tc.name, []).append(
|
||||
{
|
||||
"message_id": int(getattr(msg_row, "id", 0) or 0),
|
||||
"tool_messages_idx": int(tool_messages_idx),
|
||||
"result": dict(result or {}),
|
||||
"observed_rows": int(observed_rows_this_call),
|
||||
"call_index": int(call_index_for_tool),
|
||||
"compacted": bool(isinstance(result_for_llm, dict) and result_for_llm.get("_history_compacted")),
|
||||
}
|
||||
)
|
||||
# When threshold is reached for one SQL tool in the turn, retro-compact earlier same-tool tool messages too.
|
||||
if history_summary_threshold > 0 and should_compact_history and call_index_for_tool >= history_summary_threshold:
|
||||
running_rows = int(turn_tool_observed_rows.get(tc.name, 0))
|
||||
entries = list(local_turn_written_tool_msgs.get(tc.name) or [])
|
||||
for ent in entries:
|
||||
running_rows += int(ent.get("observed_rows") or 0)
|
||||
compacted_payload = _compact_tool_result_for_history(
|
||||
tool_name=tc.name,
|
||||
result=dict(ent.get("result") or {}),
|
||||
call_index=int(ent.get("call_index") or 0),
|
||||
threshold=history_summary_threshold,
|
||||
observed_rows_this_call=int(ent.get("observed_rows") or 0),
|
||||
observed_rows_cumulative_in_turn=int(running_rows),
|
||||
)
|
||||
compacted_content = self._json_dumps_safe(compacted_payload)
|
||||
if not bool(ent.get("compacted")):
|
||||
try:
|
||||
ctx.store.update_message_content(
|
||||
session_id=ctx.session_id,
|
||||
message_id=int(ent.get("message_id") or 0),
|
||||
content=compacted_content,
|
||||
event_payload={"tool_name": tc.name, "observed_rows": int(ent.get("observed_rows") or 0)},
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
ent["compacted"] = True
|
||||
ti = int(ent.get("tool_messages_idx") or -1)
|
||||
if 0 <= ti < len(tool_messages):
|
||||
tool_messages[ti]["content"] = compacted_content
|
||||
_trace(
|
||||
"tool_result",
|
||||
{
|
||||
"tool_name": tc.name,
|
||||
"duration_ms": duration_ms,
|
||||
"ok": bool(result.get("ok")) if isinstance(result, dict) else None,
|
||||
"error_code": str(result.get("error_code") or "") if isinstance(result, dict) else "",
|
||||
"result_bytes": _json_blob_size(result),
|
||||
"result_for_llm_bytes": len(tool_content or ""),
|
||||
"tool_log_write_ms": tool_log_write_ms,
|
||||
"tool_message_write_ms": tool_msg_write_ms,
|
||||
"truncate_ms": trunc_ms,
|
||||
"active_threads": int(threading.active_count()),
|
||||
},
|
||||
)
|
||||
if on_tool_ui:
|
||||
truncated_for_llm = bool(
|
||||
isinstance(result_for_llm, dict) and result_for_llm.get("_truncated_for_llm")
|
||||
)
|
||||
payload = {
|
||||
"name": tc.name,
|
||||
"result": result,
|
||||
"llm_wire": {
|
||||
"truncated_for_llm": truncated_for_llm,
|
||||
"max_chars": int(tool_llm_message_max_chars()),
|
||||
"result_bytes": int(_json_blob_size(result)),
|
||||
"result_for_llm_bytes": int(len(tool_content or "")),
|
||||
"truncate_ms": int(trunc_ms),
|
||||
},
|
||||
}
|
||||
on_tool_ui("tool_use_result", payload)
|
||||
return tool_messages, results_by_id
|
||||
|
||||
__all__ = [
|
||||
"ToolExecutionConfig",
|
||||
"ToolExecutionContext",
|
||||
"ToolExecutor",
|
||||
"normalize_tool_result",
|
||||
"partition_tool_use_batches",
|
||||
"tool_llm_message_max_chars",
|
||||
"truncate_tool_result_for_llm_messages",
|
||||
]
|
||||
16
runtime/chat/turn_types.py
Normal file
16
runtime/chat/turn_types.py
Normal file
|
|
@ -0,0 +1,16 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TurnRunOutcome:
|
||||
final_text: str
|
||||
tool_traces: tuple[dict[str, Any], ...] = ()
|
||||
handoff_note: str = ""
|
||||
turn_uuid: str = ""
|
||||
|
||||
|
||||
__all__ = ["TurnRunOutcome"]
|
||||
|
||||
Loading…
Add table
Add a link
Reference in a new issue