mirror of
https://github.com/hansjone/oclaw.git
synced 2026-10-09 04:40:45 +08:00
Prefer batch CLI in playbooks, cap tool persistence, and steer getManagedNe failures.
Default scheduled CLI steps to ne_ids/ume_ne_ids batch, compact oversized tool chat_message/tool_log after turns, and on getManagedNe miss guide agents to listManagedNe/UME paths instead of blind retries. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
ae70d4b89a
commit
e2fc72607e
16 changed files with 373 additions and 61 deletions
|
|
@ -10,7 +10,7 @@ from runtime.application.gateway.ops_short_intent import (
|
|||
should_send_group_mention_nudge,
|
||||
)
|
||||
from runtime.tools.base import ToolSpec
|
||||
from runtime.tools.tool_error_hints import enrich_exec_managed_ne_error
|
||||
from runtime.tools.tool_error_hints import enrich_exec_managed_ne_error, enrich_get_managed_ne_error
|
||||
|
||||
|
||||
def test_detect_ops_short_intent_english_field() -> None:
|
||||
|
|
@ -105,3 +105,22 @@ def test_enrich_exec_unreachable_nested_json() -> None:
|
|||
)
|
||||
assert out["error_class"] == "unreachable"
|
||||
assert "unreachable" in out["hint"].lower()
|
||||
|
||||
|
||||
def test_enrich_get_managed_ne_not_found() -> None:
|
||||
out = enrich_get_managed_ne_error(
|
||||
{"ok": False, "error": "netx_http_404", "error_code": "netx_http_404", "detail": "Not Found"}
|
||||
)
|
||||
assert out["error_class"] == "not_found"
|
||||
assert "listManagedNe" in out["hint"]
|
||||
assert "ume_ne_id" in out["hint"]
|
||||
assert "listManagedNe" in " ".join(out.get("next_tools") or [])
|
||||
|
||||
|
||||
def test_enrich_get_managed_ne_id_required() -> None:
|
||||
out = enrich_get_managed_ne_error(
|
||||
{"ok": False, "error": "ne_id_required", "error_code": "ne_id_required"}
|
||||
)
|
||||
assert out["error_class"] == "ne_id_required"
|
||||
assert "listManagedNe" in out["hint"]
|
||||
assert out.get("example")
|
||||
|
|
|
|||
|
|
@ -83,8 +83,38 @@ class RecipeHelpersTests(unittest.TestCase):
|
|||
cong = resolve_ops_recipe_template("congestion")
|
||||
assert cong is not None
|
||||
self.assertEqual((cong.get("source") or {}).get("template_id"), "bandwidth_congestion_daily")
|
||||
cong_blob = " ".join(str(s) for s in (cong.get("steps") or []) + (cong.get("constraints") or []))
|
||||
self.assertIn("ume_ne_ids", cong_blob)
|
||||
self.assertIn("batch", cong_blob.lower())
|
||||
license_tmpl = resolve_ops_recipe_template("license_check")
|
||||
assert license_tmpl is not None
|
||||
lic_blob = " ".join(str(s) for s in (license_tmpl.get("steps") or []))
|
||||
self.assertIn("ne_ids", lic_blob)
|
||||
self.assertIn("never one-NE", lic_blob)
|
||||
self.assertIsNone(resolve_ops_recipe_template("nope"))
|
||||
|
||||
def test_compile_injects_batch_cli_constraint(self) -> None:
|
||||
recipe = {
|
||||
"goal": "CLI check top hosts",
|
||||
"steps": [
|
||||
"Pull congestion alarms",
|
||||
"Run execManagedNe show interface on top hosts",
|
||||
],
|
||||
"success_criteria": ["Group gets summary"],
|
||||
}
|
||||
instr = compile_playbook_instruction(recipe=recipe, lang="en")
|
||||
self.assertIn("ne_ids", instr)
|
||||
self.assertIn("ume_ne_ids", instr)
|
||||
self.assertIn("batch", instr.lower())
|
||||
# Alarm-only playbook should not get CLI batch constraint.
|
||||
alarm_only = {
|
||||
"goal": "Alarm tally",
|
||||
"steps": ["aggregateUmeAlarms", "Summarize by_severity"],
|
||||
"success_criteria": ["Done"],
|
||||
}
|
||||
alarm_instr = compile_playbook_instruction(recipe=alarm_only, lang="en")
|
||||
self.assertNotIn("Multi-NE CLI default", alarm_instr)
|
||||
|
||||
def test_turn_instruction_modes(self) -> None:
|
||||
reminder = build_scheduled_turn_instruction(prompt_text="喝水", mode="scheduled", lang="zh")
|
||||
self.assertIn("提醒意图", reminder)
|
||||
|
|
|
|||
|
|
@ -1,11 +1,16 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import tempfile
|
||||
import unittest
|
||||
import uuid
|
||||
from unittest import mock
|
||||
|
||||
from runtime.chat.tool_runtime import (
|
||||
compact_turn_tool_messages_for_storage,
|
||||
tool_llm_message_max_chars,
|
||||
tool_persist_max_chars,
|
||||
truncate_tool_result_for_llm_messages,
|
||||
)
|
||||
from svc.persistence.sqlite_store import SqliteStore
|
||||
|
|
@ -27,16 +32,25 @@ class ToolLlmTruncationTests(unittest.TestCase):
|
|||
self.assertGreater(out["files_total"], len(out.get("files") or []))
|
||||
|
||||
def test_tool_llm_max_chars_env(self) -> None:
|
||||
import os
|
||||
from unittest import mock
|
||||
|
||||
with mock.patch.dict(os.environ, {"AIA_TOOL_LLM_MESSAGE_MAX_CHARS": "9000"}, clear=False):
|
||||
self.assertEqual(tool_llm_message_max_chars(), 9000)
|
||||
|
||||
def test_compact_turn_tool_messages_for_storage(self) -> None:
|
||||
import tempfile
|
||||
import uuid
|
||||
def test_tool_persist_max_chars_default(self) -> None:
|
||||
env = {
|
||||
k: v
|
||||
for k, v in os.environ.items()
|
||||
if k not in {"AIA_TOOL_PERSIST_MAX_CHARS", "AIA_TOOL_LLM_MESSAGE_MAX_CHARS"}
|
||||
}
|
||||
with mock.patch.dict(os.environ, env, clear=True):
|
||||
self.assertEqual(tool_persist_max_chars(), 24_000)
|
||||
self.assertEqual(tool_llm_message_max_chars(), 0)
|
||||
|
||||
def test_tool_persist_max_chars_disable(self) -> None:
|
||||
with mock.patch.dict(os.environ, {"AIA_TOOL_PERSIST_MAX_CHARS": "0"}, clear=False):
|
||||
self.assertEqual(tool_persist_max_chars(), 0)
|
||||
|
||||
def test_compact_turn_defaults_to_persist_cap(self) -> None:
|
||||
"""Post-turn compact runs with default 24k even when LLM wire cap is unlimited."""
|
||||
db = f"{tempfile.gettempdir()}/oclaw-test-{uuid.uuid4().hex}.sqlite"
|
||||
store = SqliteStore(db)
|
||||
sess = store.create_session("t")
|
||||
|
|
@ -56,22 +70,63 @@ class ToolLlmTruncationTests(unittest.TestCase):
|
|||
tool_calls={"tool_call_id": "c1", "name": "echo", "assistant_message_id": 1},
|
||||
turn_uuid=turn_uuid,
|
||||
)
|
||||
before = store.get_messages(session_id=sess.id, limit=20)
|
||||
before_tool = [m for m in before if m.id == row.id][0]
|
||||
self.assertNotIn("_truncated_for_llm", str(before_tool.content or ""))
|
||||
from unittest import mock
|
||||
|
||||
with mock.patch.dict("os.environ", {"AIA_TOOL_LLM_MESSAGE_MAX_CHARS": "8000"}, clear=False):
|
||||
env = {
|
||||
k: v
|
||||
for k, v in os.environ.items()
|
||||
if k not in {"AIA_TOOL_PERSIST_MAX_CHARS", "AIA_TOOL_LLM_MESSAGE_MAX_CHARS"}
|
||||
}
|
||||
with mock.patch.dict(os.environ, env, clear=True):
|
||||
stats = compact_turn_tool_messages_for_storage(
|
||||
store=store,
|
||||
session_id=sess.id,
|
||||
turn_uuid=turn_uuid,
|
||||
)
|
||||
self.assertEqual(int(stats.get("persist_cap") or 0), 24_000)
|
||||
self.assertGreaterEqual(int(stats.get("scanned") or 0), 1)
|
||||
self.assertGreaterEqual(int(stats.get("updated") or 0), 1)
|
||||
after = store.get_messages(session_id=sess.id, limit=20)
|
||||
after_tool = [m for m in after if m.id == row.id][0]
|
||||
self.assertIn("_truncated_for_llm", str(after_tool.content or ""))
|
||||
self.assertLess(len(str(after_tool.content or "")), 40_000)
|
||||
|
||||
def test_compact_turn_can_be_disabled(self) -> None:
|
||||
db = f"{tempfile.gettempdir()}/oclaw-test-{uuid.uuid4().hex}.sqlite"
|
||||
store = SqliteStore(db)
|
||||
sess = store.create_session("t")
|
||||
turn_uuid = "turn-off"
|
||||
store.add_message(
|
||||
session_id=sess.id,
|
||||
role="tool",
|
||||
content=json.dumps({"ok": True, "blob": "x" * 50_000}, ensure_ascii=False),
|
||||
tool_calls={"tool_call_id": "c1", "name": "echo"},
|
||||
turn_uuid=turn_uuid,
|
||||
)
|
||||
with mock.patch.dict(os.environ, {"AIA_TOOL_PERSIST_MAX_CHARS": "0"}, clear=False):
|
||||
stats = compact_turn_tool_messages_for_storage(
|
||||
store=store,
|
||||
session_id=sess.id,
|
||||
turn_uuid=turn_uuid,
|
||||
)
|
||||
self.assertEqual(int(stats.get("skipped") or 0), 1)
|
||||
self.assertEqual(int(stats.get("updated") or 0), 0)
|
||||
|
||||
def test_tool_log_default_cap(self) -> None:
|
||||
db = f"{tempfile.gettempdir()}/oclaw-test-{uuid.uuid4().hex}.sqlite"
|
||||
store = SqliteStore(db)
|
||||
sess = store.create_session("t")
|
||||
env = {k: v for k, v in os.environ.items() if k != "AIA_TOOL_LOG_MAX_CHARS"}
|
||||
with mock.patch.dict(os.environ, env, clear=True):
|
||||
store.add_tool_log(
|
||||
session_id=sess.id,
|
||||
tool_name="echo",
|
||||
args={},
|
||||
result={"ok": True, "blob": "y" * 200_000},
|
||||
)
|
||||
logs = store.get_tool_logs(sess.id, limit=5)
|
||||
self.assertEqual(len(logs), 1)
|
||||
blob = json.dumps(logs[0]["result"], ensure_ascii=False)
|
||||
self.assertLessEqual(len(blob), 70_000)
|
||||
self.assertTrue(logs[0]["result"].get("ok") is True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue