实现附件处理链路的统一引用化与可检索增强,避免大文件/多模态内容直接撑爆上下文并提升工具可用性。

本次补齐 text/image/video/archive 的标准化处理、会话级工具守卫、回放压缩、配置与文档对齐,并修复表格与流式输出相关体验问题。

Made-with: Cursor
This commit is contained in:
oliver 2026-04-26 15:25:50 +08:00
parent 9d2900db02
commit 37a2ef35f4
27 changed files with 3227 additions and 278 deletions

View file

@ -11,6 +11,7 @@ import logging
import time
import os
import threading
from pathlib import Path
from concurrent.futures import ThreadPoolExecutor, as_completed
from concurrent.futures import TimeoutError as FuturesTimeoutError
from dataclasses import dataclass
@ -20,7 +21,10 @@ 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
from oclaw.runtime.tools.experts.workspace.workspace_base import (
workspace_path_access_scope,
workspace_write_namespace_scope,
)
logger = logging.getLogger(__name__)
_tool_exec_log = logging.getLogger("oclaw.tool_exec")
@ -33,6 +37,170 @@ _SQL_REPLAY_COMPACT_TOOL_NAMES = {
"run_tabular_sql",
"analyze_tabular_attachment_full_scan",
}
_TABULAR_QUERY_TOOL_NAMES = {
"query_tabular_attachment",
"run_tabular_sql",
"analyze_tabular_attachment_full_scan",
}
_TEXT_QUERY_TOOL_NAMES = {
"query_text_attachment",
}
_IMAGE_QUERY_TOOL_NAMES = {
"query_image_attachment",
}
_VIDEO_QUERY_TOOL_NAMES = {
"query_video_attachment",
}
def _message_has_tabular_ref(raw_attachments: Any) -> bool:
if raw_attachments is None:
return False
obj = raw_attachments
if isinstance(raw_attachments, str):
s = str(raw_attachments or "").strip()
if not s:
return False
try:
obj = json.loads(s)
except Exception:
return False
if isinstance(obj, dict):
items = [obj]
elif isinstance(obj, list):
items = obj
else:
return False
for it in items:
if not isinstance(it, dict):
continue
if str(it.get("type") or "").strip().lower() == "tabular_ref":
return True
return False
def _message_has_text_ref(raw_attachments: Any) -> bool:
if raw_attachments is None:
return False
obj = raw_attachments
if isinstance(raw_attachments, str):
s = str(raw_attachments or "").strip()
if not s:
return False
try:
obj = json.loads(s)
except Exception:
return False
if isinstance(obj, dict):
items = [obj]
elif isinstance(obj, list):
items = obj
else:
return False
for it in items:
if not isinstance(it, dict):
continue
if str(it.get("type") or "").strip().lower() == "text_ref":
return True
return False
def _message_has_image_ref(raw_attachments: Any) -> bool:
if raw_attachments is None:
return False
obj = raw_attachments
if isinstance(raw_attachments, str):
s = str(raw_attachments or "").strip()
if not s:
return False
try:
obj = json.loads(s)
except Exception:
return False
if isinstance(obj, dict):
items = [obj]
elif isinstance(obj, list):
items = obj
else:
return False
for it in items:
if not isinstance(it, dict):
continue
t = str(it.get("type") or "").strip().lower()
if t in {"image_ref", "image", "input_image"}:
return True
return False
def _message_has_video_ref(raw_attachments: Any) -> bool:
if raw_attachments is None:
return False
obj = raw_attachments
if isinstance(raw_attachments, str):
s = str(raw_attachments or "").strip()
if not s:
return False
try:
obj = json.loads(s)
except Exception:
return False
if isinstance(obj, dict):
items = [obj]
elif isinstance(obj, list):
items = obj
else:
return False
for it in items:
if not isinstance(it, dict):
continue
t = str(it.get("type") or "").strip().lower()
if t == "video_ref":
return True
return False
def _session_has_tabular_ref(store: Any, session_id: str, *, limit: int = 300) -> bool:
try:
rows = store.get_messages(session_id=session_id, limit=max(1, int(limit)))
except Exception:
return False
for m in rows or []:
if _message_has_tabular_ref(getattr(m, "attachments", None)):
return True
return False
def _session_has_text_ref(store: Any, session_id: str, *, limit: int = 300) -> bool:
try:
rows = store.get_messages(session_id=session_id, limit=max(1, int(limit)))
except Exception:
return False
for m in rows or []:
if _message_has_text_ref(getattr(m, "attachments", None)):
return True
return False
def _session_has_image_ref(store: Any, session_id: str, *, limit: int = 300) -> bool:
try:
rows = store.get_messages(session_id=session_id, limit=max(1, int(limit)))
except Exception:
return False
for m in rows or []:
if _message_has_image_ref(getattr(m, "attachments", None)):
return True
return False
def _session_has_video_ref(store: Any, session_id: str, *, limit: int = 300) -> bool:
try:
rows = store.get_messages(session_id=session_id, limit=max(1, int(limit)))
except Exception:
return False
for m in rows or []:
if _message_has_video_ref(getattr(m, "attachments", None)):
return True
return False
def normalize_tool_result(result: Any) -> dict[str, Any]:
@ -206,6 +374,7 @@ class ToolExecutionContext:
#: 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
workspace_dir: str | None = None
turn_uuid: str | None = None
@ -235,13 +404,21 @@ class ToolExecutor:
timeout_s = 30.0
def _call() -> Any:
ws_ns = ""
raw_ws = str(ctx.workspace_dir or "").strip()
if raw_ws:
try:
wp = Path(raw_ws)
ws_ns = str(wp.name or wp.stem or "").strip()
except Exception:
ws_ns = ""
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,
):
), workspace_write_namespace_scope(ws_ns):
return tool.handler(tc.arguments)
if isinstance(timeout_s, (int, float)) and float(timeout_s) > 0:
@ -400,65 +577,98 @@ class ToolExecutor:
)
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]]] = {}
has_tabular_ref_in_session = _session_has_tabular_ref(ctx.store, ctx.session_id)
has_text_ref_in_session = _session_has_text_ref(ctx.store, ctx.session_id)
has_image_ref_in_session = _session_has_image_ref(ctx.store, ctx.session_id)
has_video_ref_in_session = _session_has_video_ref(ctx.store, ctx.session_id)
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:
if tc.name in _TABULAR_QUERY_TOOL_NAMES and not has_tabular_ref_in_session:
results_by_id[tc.id] = (
{
"ok": False,
"error_code": "tabular_ref_missing",
"error": "tabular_ref_missing",
"hint": "No tabular_ref attachment found in this session. Query tools require table_id from tabular_ref.",
},
0,
)
_trace(
"tabular_query_guard",
{
"tool_name": tc.name,
"blocked": True,
"reason": "tabular_ref_missing",
},
)
continue
if tc.name in _TEXT_QUERY_TOOL_NAMES and not has_text_ref_in_session:
results_by_id[tc.id] = (
{
"ok": False,
"error_code": "text_ref_missing",
"error": "text_ref_missing",
"hint": "No text_ref attachment found in this session. Query tools require text_id from text_ref.",
},
0,
)
_trace(
"text_query_guard",
{
"tool_name": tc.name,
"blocked": True,
"reason": "text_ref_missing",
},
)
continue
if tc.name in _IMAGE_QUERY_TOOL_NAMES and not has_image_ref_in_session:
results_by_id[tc.id] = (
{
"ok": False,
"error_code": "image_ref_missing",
"error": "image_ref_missing",
"hint": "No image_ref attachment found in this session. Query tools require attachment_id from image_ref.",
},
0,
)
_trace(
"image_query_guard",
{
"tool_name": tc.name,
"blocked": True,
"reason": "image_ref_missing",
},
)
continue
if tc.name in _VIDEO_QUERY_TOOL_NAMES and not has_video_ref_in_session:
results_by_id[tc.id] = (
{
"ok": False,
"error_code": "video_ref_missing",
"error": "video_ref_missing",
"hint": "No video_ref attachment found in this session. Query tools require attachment_id from video_ref.",
},
0,
)
_trace(
"video_query_guard",
{
"tool_name": tc.name,
"blocked": True,
"reason": "video_ref_missing",
},
)
continue
sig = f"{tc.name}:{self._json_dumps_safe(dict(tc.arguments or {}))}"
count = int(sig_seen.get(sig, 0))
if count >= budget:
@ -539,27 +749,14 @@ class ToolExecutor:
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.
# Keep full payload during the active turn. History compaction is deferred
# until the turn finishes, so current-round model context remains lossless.
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))
result_for_llm = dict(result or {})
if tc.name in _SQL_REPLAY_COMPACT_TOOL_NAMES:
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)
@ -576,47 +773,6 @@ class ToolExecutor:
)
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",
{
@ -650,6 +806,57 @@ class ToolExecutor:
on_tool_ui("tool_use_result", payload)
return tool_messages, results_by_id
def compact_turn_tool_messages_for_storage(
*,
store: Any,
session_id: str,
turn_uuid: str | None,
) -> dict[str, int]:
"""Compact persisted tool messages after turn completion.
This intentionally runs *after* the active turn so current-round model
context is not affected by truncation/compaction.
"""
tid = str(turn_uuid or "").strip()
if not tid:
return {"scanned": 0, "updated": 0}
try:
rows = store.get_messages(session_id=session_id, limit=800)
except Exception:
return {"scanned": 0, "updated": 0}
scanned = 0
updated = 0
for m in rows or []:
if str(getattr(m, "role", "") or "") != "tool":
continue
if str(getattr(m, "turn_uuid", "") or "") != tid:
continue
scanned += 1
raw = str(getattr(m, "content", "") or "")
if not raw.strip():
continue
try:
obj = json.loads(raw)
except Exception:
continue
if not isinstance(obj, dict):
continue
compacted = truncate_tool_result_for_llm_messages(obj)
if compacted == obj:
continue
try:
store.update_message_content(
session_id=session_id,
message_id=int(getattr(m, "id", 0) or 0),
content=json.dumps(compacted, ensure_ascii=False, default=str),
event_payload=getattr(m, "event_payload", None),
)
updated += 1
except Exception:
continue
return {"scanned": int(scanned), "updated": int(updated)}
__all__ = [
"ToolExecutionConfig",
"ToolExecutionContext",
@ -658,4 +865,5 @@ __all__ = [
"partition_tool_use_batches",
"tool_llm_message_max_chars",
"truncate_tool_result_for_llm_messages",
"compact_turn_tool_messages_for_storage",
]