mirror of
https://github.com/hansjone/oclaw.git
synced 2026-10-10 22:15:57 +08:00
实现附件处理链路的统一引用化与可检索增强,避免大文件/多模态内容直接撑爆上下文并提升工具可用性。
本次补齐 text/image/video/archive 的标准化处理、会话级工具守卫、回放压缩、配置与文档对齐,并修复表格与流式输出相关体验问题。 Made-with: Cursor
This commit is contained in:
parent
9d2900db02
commit
37a2ef35f4
27 changed files with 3227 additions and 278 deletions
166
tests/test_attachment_text_replay_guard.py
Normal file
166
tests/test_attachment_text_replay_guard.py
Normal file
|
|
@ -0,0 +1,166 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
from oclaw.platform.persistence.sqlite_store import SqliteStore
|
||||
from oclaw.runtime.direct_loop import _build_model_context
|
||||
from oclaw.platform.llm.chat_models import RuleBasedChatModel
|
||||
|
||||
|
||||
def test_large_text_attachment_is_guarded_in_history_context(tmp_path: Path) -> None:
|
||||
store = SqliteStore(str(tmp_path / "t.sqlite"))
|
||||
sess = store.create_session("t")
|
||||
big = "X" * 10000
|
||||
attachments = [{"type": "text", "name": "big.txt", "content": big}]
|
||||
store.add_message(session_id=sess.id, role="user", content="hi", attachments=attachments)
|
||||
|
||||
msgs = _build_model_context(
|
||||
store=store,
|
||||
session_id=sess.id,
|
||||
max_messages=50,
|
||||
system_prompt="sys",
|
||||
model=RuleBasedChatModel(),
|
||||
lang="en",
|
||||
memory_context=None,
|
||||
trace_id=None,
|
||||
parent_span_id=None,
|
||||
tools=None,
|
||||
base_url="",
|
||||
user_text="",
|
||||
prompt_build_context=None,
|
||||
active_turn_uuid="different-turn",
|
||||
)
|
||||
# Ensure guard marker appears in injected context.
|
||||
joined = json.dumps(msgs, ensure_ascii=False)
|
||||
assert "Attachment (summarized for context replay)" in joined
|
||||
assert "attachment_truncated_for_context_replay" in joined
|
||||
|
||||
|
||||
def test_text_attachment_is_collapsed_when_text_ref_present(tmp_path: Path) -> None:
|
||||
store = SqliteStore(str(tmp_path / "t2.sqlite"))
|
||||
sess = store.create_session("t")
|
||||
big = "Y" * 6000
|
||||
attachments = [
|
||||
{"type": "text", "name": "doc.txt", "content": big},
|
||||
{"type": "text_ref", "name": "doc.txt", "text_id": "a" * 64, "chars": 6000, "chunks": 4, "source_kind": "txt"},
|
||||
]
|
||||
store.add_message(session_id=sess.id, role="user", content="hi", attachments=attachments)
|
||||
|
||||
msgs = _build_model_context(
|
||||
store=store,
|
||||
session_id=sess.id,
|
||||
max_messages=50,
|
||||
system_prompt="sys",
|
||||
model=RuleBasedChatModel(),
|
||||
lang="en",
|
||||
memory_context=None,
|
||||
trace_id=None,
|
||||
parent_span_id=None,
|
||||
tools=None,
|
||||
base_url="",
|
||||
user_text="",
|
||||
prompt_build_context=None,
|
||||
active_turn_uuid="different-turn",
|
||||
)
|
||||
joined = json.dumps(msgs, ensure_ascii=False)
|
||||
assert "Attachment (collapsed; text_ref available)" in joined
|
||||
assert "attachment_collapsed_for_context_replay" in joined
|
||||
|
||||
|
||||
def test_large_image_tool_result_is_guarded_in_history_context(tmp_path: Path) -> None:
|
||||
store = SqliteStore(str(tmp_path / "t3.sqlite"))
|
||||
sess = store.create_session("t")
|
||||
long_text = "OCR-LINE\n" * 1200
|
||||
payload = {
|
||||
"ok": True,
|
||||
"task": "ocr",
|
||||
"attachment_id": "b" * 64,
|
||||
"text": long_text,
|
||||
"backend_shape": "multi",
|
||||
}
|
||||
store.add_message(session_id=sess.id, role="tool", content=json.dumps(payload, ensure_ascii=False))
|
||||
|
||||
msgs = _build_model_context(
|
||||
store=store,
|
||||
session_id=sess.id,
|
||||
max_messages=50,
|
||||
system_prompt="sys",
|
||||
model=RuleBasedChatModel(),
|
||||
lang="en",
|
||||
memory_context=None,
|
||||
trace_id=None,
|
||||
parent_span_id=None,
|
||||
tools=None,
|
||||
base_url="",
|
||||
user_text="",
|
||||
prompt_build_context=None,
|
||||
active_turn_uuid="different-turn",
|
||||
)
|
||||
joined = json.dumps(msgs, ensure_ascii=False)
|
||||
assert "image_tool_result_truncated_for_context_replay" in joined
|
||||
assert "_image_tool_result_guarded" in joined
|
||||
|
||||
|
||||
def test_small_image_tool_result_not_guarded(tmp_path: Path) -> None:
|
||||
store = SqliteStore(str(tmp_path / "t4.sqlite"))
|
||||
sess = store.create_session("t")
|
||||
payload = {
|
||||
"ok": True,
|
||||
"task": "describe",
|
||||
"attachment_id": "c" * 64,
|
||||
"text": "A concise description of an icon.",
|
||||
}
|
||||
store.add_message(session_id=sess.id, role="tool", content=json.dumps(payload, ensure_ascii=False))
|
||||
|
||||
msgs = _build_model_context(
|
||||
store=store,
|
||||
session_id=sess.id,
|
||||
max_messages=50,
|
||||
system_prompt="sys",
|
||||
model=RuleBasedChatModel(),
|
||||
lang="en",
|
||||
memory_context=None,
|
||||
trace_id=None,
|
||||
parent_span_id=None,
|
||||
tools=None,
|
||||
base_url="",
|
||||
user_text="",
|
||||
prompt_build_context=None,
|
||||
active_turn_uuid="different-turn",
|
||||
)
|
||||
joined = json.dumps(msgs, ensure_ascii=False)
|
||||
assert "image_tool_result_truncated_for_context_replay" not in joined
|
||||
|
||||
|
||||
def test_large_video_transcript_tool_result_is_guarded_in_history_context(tmp_path: Path) -> None:
|
||||
store = SqliteStore(str(tmp_path / "t5.sqlite"))
|
||||
sess = store.create_session("t")
|
||||
long_text = "LINE\n" * 2000
|
||||
payload = {
|
||||
"ok": True,
|
||||
"task": "transcript",
|
||||
"attachment_id": "d" * 64,
|
||||
"text": long_text,
|
||||
}
|
||||
store.add_message(session_id=sess.id, role="tool", content=json.dumps(payload, ensure_ascii=False))
|
||||
|
||||
msgs = _build_model_context(
|
||||
store=store,
|
||||
session_id=sess.id,
|
||||
max_messages=50,
|
||||
system_prompt="sys",
|
||||
model=RuleBasedChatModel(),
|
||||
lang="en",
|
||||
memory_context=None,
|
||||
trace_id=None,
|
||||
parent_span_id=None,
|
||||
tools=None,
|
||||
base_url="",
|
||||
user_text="",
|
||||
prompt_build_context=None,
|
||||
active_turn_uuid="different-turn",
|
||||
)
|
||||
joined = json.dumps(msgs, ensure_ascii=False)
|
||||
assert "video_tool_result_truncated_for_context_replay" in joined
|
||||
|
||||
|
|
@ -2,6 +2,8 @@ from __future__ import annotations
|
|||
|
||||
import io
|
||||
import json
|
||||
import gzip
|
||||
import tarfile
|
||||
import zipfile
|
||||
from pathlib import Path
|
||||
|
||||
|
|
@ -15,6 +17,7 @@ from oclaw.platform.files.tabular_attachment_store import (
|
|||
run_table_sql,
|
||||
save_workbook,
|
||||
)
|
||||
from oclaw.platform.files.text_attachment_store import query_text_document
|
||||
|
||||
|
||||
def _zip_bytes(files: dict[str, bytes]) -> bytes:
|
||||
|
|
@ -25,6 +28,17 @@ def _zip_bytes(files: dict[str, bytes]) -> bytes:
|
|||
return buf.getvalue()
|
||||
|
||||
|
||||
def _tar_bytes(files: dict[str, bytes], *, gz: bool = False) -> bytes:
|
||||
buf = io.BytesIO()
|
||||
mode = "w:gz" if gz else "w"
|
||||
with tarfile.open(fileobj=buf, mode=mode) as tf:
|
||||
for name, data in files.items():
|
||||
ti = tarfile.TarInfo(name=name)
|
||||
ti.size = len(data)
|
||||
tf.addfile(ti, io.BytesIO(data))
|
||||
return buf.getvalue()
|
||||
|
||||
|
||||
def test_csv_is_summarized_not_full_dump() -> None:
|
||||
payload = (
|
||||
"c1,c2,c3\n"
|
||||
|
|
@ -37,6 +51,32 @@ def test_csv_is_summarized_not_full_dump() -> None:
|
|||
assert "# Table Summary" in content
|
||||
assert "rows: 2" in content
|
||||
assert "cols: 3" in content
|
||||
assert "## Full table included (2 rows)" in content
|
||||
|
||||
|
||||
def test_long_txt_emits_text_ref_and_can_query() -> None:
|
||||
payload = ("A" * 15000 + "\nneedle\n" + "B" * 2000).encode("utf-8")
|
||||
out = process_file_data("notes.txt", payload)
|
||||
text_refs = [x for x in out if isinstance(x, dict) and str(x.get("type") or "") == "text_ref"]
|
||||
assert text_refs
|
||||
text_id = str(text_refs[0].get("text_id") or "")
|
||||
assert text_id
|
||||
got = query_text_document(text_id=text_id, query="needle", top_k=3, offset=0)
|
||||
assert bool(got.get("ok"))
|
||||
rows = got.get("rows") or []
|
||||
assert rows
|
||||
assert any("needle" in str(x.get("content") or "") for x in rows)
|
||||
|
||||
|
||||
def test_video_upload_emits_video_ref() -> None:
|
||||
# Minimal MP4 header-like bytes; parser should not crash and should still store as video_ref.
|
||||
payload = b"\x00\x00\x00\x18ftypmp42\x00\x00\x00\x00mp42isom"
|
||||
out = process_file_data("clip.mp4", payload)
|
||||
vrefs = [x for x in out if isinstance(x, dict) and str(x.get("type") or "") == "video_ref"]
|
||||
assert vrefs
|
||||
vr = vrefs[0]
|
||||
assert str(vr.get("attachment_id") or "")
|
||||
assert str(vr.get("mime") or "").startswith("video/")
|
||||
|
||||
|
||||
def test_zip_file_count_limit_returns_error_attachment() -> None:
|
||||
|
|
@ -48,6 +88,104 @@ def test_zip_file_count_limit_returns_error_attachment() -> None:
|
|||
assert "too many files" in str(first.get("content") or "").lower()
|
||||
|
||||
|
||||
def test_tar_path_traversal_is_blocked() -> None:
|
||||
payload = _tar_bytes({"../evil.txt": b"x"})
|
||||
out = process_file_data("unsafe.tar", payload)
|
||||
assert out
|
||||
first = out[0]
|
||||
assert str(first.get("name") or "") == "archive-error"
|
||||
assert "unsafe path" in str(first.get("content") or "").lower()
|
||||
|
||||
|
||||
def test_tgz_is_parsed_and_members_processed() -> None:
|
||||
payload = _tar_bytes({"ok.txt": b"hello"}, gz=True)
|
||||
out = process_file_data("bundle.tgz", payload)
|
||||
assert out
|
||||
assert any(str(x.get("type") or "") == "text" and "hello" in str(x.get("content") or "") for x in out)
|
||||
|
||||
|
||||
def test_tar_link_entry_is_explicitly_rejected() -> None:
|
||||
buf = io.BytesIO()
|
||||
with tarfile.open(fileobj=buf, mode="w") as tf:
|
||||
ti = tarfile.TarInfo(name="target.txt")
|
||||
data = b"ok"
|
||||
ti.size = len(data)
|
||||
tf.addfile(ti, io.BytesIO(data))
|
||||
lnk = tarfile.TarInfo(name="sym")
|
||||
lnk.type = tarfile.SYMTYPE
|
||||
lnk.linkname = "target.txt"
|
||||
tf.addfile(lnk)
|
||||
out = process_file_data("has-link.tar", buf.getvalue())
|
||||
assert out
|
||||
first = out[0]
|
||||
assert str(first.get("name") or "") == "archive-error"
|
||||
assert str(first.get("error_code") or "") == "archive_link_entry_forbidden"
|
||||
|
||||
|
||||
def test_gz_single_file_is_decompressed_and_processed() -> None:
|
||||
payload = gzip.compress(b"hello-gz")
|
||||
out = process_file_data("note.txt.gz", payload)
|
||||
assert out
|
||||
assert any(str(x.get("type") or "") == "text" and "hello-gz" in str(x.get("content") or "") for x in out)
|
||||
|
||||
|
||||
def test_rar_returns_explicit_unsupported_error() -> None:
|
||||
out = process_file_data("bundle.rar", b"not-a-real-rar")
|
||||
assert out
|
||||
first = out[0]
|
||||
assert str(first.get("name") or "") == "archive-error"
|
||||
assert "unsupported archive format" in str(first.get("content") or "").lower()
|
||||
|
||||
|
||||
def test_7z_returns_explicit_unsupported_error() -> None:
|
||||
out = process_file_data("bundle.7z", b"not-a-real-7z")
|
||||
assert out
|
||||
first = out[0]
|
||||
assert str(first.get("name") or "") == "archive-error"
|
||||
assert "unsupported archive format" in str(first.get("content") or "").lower()
|
||||
assert str(first.get("error_code") or "") == "archive_unsupported_format"
|
||||
|
||||
|
||||
def test_archive_signature_detection_overrides_misleading_extension() -> None:
|
||||
payload = _zip_bytes({"ok.txt": b"sig-ok"})
|
||||
out = process_file_data("misleading.tar", payload)
|
||||
assert out
|
||||
assert any(str(x.get("type") or "") == "text" and "sig-ok" in str(x.get("content") or "") for x in out)
|
||||
|
||||
|
||||
def test_archive_limits_can_be_overridden_by_config(tmp_path: Path, monkeypatch) -> None:
|
||||
cfg = {
|
||||
"plugins": {
|
||||
"entries": {
|
||||
"memory-wiki": {
|
||||
"auto": {
|
||||
"attachments": {
|
||||
"tabular": {
|
||||
"archive_max_depth": 1,
|
||||
"archive_max_file_count": 2,
|
||||
"archive_max_entry_bytes": 20,
|
||||
"archive_max_total_uncompressed_bytes": 30,
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
cfg_path = tmp_path / "oclaw.json"
|
||||
cfg_path.write_text(json.dumps(cfg), encoding="utf-8")
|
||||
monkeypatch.setenv("AIA_OCLAW_CONFIG_PATH", str(cfg_path))
|
||||
fa._attachments_limits.cache_clear()
|
||||
try:
|
||||
payload = _zip_bytes({"a.txt": b"x" * 50})
|
||||
out = process_file_data("limited.zip", payload)
|
||||
finally:
|
||||
fa._attachments_limits.cache_clear()
|
||||
assert out
|
||||
first = out[0]
|
||||
assert "entry too large" in str(first.get("content") or "").lower()
|
||||
|
||||
|
||||
def test_nested_zip_depth_limit_is_enforced() -> None:
|
||||
z4 = _zip_bytes({"deep.txt": b"hello"})
|
||||
z3 = _zip_bytes({"z4.zip": z4})
|
||||
|
|
@ -197,6 +335,14 @@ def test_large_csv_emits_tabular_ref_and_can_query() -> None:
|
|||
assert str(got.get("engine") or "") in {"builtin_sqlite", "mcp_sqlite"}
|
||||
|
||||
|
||||
def test_medium_csv_enters_tool_mode_with_default_threshold() -> None:
|
||||
rows = ["c1,c2"] + [f"{i},row-{i}" for i in range(0, 6000)]
|
||||
payload = ("\n".join(rows)).encode("utf-8")
|
||||
out = process_file_data("medium-default.csv", payload)
|
||||
tab_refs = [x for x in out if isinstance(x, dict) and str(x.get("type") or "") == "tabular_ref"]
|
||||
assert tab_refs
|
||||
|
||||
|
||||
def test_large_csv_can_aggregate_grouped_sum() -> None:
|
||||
rows = ["dept,amount"] + [f"{'A' if i % 2 == 0 else 'B'},{i}" for i in range(0, 25010)]
|
||||
payload = ("\n".join(rows)).encode("utf-8")
|
||||
|
|
@ -324,6 +470,20 @@ def test_query_can_target_specific_excel_sheet() -> None:
|
|||
assert rows and str(rows[0].get("k") or "") == "b"
|
||||
|
||||
|
||||
def test_multi_sheet_query_without_sheet_returns_sheet_hint() -> None:
|
||||
s1 = pd.DataFrame([["a", "1"]], columns=["k", "v"])
|
||||
s2 = pd.DataFrame([["b", "2"]], columns=["k", "v"])
|
||||
meta = save_workbook(attachment_id="aid3", name="multi2.xlsx", sheets={"Main": s1, "Backup": s2})
|
||||
tid = str(meta.get("table_id") or "")
|
||||
got = query_table(table_id=tid, columns=["k", "v"], limit=5, offset=0)
|
||||
assert bool(got.get("ok"))
|
||||
hint = got.get("sheet_hint") or {}
|
||||
assert isinstance(hint, dict)
|
||||
names = hint.get("available_sheets") or []
|
||||
assert "Main" in names and "Backup" in names
|
||||
assert str(hint.get("default_sheet") or "") == "Main"
|
||||
|
||||
|
||||
def test_full_scan_analyzes_all_rows_and_returns_audit() -> None:
|
||||
rows = ["dept,score"] + [f"{'A' if i % 2 == 0 else 'B'},{i % 5}" for i in range(0, 25025)]
|
||||
payload = ("\n".join(rows)).encode("utf-8")
|
||||
|
|
@ -345,6 +505,25 @@ def test_full_scan_analyzes_all_rows_and_returns_audit() -> None:
|
|||
assert isinstance(tops, list) and len(tops) <= 2
|
||||
|
||||
|
||||
def test_tabular_tools_reject_non_table_id_filename() -> None:
|
||||
bad_id = "日志.xlsx"
|
||||
q = query_table(table_id=bad_id, columns=["a"], limit=5, offset=0)
|
||||
assert not bool(q.get("ok"))
|
||||
assert str(q.get("error") or "") == "table_id_invalid_format"
|
||||
|
||||
ag = aggregate_table(table_id=bad_id, metric="count")
|
||||
assert not bool(ag.get("ok"))
|
||||
assert str(ag.get("error") or "") == "table_id_invalid_format"
|
||||
|
||||
sql = run_table_sql(table_id=bad_id, sql="SELECT 1", limit=5)
|
||||
assert not bool(sql.get("ok"))
|
||||
assert str(sql.get("error") or "") == "table_id_invalid_format"
|
||||
|
||||
fs = analyze_table_full_scan(table_id=bad_id, columns=["a"], top_values_limit=1)
|
||||
assert not bool(fs.get("ok"))
|
||||
assert str(fs.get("error") or "") == "table_id_invalid_format"
|
||||
|
||||
|
||||
def test_xlsx_zip_safety_blocks_unsafe_paths() -> None:
|
||||
payload = _zip_bytes({"../xl/workbook.xml": b"x"})
|
||||
out = process_file_data("bad.xlsx", payload)
|
||||
|
|
|
|||
|
|
@ -123,3 +123,272 @@ def test_repeated_non_sql_tools_are_not_compacted(tmp_path: Path) -> None:
|
|||
assert payloads
|
||||
assert not any(bool(p.get("_history_compacted")) for p in payloads)
|
||||
|
||||
|
||||
def test_tabular_tools_blocked_without_tabular_ref(tmp_path: Path) -> None:
|
||||
store = SqliteStore(str(tmp_path / "g4.sqlite"))
|
||||
sess = store.create_session("t")
|
||||
# A normal user turn without tabular_ref attachments.
|
||||
store.add_message(session_id=sess.id, role="user", content="analyze this file", attachments=[{"type": "text", "name": "x"}])
|
||||
|
||||
def _handler(_args):
|
||||
return {"ok": True, "rows": []}
|
||||
|
||||
reg = ToolRegistry(
|
||||
[
|
||||
ToolSpec(
|
||||
name="query_tabular_attachment",
|
||||
description="query table",
|
||||
parameters={"type": "object", "properties": {"table_id": {"type": "string"}}},
|
||||
handler=_handler,
|
||||
read_only=True,
|
||||
)
|
||||
]
|
||||
)
|
||||
tool_uses = [LLMToolCall(id="c1", name="query_tabular_attachment", arguments={"table_id": "日志.xlsx"})]
|
||||
_, results = ToolExecutor().execute_tool_uses(
|
||||
ctx=ToolExecutionContext(store=store, tools=reg, session_id=sess.id),
|
||||
assistant_msg_id=1,
|
||||
tool_uses=tool_uses,
|
||||
)
|
||||
blocked, _ = results["c1"]
|
||||
assert not bool(blocked.get("ok"))
|
||||
assert str(blocked.get("error_code") or "") == "tabular_ref_missing"
|
||||
|
||||
|
||||
def test_tabular_tools_allowed_with_tabular_ref(tmp_path: Path) -> None:
|
||||
store = SqliteStore(str(tmp_path / "g5.sqlite"))
|
||||
sess = store.create_session("t")
|
||||
store.add_message(
|
||||
session_id=sess.id,
|
||||
role="user",
|
||||
content="uploaded table",
|
||||
attachments=[{"type": "tabular_ref", "table_id": "a" * 64, "name": "ok.xlsx"}],
|
||||
)
|
||||
calls = {"n": 0}
|
||||
|
||||
def _handler(_args):
|
||||
calls["n"] += 1
|
||||
return {"ok": True, "rows": [{"x": 1}]}
|
||||
|
||||
reg = ToolRegistry(
|
||||
[
|
||||
ToolSpec(
|
||||
name="query_tabular_attachment",
|
||||
description="query table",
|
||||
parameters={"type": "object", "properties": {"table_id": {"type": "string"}}},
|
||||
handler=_handler,
|
||||
read_only=True,
|
||||
)
|
||||
]
|
||||
)
|
||||
tool_uses = [LLMToolCall(id="c1", name="query_tabular_attachment", arguments={"table_id": "a" * 64})]
|
||||
_, results = ToolExecutor().execute_tool_uses(
|
||||
ctx=ToolExecutionContext(store=store, tools=reg, session_id=sess.id),
|
||||
assistant_msg_id=1,
|
||||
tool_uses=tool_uses,
|
||||
)
|
||||
ok_res, _ = results["c1"]
|
||||
assert bool(ok_res.get("ok"))
|
||||
assert calls["n"] == 1
|
||||
|
||||
|
||||
def test_text_tools_blocked_without_text_ref(tmp_path: Path) -> None:
|
||||
store = SqliteStore(str(tmp_path / "g6.sqlite"))
|
||||
sess = store.create_session("t")
|
||||
store.add_message(session_id=sess.id, role="user", content="summarize", attachments=[{"type": "text", "name": "a.txt"}])
|
||||
|
||||
def _handler(_args):
|
||||
return {"ok": True, "rows": []}
|
||||
|
||||
reg = ToolRegistry(
|
||||
[
|
||||
ToolSpec(
|
||||
name="query_text_attachment",
|
||||
description="query text",
|
||||
parameters={"type": "object", "properties": {"text_id": {"type": "string"}}},
|
||||
handler=_handler,
|
||||
read_only=True,
|
||||
)
|
||||
]
|
||||
)
|
||||
tool_uses = [LLMToolCall(id="c1", name="query_text_attachment", arguments={"text_id": "a" * 64})]
|
||||
_, results = ToolExecutor().execute_tool_uses(
|
||||
ctx=ToolExecutionContext(store=store, tools=reg, session_id=sess.id),
|
||||
assistant_msg_id=1,
|
||||
tool_uses=tool_uses,
|
||||
)
|
||||
blocked, _ = results["c1"]
|
||||
assert not bool(blocked.get("ok"))
|
||||
assert str(blocked.get("error_code") or "") == "text_ref_missing"
|
||||
|
||||
|
||||
def test_text_tools_allowed_with_text_ref(tmp_path: Path) -> None:
|
||||
store = SqliteStore(str(tmp_path / "g7.sqlite"))
|
||||
sess = store.create_session("t")
|
||||
store.add_message(
|
||||
session_id=sess.id,
|
||||
role="user",
|
||||
content="uploaded long text",
|
||||
attachments=[{"type": "text_ref", "text_id": "b" * 64, "name": "long.txt"}],
|
||||
)
|
||||
calls = {"n": 0}
|
||||
|
||||
def _handler(_args):
|
||||
calls["n"] += 1
|
||||
return {"ok": True, "rows": [{"x": 1}]}
|
||||
|
||||
reg = ToolRegistry(
|
||||
[
|
||||
ToolSpec(
|
||||
name="query_text_attachment",
|
||||
description="query text",
|
||||
parameters={"type": "object", "properties": {"text_id": {"type": "string"}}},
|
||||
handler=_handler,
|
||||
read_only=True,
|
||||
)
|
||||
]
|
||||
)
|
||||
tool_uses = [LLMToolCall(id="c1", name="query_text_attachment", arguments={"text_id": "b" * 64})]
|
||||
_, results = ToolExecutor().execute_tool_uses(
|
||||
ctx=ToolExecutionContext(store=store, tools=reg, session_id=sess.id),
|
||||
assistant_msg_id=1,
|
||||
tool_uses=tool_uses,
|
||||
)
|
||||
ok_res, _ = results["c1"]
|
||||
assert bool(ok_res.get("ok"))
|
||||
assert calls["n"] == 1
|
||||
|
||||
|
||||
def test_image_tools_blocked_without_image_ref(tmp_path: Path) -> None:
|
||||
store = SqliteStore(str(tmp_path / "g8.sqlite"))
|
||||
sess = store.create_session("t")
|
||||
store.add_message(session_id=sess.id, role="user", content="analyze image", attachments=[{"type": "text", "name": "x"}])
|
||||
|
||||
def _handler(_args):
|
||||
return {"ok": True}
|
||||
|
||||
reg = ToolRegistry(
|
||||
[
|
||||
ToolSpec(
|
||||
name="query_image_attachment",
|
||||
description="query image",
|
||||
parameters={"type": "object", "properties": {"attachment_id": {"type": "string"}}},
|
||||
handler=_handler,
|
||||
read_only=True,
|
||||
)
|
||||
]
|
||||
)
|
||||
tool_uses = [LLMToolCall(id="c1", name="query_image_attachment", arguments={"attachment_id": "abc"})]
|
||||
_, results = ToolExecutor().execute_tool_uses(
|
||||
ctx=ToolExecutionContext(store=store, tools=reg, session_id=sess.id),
|
||||
assistant_msg_id=1,
|
||||
tool_uses=tool_uses,
|
||||
)
|
||||
blocked, _ = results["c1"]
|
||||
assert not bool(blocked.get("ok"))
|
||||
assert str(blocked.get("error_code") or "") == "image_ref_missing"
|
||||
|
||||
|
||||
def test_image_tools_allowed_with_image_ref(tmp_path: Path) -> None:
|
||||
store = SqliteStore(str(tmp_path / "g9.sqlite"))
|
||||
sess = store.create_session("t")
|
||||
store.add_message(
|
||||
session_id=sess.id,
|
||||
role="user",
|
||||
content="uploaded image",
|
||||
attachments=[{"type": "image_ref", "attachment_id": "abc", "mime": "image/png"}],
|
||||
)
|
||||
calls = {"n": 0}
|
||||
|
||||
def _handler(_args):
|
||||
calls["n"] += 1
|
||||
return {"ok": True}
|
||||
|
||||
reg = ToolRegistry(
|
||||
[
|
||||
ToolSpec(
|
||||
name="query_image_attachment",
|
||||
description="query image",
|
||||
parameters={"type": "object", "properties": {"attachment_id": {"type": "string"}}},
|
||||
handler=_handler,
|
||||
read_only=True,
|
||||
)
|
||||
]
|
||||
)
|
||||
tool_uses = [LLMToolCall(id="c1", name="query_image_attachment", arguments={"attachment_id": "abc"})]
|
||||
_, results = ToolExecutor().execute_tool_uses(
|
||||
ctx=ToolExecutionContext(store=store, tools=reg, session_id=sess.id),
|
||||
assistant_msg_id=1,
|
||||
tool_uses=tool_uses,
|
||||
)
|
||||
ok_res, _ = results["c1"]
|
||||
assert bool(ok_res.get("ok"))
|
||||
assert calls["n"] == 1
|
||||
|
||||
|
||||
def test_video_tools_blocked_without_video_ref(tmp_path: Path) -> None:
|
||||
store = SqliteStore(str(tmp_path / "g10.sqlite"))
|
||||
sess = store.create_session("t")
|
||||
store.add_message(session_id=sess.id, role="user", content="analyze video", attachments=[{"type": "text", "name": "x"}])
|
||||
|
||||
def _handler(_args):
|
||||
return {"ok": True}
|
||||
|
||||
reg = ToolRegistry(
|
||||
[
|
||||
ToolSpec(
|
||||
name="query_video_attachment",
|
||||
description="query video",
|
||||
parameters={"type": "object", "properties": {"attachment_id": {"type": "string"}}},
|
||||
handler=_handler,
|
||||
read_only=True,
|
||||
)
|
||||
]
|
||||
)
|
||||
tool_uses = [LLMToolCall(id="c1", name="query_video_attachment", arguments={"attachment_id": "abc"})]
|
||||
_, results = ToolExecutor().execute_tool_uses(
|
||||
ctx=ToolExecutionContext(store=store, tools=reg, session_id=sess.id),
|
||||
assistant_msg_id=1,
|
||||
tool_uses=tool_uses,
|
||||
)
|
||||
blocked, _ = results["c1"]
|
||||
assert not bool(blocked.get("ok"))
|
||||
assert str(blocked.get("error_code") or "") == "video_ref_missing"
|
||||
|
||||
|
||||
def test_video_tools_allowed_with_video_ref(tmp_path: Path) -> None:
|
||||
store = SqliteStore(str(tmp_path / "g11.sqlite"))
|
||||
sess = store.create_session("t")
|
||||
store.add_message(
|
||||
session_id=sess.id,
|
||||
role="user",
|
||||
content="uploaded video",
|
||||
attachments=[{"type": "video_ref", "attachment_id": "abc", "mime": "video/mp4"}],
|
||||
)
|
||||
calls = {"n": 0}
|
||||
|
||||
def _handler(_args):
|
||||
calls["n"] += 1
|
||||
return {"ok": True}
|
||||
|
||||
reg = ToolRegistry(
|
||||
[
|
||||
ToolSpec(
|
||||
name="query_video_attachment",
|
||||
description="query video",
|
||||
parameters={"type": "object", "properties": {"attachment_id": {"type": "string"}}},
|
||||
handler=_handler,
|
||||
read_only=True,
|
||||
)
|
||||
]
|
||||
)
|
||||
tool_uses = [LLMToolCall(id="c1", name="query_video_attachment", arguments={"attachment_id": "abc"})]
|
||||
_, results = ToolExecutor().execute_tool_uses(
|
||||
ctx=ToolExecutionContext(store=store, tools=reg, session_id=sess.id),
|
||||
assistant_msg_id=1,
|
||||
tool_uses=tool_uses,
|
||||
)
|
||||
ok_res, _ = results["c1"]
|
||||
assert bool(ok_res.get("ok"))
|
||||
assert calls["n"] == 1
|
||||
|
||||
|
|
|
|||
19
tests/test_video_query_tool.py
Normal file
19
tests/test_video_query_tool.py
Normal file
|
|
@ -0,0 +1,19 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import oclaw.runtime.tools.experts.generalist.video_query as vq
|
||||
|
||||
|
||||
def test_normalized_video_meta_extracts_top_level_fields() -> None:
|
||||
obj = {
|
||||
"format": {"duration": "12.50"},
|
||||
"streams": [
|
||||
{"codec_type": "audio", "sample_rate": "48000"},
|
||||
{"codec_type": "video", "width": 1920, "height": 1080, "avg_frame_rate": "30000/1001"},
|
||||
],
|
||||
}
|
||||
out = vq._normalized_video_meta(obj)
|
||||
assert float(out.get("duration_sec") or 0.0) > 12.0
|
||||
assert int(out.get("width") or 0) == 1920
|
||||
assert int(out.get("height") or 0) == 1080
|
||||
assert float(out.get("fps") or 0.0) > 29.0
|
||||
|
||||
Loading…
Add table
Add a link
Reference in a new issue