diff --git a/.gitignore b/.gitignore index a17f8428..5f8d5a4e 100644 --- a/.gitignore +++ b/.gitignore @@ -37,6 +37,7 @@ _local/*.env _local/*.json _local/*.txt data/channel_sidecar/ +data/wiki/ data/**/node_modules/ data/**/*.log data/**/*.err.log diff --git a/runtime/agent_core_attempt.py b/runtime/agent_core_attempt.py index 29355e8b..b9feb771 100644 --- a/runtime/agent_core_attempt.py +++ b/runtime/agent_core_attempt.py @@ -15,6 +15,13 @@ def _workspace_owner_session_id_from_msg(msg: StandardMessage) -> str | None: return w or None +def _should_after_turn_memory(msg: StandardMessage) -> bool: + # Wiki memory must be agent-initiated only. + # Disable passive/automatic capture path in agent core. + _ = msg + return False + + ALL_ATTEMPT_ERROR_CODES = ( "relay_envelope_invalid", "relay_envelope_unsupported_version", @@ -56,6 +63,7 @@ class AttemptRunnerInput: skill_binding_role: str | None = None wire_policy_role: str | None = None prompt_build_context: dict[str, Any] | None = None + turn_uuid: str | None = None @dataclass(frozen=True) @@ -125,33 +133,35 @@ def run_attempt(*, store: Any, data: AttemptRunnerInput) -> AttemptRunnerOutput: skill_binding_role=data.skill_binding_role, wire_policy_role=data.wire_policy_role, prompt_build_context=data.prompt_build_context, + turn_uuid=data.turn_uuid, ) - after_turn_memory( - store=store, - session_id=data.msg.session_id, - tenant_id=data.msg.tenant_id, - user_id=data.msg.user_id, - user_text=data.msg.text, - assistant_text=outcome.final_text, - turn_uuid=outcome.turn_uuid, - ) - if data.trace_id: - try: - store.add_trace_event( - session_id=data.msg.session_id, - trace_id=str(data.trace_id), - span_id=new_span_id(), - parent_span_id=data.parent_span_id, - event_type="after_turn_memory", - payload={ - "pipeline": "oclaw_agent_core", - "oc_stage": "memory_done", - "run_id": str(data.run_id or ""), - "attempt_no": int(data.attempt_no), - }, - ) - except Exception: - pass + if _should_after_turn_memory(data.msg): + after_turn_memory( + store=store, + session_id=data.msg.session_id, + tenant_id=data.msg.tenant_id, + user_id=data.msg.user_id, + user_text=data.msg.text, + assistant_text=outcome.final_text, + turn_uuid=outcome.turn_uuid, + ) + if data.trace_id: + try: + store.add_trace_event( + session_id=data.msg.session_id, + trace_id=str(data.trace_id), + span_id=new_span_id(), + parent_span_id=data.parent_span_id, + event_type="after_turn_memory", + payload={ + "pipeline": "oclaw_agent_core", + "oc_stage": "memory_done", + "run_id": str(data.run_id or ""), + "attempt_no": int(data.attempt_no), + }, + ) + except Exception: + pass st = AttemptState( attempt_no=int(data.attempt_no), status="success", diff --git a/runtime/agent_core_run.py b/runtime/agent_core_run.py index 0bfa9f7a..6e92be3d 100644 --- a/runtime/agent_core_run.py +++ b/runtime/agent_core_run.py @@ -152,6 +152,7 @@ def _relay_envelope_stats(msg: StandardMessage) -> dict[str, Any]: def run_agent_core(*, store: Any, data: AgentCoreRunInput) -> AgentCoreRunOutput: run_id = str(data.run_id or "").strip() or str(uuid.uuid4()) + user_turn_uuid = str(uuid.uuid4()) attempts: list[AttemptState] = [] compact_count = 0 outcome = TurnRunOutcome(final_text="", tool_traces=tuple(), handoff_note="", turn_uuid="") @@ -215,6 +216,7 @@ def run_agent_core(*, store: Any, data: AgentCoreRunInput) -> AgentCoreRunOutput skill_binding_role=data.skill_binding_role, wire_policy_role=data.wire_policy_role, prompt_build_context=(dict(data.msg.metadata or {}) if isinstance(data.msg.metadata, dict) else None), + turn_uuid=user_turn_uuid, ), ) attempts.append(out.state) diff --git a/runtime/skill_executor.py b/runtime/skill_executor.py index a22ec742..79eac744 100644 --- a/runtime/skill_executor.py +++ b/runtime/skill_executor.py @@ -167,6 +167,7 @@ class SkillExecutor: workspace_owner_session_id=ctx.workspace_owner_session_id, path_policy_tenant_id=ctx.path_policy_tenant_id, path_policy_user_id=ctx.path_policy_user_id, + workspace_dir=ctx.workspace_dir, turn_uuid=ctx.turn_uuid, ), assistant_msg_id=assistant_msg_id, diff --git a/runtime/tools/experts/workspace/fs_tools.py b/runtime/tools/experts/workspace/fs_tools.py index 94a16ae6..08f2edfb 100644 --- a/runtime/tools/experts/workspace/fs_tools.py +++ b/runtime/tools/experts/workspace/fs_tools.py @@ -5,7 +5,9 @@ from pathlib import Path from typing import Any from oclaw.runtime.tools.base import ToolSpec -from oclaw.runtime.tools.experts.workspace.workspace_base import resolve_workspace_path +from oclaw.runtime.tools.experts.workspace.workspace_base import ( + resolve_workspace_path, +) def read_file_tool() -> ToolSpec: @@ -60,20 +62,28 @@ def read_file_tool() -> ToolSpec: def write_file_tool() -> ToolSpec: + def _sandbox_base_dir() -> Path: + return Path("data") / "workspace" + def _normalize_write_path(path: str) -> str: raw = str(path or "").strip().strip('"').strip("'") if not raw: raise ValueError("path_required") p = Path(raw) + base = _sandbox_base_dir() + # Enforce sandbox for absolute paths as well. if p.is_absolute(): - return str(p) - # Keep generated files out of repo root: default relative writes go under data//... - ws = resolve_workspace_path(".") - ws_name = ws.name or "workspace" + # Collapse absolute user path into sandbox-relative target to prevent + # writes to repo root or arbitrary host locations. + name = str(p.name or "").strip() + if not name: + raise ValueError("path_required") + return str(base / name) + # Keep generated files out of repo root: default relative writes go under data/workspace/... rel = raw.lstrip("./\\") if not rel: raise ValueError("path_required") - return str(Path("data") / ws_name / rel) + return str(base / rel) def handler(args: dict[str, Any]) -> dict[str, Any]: path = str(args.get("path") or "").strip() diff --git a/runtime/tools/experts/workspace/shell_tools.py b/runtime/tools/experts/workspace/shell_tools.py index 5b037164..9bb071e1 100644 --- a/runtime/tools/experts/workspace/shell_tools.py +++ b/runtime/tools/experts/workspace/shell_tools.py @@ -1,15 +1,26 @@ from __future__ import annotations import subprocess +import re +from pathlib import Path from typing import Any from oclaw.platform.config.paths import db_path from oclaw.platform.persistence.sqlite_store import SqliteStore from oclaw.runtime.tools.base import ToolSpec -from oclaw.runtime.tools.experts.workspace.workspace_base import resolve_workspace_path, truncate_text +from oclaw.runtime.tools.experts.workspace.workspace_base import ( + resolve_workspace_path, + truncate_text, + workspace_root, +) def run_command_tool() -> ToolSpec: + _LEADING_CD_CHAIN_RE = re.compile( + r"^\s*(?:(?:[A-Za-z]:)\s*&&\s*)?(?:@echo\s+off\s*&&\s*)?cd\s+(?:/d\s+)?(?:\"[^\"]+\"|[^&]+?)\s*&&\s*(.+)$", + re.IGNORECASE | re.DOTALL, + ) + def _run_command_enabled() -> bool: import os @@ -33,6 +44,84 @@ def run_command_tool() -> ToolSpec: def handler(args: dict[str, Any]) -> dict[str, Any]: import os + def _default_exec_dir() -> str: + return str(resolve_workspace_path("data/workspace")) + + def _strip_leading_cd_chain(cmd: str) -> tuple[str, bool]: + raw = str(cmd or "") + changed = False + out = raw + # Strip repeated leading "cd ... &&" so default sandbox cwd cannot be bypassed by habit. + for _ in range(3): + m = _LEADING_CD_CHAIN_RE.match(out) + if not m: + break + tail = str(m.group(1) or "").strip() + if not tail: + break + out = tail + changed = True + return out, changed + + def _rewrite_workspace_absolute_refs(cmd: str, *, workdir: str) -> tuple[str, bool]: + raw = str(cmd or "") + root = str(workspace_root()) + if not raw or not root: + return raw, False + root_norm = root.rstrip("\\/") + changed = False + out = raw + marker = root_norm + "\\" + if marker.lower() not in out.lower(): + return out, False + idx = 0 + rebuilt = [] + low = out.lower() + marker_low = marker.lower() + while True: + pos = low.find(marker_low, idx) + if pos < 0: + rebuilt.append(out[idx:]) + break + rebuilt.append(out[idx:pos]) + tail_start = pos + len(marker) + tail_end = tail_start + while tail_end < len(out) and out[tail_end] not in ('"', "'", " ", "\t", "\r", "\n"): + tail_end += 1 + rel_tail = out[tail_start:tail_end] + candidate = str(Path(workdir) / rel_tail) + if Path(candidate).exists(): + rebuilt.append(candidate) + changed = True + else: + rebuilt.append(out[pos:tail_end]) + idx = tail_end + return "".join(rebuilt), changed + + def _rewrite_python_script_arg(cmd: str, *, workdir: str) -> tuple[str, bool]: + raw = str(cmd or "").strip() + if not raw: + return raw, False + m = re.match(r'^\s*(python|py)\s+("([^"]+\.py)"|([^\s]+\.py))(\s+.*)?$', raw, flags=re.IGNORECASE) + if not m: + return raw, False + script = str(m.group(3) or m.group(4) or "").strip() + if not script: + return raw, False + # Absolute path is handled by workspace-absolute rewrite already. + sp = Path(script) + if sp.is_absolute(): + return raw, False + base = str(Path(script).name or "").strip() + if not base: + return raw, False + # Deterministic policy: always bind python script arg to sandbox root. + rel = base + quote = '"' if " " in rel else "" + prefix = str(m.group(1) or "python") + rest = str(m.group(5) or "") + return f"{prefix} {quote}{rel}{quote}{rest}", True + if not _run_command_enabled(): return { "ok": False, @@ -45,7 +134,21 @@ def run_command_tool() -> ToolSpec: max_output_chars = int(args.get("max_output_chars") or 20000) if not command: return {"ok": False, "error": "command_required"} - workdir = resolve_workspace_path(cwd or ".") + normalized_cd_removed = False + command_rewritten = False + cwd_redirected_to_sandbox = False + script_path_rewritten = False + original_command = command + # Hard policy: always execute in sandbox root, ignore caller-supplied cwd. + workdir = _default_exec_dir() + cwd_redirected_to_sandbox = bool(str(cwd or "").strip() and str(cwd).strip() not in {".", "./", ".\\"}) + command, normalized_cd_removed = _strip_leading_cd_chain(command) + command, command_rewritten = _rewrite_workspace_absolute_refs(command, workdir=workdir) + command, script_path_rewritten = _rewrite_python_script_arg(command, workdir=workdir) + try: + os.makedirs(workdir, exist_ok=True) + except Exception: + pass try: run_kwargs: dict[str, Any] = { "cwd": str(workdir), @@ -64,14 +167,39 @@ def run_command_tool() -> ToolSpec: command, **run_kwargs, ) - out = (cp.stdout or "") + (("\n" + cp.stderr) if cp.stderr else "") - out = truncate_text(out, limit=max(1000, min(max_output_chars, 200000))) + stdout_text = str(cp.stdout or "") + stderr_text = str(cp.stderr or "") + out_raw = stdout_text + (("\n" + stderr_text) if stderr_text else "") + out_limit = max(1000, min(max_output_chars, 200000)) + out = truncate_text(out_raw, limit=out_limit) + out_truncated = len(out_raw) > out_limit + out_empty = (len(str(out_raw or "").strip()) == 0) + exit_code = int(cp.returncode) + ok_flag = exit_code == 0 return { - "ok": True, + "ok": bool(ok_flag), "command": command, "cwd": str(workdir), - "exit_code": int(cp.returncode), + "exit_code": exit_code, + "stdout": stdout_text, + "stderr": stderr_text, "output": out, + "output_chars": int(len(out_raw)), + "output_empty": bool(out_empty), + "output_truncated": bool(out_truncated), + "output_not_truncated": bool(not out_truncated), + "normalized_cd_removed": bool(normalized_cd_removed), + "cwd_redirected_to_sandbox": bool(cwd_redirected_to_sandbox), + "command_rewritten": bool(command_rewritten), + "script_path_rewritten": bool(script_path_rewritten), + "original_command": original_command, + "error_code": ("" if ok_flag else "command_exit_nonzero"), + "output_hint": ( + "Command produced empty stdout/stderr; this is not system truncation. " + "Do not claim truncation unless output_truncated=true." + if out_empty + else "" + ), } except subprocess.TimeoutExpired as e: partial = "" @@ -86,9 +214,26 @@ def run_command_tool() -> ToolSpec: "cwd": str(workdir), "timeout_s": timeout_s, "output": truncate_text(partial, limit=max_output_chars), + "output_empty": len(str(partial or "").strip()) == 0, + "output_truncated": len(str(partial or "")) > int(max_output_chars or 0), + "normalized_cd_removed": bool(normalized_cd_removed), + "cwd_redirected_to_sandbox": bool(cwd_redirected_to_sandbox), + "command_rewritten": bool(command_rewritten), + "script_path_rewritten": bool(script_path_rewritten), + "original_command": original_command, } except Exception as e: - return {"ok": False, "error": f"{type(e).__name__}: {e}", "command": command, "cwd": str(workdir)} + return { + "ok": False, + "error": f"{type(e).__name__}: {e}", + "command": command, + "cwd": str(workdir), + "normalized_cd_removed": bool(normalized_cd_removed), + "cwd_redirected_to_sandbox": bool(cwd_redirected_to_sandbox), + "command_rewritten": bool(command_rewritten), + "script_path_rewritten": bool(script_path_rewritten), + "original_command": original_command, + } return ToolSpec( name="run_command", diff --git a/runtime/tools/experts/workspace/workspace_base.py b/runtime/tools/experts/workspace/workspace_base.py index 267efd83..48f45c9a 100644 --- a/runtime/tools/experts/workspace/workspace_base.py +++ b/runtime/tools/experts/workspace/workspace_base.py @@ -179,6 +179,29 @@ def current_workspace_path_access() -> WorkspacePathAccess: return access_from_env() +@contextmanager +def workspace_write_namespace_scope(namespace: str | None) -> Iterator[str]: + prev = getattr(_TLS, "write_namespace", None) + ns = str(namespace or "").strip() + _TLS.write_namespace = ns + try: + yield ns + finally: + if prev is None: + if hasattr(_TLS, "write_namespace"): + delattr(_TLS, "write_namespace") + else: + _TLS.write_namespace = prev + + +def current_workspace_write_namespace() -> str: + ns = str(getattr(_TLS, "write_namespace", "") or "").strip() + if ns: + return ns + root = workspace_root() + return str(root.name or "workspace").strip() or "workspace" + + def clear_workspace_path_access_for_tests() -> None: if hasattr(_TLS, "access"): delattr(_TLS, "access") @@ -253,9 +276,11 @@ __all__ = [ "build_workspace_path_access", "clear_workspace_path_access_for_tests", "current_workspace_path_access", + "current_workspace_write_namespace", "resolve_workspace_path", "sanitize_git_ref", "truncate_text", + "workspace_write_namespace_scope", "workspace_path_access_scope", "workspace_root", ] diff --git a/runtime/worker.py b/runtime/worker.py index 2d1515fc..3e060267 100644 --- a/runtime/worker.py +++ b/runtime/worker.py @@ -14,6 +14,128 @@ from oclaw.runtime.types import StandardMessage _LOCK = threading.Lock() _THREAD: threading.Thread | None = None +_SESSION_TITLE_MAX_LEN = 120 +_AUTO_TITLE_STAGE_KEY_PREFIX = "AIA_SESSION_AUTO_TITLE_STAGE:" +_TITLE_TRIGGER_ROUND = 3 +_TITLE_BODIES_MAX_CHARS = 4000 + + +def _maybe_rename_from_first_user_message(*, store: Any, session_id: str, user_text: str, attachments: list[dict[str, Any]] | None) -> None: + sid = str(session_id or "").strip() + if not sid: + return + try: + stage_raw = str(store.get_setting(f"{_AUTO_TITLE_STAGE_KEY_PREFIX}{sid}") or "").strip() + except Exception: + stage_raw = "" + if stage_raw in ("1", "3"): + return + try: + sess = store.get_session(sid) + except Exception: + sess = None + if not sess: + return + cur_title = str(getattr(sess, "title", "") or "").strip() + if cur_title not in ("新会话", "New Chat"): + return + try: + rows = store.get_messages(session_id=sid, limit=20) + except Exception: + rows = [] + user_count = 0 + for r in rows or []: + if str(getattr(r, "role", "") or "").strip().lower() == "user": + user_count += 1 + if user_count > 1: + return + title = str(user_text or "").strip().replace("\n", " ") + if not title: + atts = attachments if isinstance(attachments, list) else [] + if atts and isinstance(atts[0], dict): + title = str(atts[0].get("name") or "").strip() + if not title: + return + try: + store.rename_session(sid, title[:_SESSION_TITLE_MAX_LEN]) + try: + store.set_setting(f"{_AUTO_TITLE_STAGE_KEY_PREFIX}{sid}", "1") + except Exception: + pass + except Exception: + pass + + +def _maybe_generate_title_on_third_round(*, store: Any, msg: StandardMessage, model: Any | None) -> None: + """Generate title once on round-3 using user text only (no tools/reasoning context).""" + if model is None or not callable(getattr(model, "chat", None)): + return + sid = str(msg.session_id or "").strip() + if not sid: + return + try: + stage_raw = str(store.get_setting(f"{_AUTO_TITLE_STAGE_KEY_PREFIX}{sid}") or "").strip() + except Exception: + stage_raw = "" + try: + sess = store.get_session(sid) + except Exception: + sess = None + if not sess: + return + cur_title = str(getattr(sess, "title", "") or "").strip() + # Two-stage naming: + # - stage "1": renamed from first user message + # - stage "3": renamed on third user message (final) + if stage_raw == "3": + return + if (cur_title not in ("新会话", "New Chat")) and (stage_raw != "1"): + return + try: + rows = store.get_messages(session_id=sid, limit=200) + except Exception: + rows = [] + bodies: list[str] = [] + for r in rows or []: + role = str(getattr(r, "role", "") or "").strip().lower() + if role != "user": + continue + txt = str(getattr(r, "content", "") or "").strip() + if txt: + bodies.append(txt) + cur_txt = str(msg.text or "").strip() + if cur_txt: + bodies.append(cur_txt) + if len(bodies) != _TITLE_TRIGGER_ROUND: + return + body = "\n".join(f"{i+1}. {t}" for i, t in enumerate(bodies)) + body = body[:_TITLE_BODIES_MAX_CHARS] + try: + lang_is_en = str(msg.metadata.get("lang") if isinstance(msg.metadata, dict) else "").lower().startswith("en") + sys = ( + "Generate a concise chat title from these user messages only. " + "Use the dominant language used by the user content body. " + "Return title text only, no quotes, no markdown, max 18 chars." + if lang_is_en + else "仅基于以下用户正文生成简短会话标题。请使用对话内容主体语言命名。" + "只返回标题文本,不要引号,不要markdown,最多18个字。" + ) + resp = model.chat( + [{"role": "system", "content": sys}, {"role": "user", "content": body}], + [], + on_token=None, + ) + title = str(getattr(resp, "content", "") or "").strip().replace("\n", " ") + title = title.strip("\"'` ").strip() + if not title: + return + store.rename_session(sid, title[:_SESSION_TITLE_MAX_LEN]) + try: + store.set_setting(f"{_AUTO_TITLE_STAGE_KEY_PREFIX}{sid}", "3") + except Exception: + pass + except Exception: + return def ensure_worker_started(*, store: Any, poll_interval_s: float = 1.0) -> str: @@ -157,6 +279,17 @@ def _worker_loop(*, store: Any, worker_id: str, poll_interval_s: float) -> None: attachments=list(attachments or []), metadata=dict(metadata or {}), ) + _maybe_rename_from_first_user_message( + store=store, + session_id=session_id, + user_text=user_text, + attachments=list(attachments or []), + ) + _maybe_generate_title_on_third_round( + store=store, + msg=msg, + model=getattr(executor, "model", None), + ) run_out = run_agent_core( store=store, data=AgentCoreRunInput( diff --git a/runtime/workspaces/main/ROLE_SYSTEM.md b/runtime/workspaces/main/ROLE_SYSTEM.md index 3355df09..565a7065 100644 --- a/runtime/workspaces/main/ROLE_SYSTEM.md +++ b/runtime/workspaces/main/ROLE_SYSTEM.md @@ -10,6 +10,7 @@ ## 下发规则(何时调用专家) - 每轮都必须选择并下发一个专家(固定或动态)。 +- 必须结合上下文分析用户意图,将全量已知信息交给专家进行处理,不能仅转述用户当前的问题。 - 简单任务也要下发 `generalist`,不要由主控直接产出最终答案。 - 唯一例外:`route.kind="manager_memory"`,用于主控直接执行“记忆写入”动作(不是通用任务直出)。 @@ -35,6 +36,7 @@ - 仅当 `route.kind="manager_memory"` 时,记忆写入与对话回复可同轮并行:写入使用 `dispatch.memory_write_text`,对话回复使用 `dispatch.instruction_text`。 - 记忆写入不得改变本轮对话输出语义;回复内容以用户问题与业务目标为准。 - 若提供 `dispatch.post_reply_memory_write_text`,其语义是“回程补写记忆”,与用户可见回复解耦。 +- 记忆写入必须由主控主动显式触发(`route.kind="manager_memory"` + `dispatch.memory_write_text` 或 `dispatch.post_reply_memory_write_text`);禁止依赖任何被动/自动兜底写入机制。 ## 决策解释(为何写入 / 为何注入) - 为何写入记忆:把“本轮产生且未来可复用”的稳定结论沉淀到 wiki,减少后续重复澄清与重复决策。 diff --git a/tests/test_admin_chat_stream_async_task.py b/tests/test_admin_chat_stream_async_task.py index c675c081..01d78cfd 100644 --- a/tests/test_admin_chat_stream_async_task.py +++ b/tests/test_admin_chat_stream_async_task.py @@ -127,6 +127,28 @@ class AdminChatStreamAsyncTaskTests(unittest.TestCase): self.assertEqual(str(body.get("mode") or ""), "async_task") self.assertTrue(str(body.get("task_id") or "").strip()) + def test_async_task_first_turn_renames_new_chat_title(self) -> None: + token = self._login() + headers = { + "authorization": f"Bearer {token}", + "accept": "application/json", + "content-type": "application/json", + } + created = self.client.post("/admin/api/chat/sessions", headers=headers, json={}).json() + sid = str(((created.get("session") or {}).get("id")) or "") + self.assertTrue(sid) + send = self.client.post( + f"/admin/api/chat/sessions/{sid}/messages", + headers=headers, + json={"text": "请总结并发送到项目群"}, + ) + self.assertEqual(send.status_code, 200) + self.assertTrue(send.json().get("ok"), send.json()) + sess = self.store.get_session(sid) + self.assertIsNotNone(sess) + title = str(getattr(sess, "title", "") or "") + self.assertNotIn(title, {"新会话", "New Chat"}) + def test_async_task_payload_contains_selected_specialist(self) -> None: token = self._login() headers = { @@ -209,6 +231,71 @@ class AdminChatStreamAsyncTaskTests(unittest.TestCase): self.assertEqual(str(payload.get("requested_specialist") or ""), "ops") self.assertEqual(str(payload.get("memory_mode") or ""), "store_only") + def test_new_session_inherits_user_mode_preference(self) -> None: + token = self._login() + headers = { + "authorization": f"Bearer {token}", + "accept": "application/json", + "content-type": "application/json", + } + set_resp = self.client.post( + f"/admin/api/chat/sessions/{self.session_id}/mode", + headers=headers, + json={"interaction_mode": "expert", "specialist": "ops", "memory_mode": "store_only"}, + ) + self.assertEqual(set_resp.status_code, 200) + + create_resp = self.client.post( + "/admin/api/chat/sessions", + headers=headers, + json={"title": "new one"}, + ) + self.assertEqual(create_resp.status_code, 200) + created = create_resp.json() + self.assertTrue(created.get("ok"), created) + new_session_id = str(((created.get("session") or {}).get("id")) or "") + self.assertTrue(new_session_id) + + mode_resp = self.client.get( + f"/admin/api/chat/sessions/{new_session_id}/mode", + headers={"authorization": f"Bearer {token}", "accept": "application/json"}, + ) + self.assertEqual(mode_resp.status_code, 200) + mode_body = mode_resp.json() + self.assertTrue(mode_body.get("ok"), mode_body) + self.assertEqual(str(mode_body.get("interaction_mode") or ""), "expert") + self.assertEqual(str(mode_body.get("specialist") or ""), "ops") + self.assertEqual(str(mode_body.get("memory_mode") or ""), "store_only") + + def test_new_session_default_mode_is_expert_generalist(self) -> None: + token = self._login() + headers = { + "authorization": f"Bearer {token}", + "accept": "application/json", + "content-type": "application/json", + } + create_resp = self.client.post( + "/admin/api/chat/sessions", + headers=headers, + json={"title": "default mode chat"}, + ) + self.assertEqual(create_resp.status_code, 200) + created = create_resp.json() + self.assertTrue(created.get("ok"), created) + new_session_id = str(((created.get("session") or {}).get("id")) or "") + self.assertTrue(new_session_id) + + mode_resp = self.client.get( + f"/admin/api/chat/sessions/{new_session_id}/mode", + headers={"authorization": f"Bearer {token}", "accept": "application/json"}, + ) + self.assertEqual(mode_resp.status_code, 200) + mode_body = mode_resp.json() + self.assertTrue(mode_body.get("ok"), mode_body) + self.assertEqual(str(mode_body.get("interaction_mode") or ""), "expert") + self.assertEqual(str(mode_body.get("specialist") or ""), "generalist") + self.assertEqual(str(mode_body.get("memory_mode") or ""), "default") + def test_admin_dynamic_expert_stats_endpoint(self) -> None: token = self._login() headers = { @@ -322,6 +409,13 @@ class AdminChatStreamAsyncTaskTests(unittest.TestCase): self.assertTrue(get_body.get("ok"), get_body) self.assertEqual(int((get_body.get("limits") or {}).get("max_rows_read") or 0), 5000) self.assertEqual(int((get_body.get("limits") or {}).get("sql_timeout_ms") or 0), 8000) + # newly added replay caps should exist with sane defaults + self.assertGreaterEqual(int((get_body.get("limits") or {}).get("image_result_replay_cap_chars") or 0), 600) + self.assertGreaterEqual(int((get_body.get("limits") or {}).get("video_result_replay_cap_chars") or 0), 600) + self.assertGreaterEqual(int((get_body.get("limits") or {}).get("video_transcript_chunk_size") or 0), 1) + self.assertGreaterEqual(int((get_body.get("limits") or {}).get("video_transcript_chunk_overlap") or 0), 0) + self.assertGreaterEqual(int((get_body.get("limits") or {}).get("archive_max_depth") or 0), 1) + self.assertGreaterEqual(int((get_body.get("limits") or {}).get("archive_max_file_count") or 0), 1) set_resp = self.client.post( "/admin/api/chat/settings/attachment-limits", @@ -334,6 +428,12 @@ class AdminChatStreamAsyncTaskTests(unittest.TestCase): "max_excel_sheets": 8, "large_table_preview_rows": 33, "sql_timeout_ms": 1234, + "video_transcript_chunk_size": 1700, + "video_transcript_chunk_overlap": 240, + "archive_max_depth": 3, + "archive_max_file_count": 555, + "archive_max_entry_bytes": 123456, + "archive_max_total_uncompressed_bytes": 654321, } }, ) @@ -347,6 +447,12 @@ class AdminChatStreamAsyncTaskTests(unittest.TestCase): self.assertEqual(int(limits.get("max_excel_sheets") or 0), 8) self.assertEqual(int(limits.get("large_table_preview_rows") or 0), 33) self.assertEqual(int(limits.get("sql_timeout_ms") or 0), 1234) + self.assertEqual(int(limits.get("video_transcript_chunk_size") or 0), 1700) + self.assertEqual(int(limits.get("video_transcript_chunk_overlap") or 0), 240) + self.assertEqual(int(limits.get("archive_max_depth") or 0), 3) + self.assertEqual(int(limits.get("archive_max_file_count") or 0), 555) + self.assertEqual(int(limits.get("archive_max_entry_bytes") or 0), 123456) + self.assertEqual(int(limits.get("archive_max_total_uncompressed_bytes") or 0), 654321) if __name__ == "__main__": diff --git a/tests/test_agent_core_attempt_memory_gate.py b/tests/test_agent_core_attempt_memory_gate.py new file mode 100644 index 00000000..823e2602 --- /dev/null +++ b/tests/test_agent_core_attempt_memory_gate.py @@ -0,0 +1,33 @@ +from __future__ import annotations + +from oclaw.runtime.agent_core_attempt import _should_after_turn_memory +from oclaw.runtime.types import StandardMessage + + +def _msg(interaction_mode: str, extra: dict | None = None) -> StandardMessage: + md = {"interaction_mode": interaction_mode} + if isinstance(extra, dict): + md.update(extra) + return StandardMessage( + session_id="s1", + tenant_id="t1", + user_id="u1", + role="member", + channel="admin_chat", + text="hello", + attachments=[], + metadata=md, + ) + + +def test_after_turn_memory_disabled_for_comprehensive_mode() -> None: + assert _should_after_turn_memory(_msg("comprehensive")) is False + + +def test_after_turn_memory_disabled_for_expert_mode_by_default() -> None: + assert _should_after_turn_memory(_msg("expert")) is False + + +def test_after_turn_memory_cannot_be_forced_by_metadata() -> None: + assert _should_after_turn_memory(_msg("expert", {"enable_after_turn_memory": True})) is False + diff --git a/tests/test_oclaw_gateway_trace.py b/tests/test_oclaw_gateway_trace.py index 166b0fd5..b2707501 100644 --- a/tests/test_oclaw_gateway_trace.py +++ b/tests/test_oclaw_gateway_trace.py @@ -275,6 +275,65 @@ def test_gateway_comprehensive_mode_manager_first_selects_specialist(monkeypatch assert chosen.get("sid") == "image" assert captured.get("exec_text") == "Please edit the image background." + +def test_gateway_comprehensive_mode_writes_task_assignment_reasoning(monkeypatch: pytest.MonkeyPatch) -> None: + written: list[dict] = [] + + class Store: + def get_setting(self, _k: str) -> str: + return "" + + def add_trace_event(self, **_kwargs: object) -> None: + return None + + def add_message(self, **kwargs: object) -> None: + written.append(dict(kwargs)) + + class _ManagerModel: + def chat(self, _messages, _tools, *, on_token=None): + return LLMResponse( + content='{"route":{"specialist":"ops","reason":"ops task"},"dispatch":{"instruction_text":"请检查并修复网关启动失败。"}}', + tool_calls=[], + ) + + class _Exec: + def __init__(self, model=None): + self.model = model + self.tools = object() + self.system_prompt = "" + + monkeypatch.setattr( + "oclaw.runtime.gateway.get_manager_prompt_prebuild", + lambda **kwargs: { + "manager_context": "manager", + "allowed_fixed": ("generalist", "ops", "image", "memory"), + "allowed_fixed_quoted": '"generalist", "ops", "image", "memory"', + }, + ) + monkeypatch.setattr( + "oclaw.runtime.gateway.run_agent_core", + lambda **kwargs: SimpleNamespace(outcome=SimpleNamespace(final_text="specialist_answer")), + ) + + gw = OclawGateway(store=Store()) + msg = StandardMessage( + session_id="sid-assign-1", + tenant_id="t1", + user_id="u1", + role="user", + channel="admin_chat", + text="网关起不来,帮我修复", + attachments=[], + metadata={"interaction_mode": "comprehensive"}, + ) + _ = gw.handle_turn(msg=msg, lang="zh", executor=_Exec(model=_ManagerModel())) + reasoning_rows = [x for x in written if str(x.get("event_type") or "") == "reasoning"] + assert reasoning_rows, "expected task-assignment reasoning row" + content = str(reasoning_rows[-1].get("content") or "") + assert "任务分配" in content + assert "specialist=ops" in content + assert "请检查并修复网关启动失败" in content + def test_gateway_comprehensive_mode_has_manager_final_pass(monkeypatch: pytest.MonkeyPatch) -> None: class Store: def get_setting(self, _k: str) -> str: diff --git a/tests/test_oclaw_tool_result_guard.py b/tests/test_oclaw_tool_result_guard.py index 3e1757fa..adcee740 100644 --- a/tests/test_oclaw_tool_result_guard.py +++ b/tests/test_oclaw_tool_result_guard.py @@ -41,3 +41,40 @@ def test_oclaw_tool_result_context_guard_truncates_large_tool_message(tmp_path: assert len(guarded) <= _OCLAW_TOOL_RESULT_HARD_CAP_CHARS + 2000 assert "_tool_result_guarded" in guarded + +def test_oclaw_tool_result_context_guard_skips_active_turn_tool_messages(tmp_path: Path) -> None: + store = SqliteStore(str(tmp_path / "ops.sqlite")) + sess = store.create_session("t") + turn_uuid = "turn-active-1" + store.add_message( + session_id=sess.id, + role="assistant", + content="", + tool_calls=[{"id": "c1", "name": "echo", "arguments": {"x": 1}}], + turn_uuid=turn_uuid, + ) + huge = "X" * (_OCLAW_TOOL_RESULT_HARD_CAP_CHARS + 5000) + store.add_message( + session_id=sess.id, + role="tool", + content='{"ok": true, "blob": "' + huge + '"}', + tool_calls={"tool_call_id": "c1", "name": "echo", "assistant_message_id": 1}, + turn_uuid=turn_uuid, + ) + msgs = _build_model_context( + store=store, + session_id=sess.id, + max_messages=50, + system_prompt="sys", + model=RuleBasedChatModel(), + lang="zh", + memory_context=None, + trace_id="t1", + parent_span_id=None, + active_turn_uuid=turn_uuid, + ) + tool_msgs = [m for m in msgs if m.get("role") == "tool"] + assert tool_msgs, msgs + raw = str(tool_msgs[-1].get("content") or "") + assert "_tool_result_guarded" not in raw + diff --git a/tests/test_oclaw_worker_title.py b/tests/test_oclaw_worker_title.py new file mode 100644 index 00000000..9f74ba1a --- /dev/null +++ b/tests/test_oclaw_worker_title.py @@ -0,0 +1,59 @@ +from __future__ import annotations + +import hashlib +from pathlib import Path + +from oclaw.platform.persistence.sqlite_store import SqliteStore +from oclaw.runtime.types import StandardMessage +from oclaw.runtime.worker import _maybe_generate_title_on_third_round + + +class _DummyResp: + def __init__(self, content: str) -> None: + self.content = content + + +class _DummyModel: + def chat(self, _messages, _tools, on_token=None): # noqa: ANN001 + return _DummyResp("第三轮标题") + + +def test_worker_third_round_title_generation_updates_stage3(tmp_path: Path) -> None: + store = SqliteStore(str(tmp_path / "worker-title.sqlite")) + tenant = store.create_tenant("Team") + user = store.create_user_account( + tenant_id=str(tenant["id"]), + username="tester", + display_name="Tester", + role="owner", + password_hash=hashlib.sha256("test-pass".encode("utf-8")).hexdigest(), + is_active=True, + ) + session = store.create_session_for_user( + title="第一轮标题", + tenant_id=str(tenant["id"]), + user_id=str(user["id"]), + ) + sid = str(session.id) + store.set_setting(f"AIA_SESSION_AUTO_TITLE_STAGE:{sid}", "1") + store.add_message(session_id=sid, role="user", content="第一轮问题") + store.add_message(session_id=sid, role="assistant", content="第一轮回答") + store.add_message(session_id=sid, role="user", content="第二轮问题") + store.add_message(session_id=sid, role="assistant", content="第二轮回答") + + msg = StandardMessage( + session_id=sid, + tenant_id=str(tenant["id"]), + user_id=str(user["id"]), + role="member", + channel="admin_chat", + text="第三轮问题", + attachments=[], + metadata={"lang": "zh"}, + ) + _maybe_generate_title_on_third_round(store=store, msg=msg, model=_DummyModel()) + + renamed = store.get_session(sid) + assert renamed is not None + assert str(getattr(renamed, "title", "") or "") == "第三轮标题" + assert str(store.get_setting(f"AIA_SESSION_AUTO_TITLE_STAGE:{sid}") or "") == "3" diff --git a/tests/test_tool_llm_truncation.py b/tests/test_tool_llm_truncation.py index df0918d7..806f9676 100644 --- a/tests/test_tool_llm_truncation.py +++ b/tests/test_tool_llm_truncation.py @@ -3,7 +3,12 @@ from __future__ import annotations import json import unittest -from oclaw.runtime.chat.tool_runtime import truncate_tool_result_for_llm_messages, tool_llm_message_max_chars +from oclaw.runtime.chat.tool_runtime import ( + compact_turn_tool_messages_for_storage, + tool_llm_message_max_chars, + truncate_tool_result_for_llm_messages, +) +from oclaw.platform.persistence.sqlite_store import SqliteStore class ToolLlmTruncationTests(unittest.TestCase): @@ -28,6 +33,46 @@ class ToolLlmTruncationTests(unittest.TestCase): 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 + + db = f"{tempfile.gettempdir()}/oclaw-test-{uuid.uuid4().hex}.sqlite" + store = SqliteStore(db) + sess = store.create_session("t") + turn_uuid = "turn-1" + store.add_message( + session_id=sess.id, + role="assistant", + content="", + tool_calls=[{"id": "c1", "name": "echo", "arguments": {"x": 1}}], + turn_uuid=turn_uuid, + ) + huge = "x" * 200_000 + row = store.add_message( + session_id=sess.id, + role="tool", + content=json.dumps({"ok": True, "blob": huge}, ensure_ascii=False), + 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): + stats = compact_turn_tool_messages_for_storage( + store=store, + session_id=sess.id, + turn_uuid=turn_uuid, + ) + 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 "")) + if __name__ == "__main__": unittest.main() diff --git a/tests/test_tool_pairing_messages.py b/tests/test_tool_pairing_messages.py index 07ffffa3..f83b693c 100644 --- a/tests/test_tool_pairing_messages.py +++ b/tests/test_tool_pairing_messages.py @@ -165,3 +165,52 @@ def test_signature_metadata_can_be_forced_on_via_env(monkeypatch) -> None: tc = (assistant.get("tool_calls") or [])[0] assert tc.get("extra_content", {}).get("google", {}).get("thought_signature") == "sig_abc" + +def test_current_turn_tool_content_not_clipped_with_details_phrase() -> None: + model = RuleBasedChatModel() + rows = [ + _Msg( + "assistant", + "", + tool_calls=json.dumps([{"id": "call_now", "name": "t", "arguments": {}}], ensure_ascii=False), + event_type="tool_call", + turn_uuid="turn-now", + ), + _Msg( + "tool", + json.dumps({"ok": True, "blob": "x" * 1200}, ensure_ascii=False), + tool_calls=json.dumps({"tool_call_id": "call_now", "name": "t"}, ensure_ascii=False), + event_type="tool_result", + turn_uuid="turn-now", + ), + ] + msgs = build_llm_messages(store_messages=rows, system_prompt="s", model=model, lang="zh", tool_context_truncate_enabled=True) + tool_rows = [m for m in msgs if m.get("role") == "tool"] + assert tool_rows + assert "详情请重新阅读" not in str(tool_rows[-1].get("content") or "") + + +def test_historical_tool_content_can_include_details_phrase_when_clamped(monkeypatch) -> None: + monkeypatch.setenv("AIA_REPLAY_TOOL_FULL_ROUNDS", "0") + model = RuleBasedChatModel() + rows = [ + _Msg( + "assistant", + "", + tool_calls=json.dumps([{"id": "call_hist", "name": "t", "arguments": {}}], ensure_ascii=False), + event_type="tool_call", + turn_uuid="turn-hist", + ), + _Msg( + "tool", + json.dumps({"ok": True, "blob": "x" * 1600}, ensure_ascii=False), + tool_calls=json.dumps({"tool_call_id": "call_hist", "name": "t"}, ensure_ascii=False), + event_type="tool_result", + turn_uuid="turn-hist", + ), + ] + msgs = build_llm_messages(store_messages=rows, system_prompt="s", model=model, lang="zh", tool_context_truncate_enabled=True) + tool_rows = [m for m in msgs if m.get("role") == "tool"] + assert tool_rows + assert "详情请重新阅读" in str(tool_rows[-1].get("content") or "") + diff --git a/tests/test_workspace_path_guard.py b/tests/test_workspace_path_guard.py index 17a03d63..fabcd2a8 100644 --- a/tests/test_workspace_path_guard.py +++ b/tests/test_workspace_path_guard.py @@ -13,12 +13,15 @@ from oclaw.interfaces.http.fastapi_app import create_app from oclaw.platform.config.paths import db_path from oclaw.platform.persistence.sqlite_store import SqliteStore from oclaw.runtime.tools.experts.workspace.fs_tools import list_files_tool, write_file_tool +from oclaw.runtime.tools.experts.workspace.shell_tools import run_command_tool from oclaw.runtime.tools.experts.workspace.workspace_base import ( access_from_env, build_workspace_path_access, clear_workspace_path_access_for_tests, + current_workspace_write_namespace, resolve_workspace_path, workspace_path_access_scope, + workspace_write_namespace_scope, ) @@ -111,10 +114,223 @@ class WorkspacePathGuardTests(unittest.TestCase): with workspace_path_access_scope(None, None): r = spec.handler({"path": "generated.py", "content": "print('ok')\n", "mode": "overwrite"}) self.assertTrue(r.get("ok"), r) - expected = (self.root / "data" / self.root.name / "generated.py").resolve() + expected = (self.root / "data" / "workspace" / "generated.py").resolve() self.assertEqual(str(expected), str(r.get("path"))) self.assertTrue(expected.exists()) + def test_write_file_relative_path_uses_workspace_namespace_scope(self) -> None: + with mock.patch.dict( + os.environ, + {"OPS_WORKSPACE_ROOT": str(self.root), "OPS_WORKSPACE_EXTRA_ROOTS": "", "OPS_WORKSPACE_ALLOW_ANY_PATH": ""}, + clear=False, + ): + clear_workspace_path_access_for_tests() + spec = write_file_tool() + with workspace_path_access_scope(None, None), workspace_write_namespace_scope("ops"): + self.assertEqual(current_workspace_write_namespace(), "ops") + r = spec.handler({"path": "generated.py", "content": "print('ok')\n", "mode": "overwrite"}) + self.assertTrue(r.get("ok"), r) + expected = (self.root / "data" / "workspace" / "generated.py").resolve() + self.assertEqual(str(expected), str(r.get("path"))) + self.assertTrue(expected.exists()) + + def test_write_file_absolute_path_is_forced_into_workspace_sandbox(self) -> None: + with mock.patch.dict( + os.environ, + {"OPS_WORKSPACE_ROOT": str(self.root), "OPS_WORKSPACE_EXTRA_ROOTS": "", "OPS_WORKSPACE_ALLOW_ANY_PATH": ""}, + clear=False, + ): + clear_workspace_path_access_for_tests() + spec = write_file_tool() + abs_target = str((self.root / "count_items.py").resolve()) + with workspace_path_access_scope(None, None), workspace_write_namespace_scope("ops"): + r = spec.handler({"path": abs_target, "content": "print('ok')\n", "mode": "overwrite"}) + self.assertTrue(r.get("ok"), r) + expected = (self.root / "data" / "workspace" / "count_items.py").resolve() + self.assertEqual(str(expected), str(r.get("path"))) + self.assertTrue(expected.exists()) + + def test_run_command_default_cwd_uses_workspace_namespace_sandbox(self) -> None: + with mock.patch.dict( + os.environ, + { + "OPS_WORKSPACE_ROOT": str(self.root), + "OPS_WORKSPACE_EXTRA_ROOTS": "", + "OPS_WORKSPACE_ALLOW_ANY_PATH": "", + "AIA_ENABLE_RUN_COMMAND": "1", + }, + clear=False, + ): + clear_workspace_path_access_for_tests() + spec = run_command_tool() + with workspace_path_access_scope(None, None), workspace_write_namespace_scope("ops"): + r = spec.handler({"command": "python -c \"print('ok')\""}) + self.assertTrue(r.get("ok"), r) + expected_cwd = (self.root / "data" / "workspace").resolve() + self.assertEqual(str(expected_cwd), str(r.get("cwd"))) + + def test_run_command_strips_leading_cd_chain_in_default_sandbox(self) -> None: + with mock.patch.dict( + os.environ, + { + "OPS_WORKSPACE_ROOT": str(self.root), + "OPS_WORKSPACE_EXTRA_ROOTS": "", + "OPS_WORKSPACE_ALLOW_ANY_PATH": "", + "AIA_ENABLE_RUN_COMMAND": "1", + }, + clear=False, + ): + clear_workspace_path_access_for_tests() + spec = run_command_tool() + cmd = f'cd /d "{self.root}" && python -c "print(123)"' + with workspace_path_access_scope(None, None), workspace_write_namespace_scope("ops"): + r = spec.handler({"command": cmd}) + self.assertTrue(r.get("ok"), r) + self.assertTrue(bool(r.get("normalized_cd_removed"))) + expected_cwd = (self.root / "data" / "workspace").resolve() + self.assertEqual(str(expected_cwd), str(r.get("cwd"))) + self.assertIn("123", str(r.get("output") or "")) + + def test_run_command_strips_windows_drive_prefix_cd_chain(self) -> None: + with mock.patch.dict( + os.environ, + { + "OPS_WORKSPACE_ROOT": str(self.root), + "OPS_WORKSPACE_EXTRA_ROOTS": "", + "OPS_WORKSPACE_ALLOW_ANY_PATH": "", + "AIA_ENABLE_RUN_COMMAND": "1", + }, + clear=False, + ): + clear_workspace_path_access_for_tests() + ws_root = self.root / "data" / "workspace" + ws_root.mkdir(parents=True, exist_ok=True) + (ws_root / "count_directory.py").write_text("print('drive-cd-ok')\n", encoding="utf-8") + spec = run_command_tool() + cmd = f'D: && cd /d "{self.root}" && python count_directory.py' + with workspace_path_access_scope(None, None): + r = spec.handler({"command": cmd}) + self.assertTrue(r.get("ok"), r) + self.assertTrue(bool(r.get("normalized_cd_removed")), r) + self.assertTrue(bool(r.get("script_path_rewritten")), r) + self.assertIn("drive-cd-ok", str(r.get("output") or "")) + + def test_run_command_output_flags_distinguish_empty_from_truncation(self) -> None: + with mock.patch.dict( + os.environ, + { + "OPS_WORKSPACE_ROOT": str(self.root), + "OPS_WORKSPACE_EXTRA_ROOTS": "", + "OPS_WORKSPACE_ALLOW_ANY_PATH": "", + "AIA_ENABLE_RUN_COMMAND": "1", + }, + clear=False, + ): + clear_workspace_path_access_for_tests() + spec = run_command_tool() + with workspace_path_access_scope(None, None), workspace_write_namespace_scope("ops"): + r = spec.handler({"command": 'python -c "pass"'}) + self.assertTrue(r.get("ok"), r) + self.assertTrue(bool(r.get("output_empty"))) + self.assertFalse(bool(r.get("output_truncated"))) + self.assertTrue(bool(r.get("output_not_truncated"))) + self.assertEqual(str(r.get("error_code") or ""), "") + + def test_run_command_nonzero_exit_marks_failure_not_truncation(self) -> None: + with mock.patch.dict( + os.environ, + { + "OPS_WORKSPACE_ROOT": str(self.root), + "OPS_WORKSPACE_EXTRA_ROOTS": "", + "OPS_WORKSPACE_ALLOW_ANY_PATH": "", + "AIA_ENABLE_RUN_COMMAND": "1", + }, + clear=False, + ): + clear_workspace_path_access_for_tests() + spec = run_command_tool() + with workspace_path_access_scope(None, None), workspace_write_namespace_scope("ops"): + r = spec.handler({"command": 'python -c "import sys; sys.exit(3)"'}) + self.assertFalse(bool(r.get("ok"))) + self.assertEqual(int(r.get("exit_code") or 0), 3) + self.assertEqual(str(r.get("error_code") or ""), "command_exit_nonzero") + self.assertFalse(bool(r.get("output_truncated"))) + + def test_run_command_rewrites_workspace_absolute_script_path_to_sandbox(self) -> None: + with mock.patch.dict( + os.environ, + { + "OPS_WORKSPACE_ROOT": str(self.root), + "OPS_WORKSPACE_EXTRA_ROOTS": "", + "OPS_WORKSPACE_ALLOW_ANY_PATH": "", + "AIA_ENABLE_RUN_COMMAND": "1", + }, + clear=False, + ): + clear_workspace_path_access_for_tests() + # Prepare script inside sandbox, but command will reference repo-root absolute path. + ws_script = self.root / "data" / "workspace" / "count_files.py" + ws_script.parent.mkdir(parents=True, exist_ok=True) + ws_script.write_text("print('sandbox-ok')\n", encoding="utf-8") + spec = run_command_tool() + absolute_repo_script = str((self.root / "count_files.py").resolve()) + with workspace_path_access_scope(None, None), workspace_write_namespace_scope("ops"): + r = spec.handler({"command": f'python "{absolute_repo_script}"'}) + self.assertTrue(r.get("ok"), r) + self.assertTrue(bool(r.get("command_rewritten")), r) + self.assertIn("sandbox-ok", str(r.get("output") or "")) + + def test_run_command_explicit_repo_root_cwd_is_redirected_to_sandbox(self) -> None: + with mock.patch.dict( + os.environ, + { + "OPS_WORKSPACE_ROOT": str(self.root), + "OPS_WORKSPACE_EXTRA_ROOTS": "", + "OPS_WORKSPACE_ALLOW_ANY_PATH": "", + "AIA_ENABLE_RUN_COMMAND": "1", + }, + clear=False, + ): + clear_workspace_path_access_for_tests() + ws_script = self.root / "data" / "workspace" / "count_files.py" + ws_script.parent.mkdir(parents=True, exist_ok=True) + ws_script.write_text("print('redirect-ok')\n", encoding="utf-8") + spec = run_command_tool() + abs_repo_script = str((self.root / "count_files.py").resolve()) + with workspace_path_access_scope(None, None): + r = spec.handler( + { + "command": f'python "{abs_repo_script}"', + "cwd": str(self.root), + } + ) + self.assertTrue(r.get("ok"), r) + self.assertTrue(bool(r.get("cwd_redirected_to_sandbox")), r) + self.assertTrue(bool(r.get("command_rewritten")), r) + self.assertIn("redirect-ok", str(r.get("output") or "")) + + def test_run_command_rewrites_relative_python_script_to_sandbox_root(self) -> None: + with mock.patch.dict( + os.environ, + { + "OPS_WORKSPACE_ROOT": str(self.root), + "OPS_WORKSPACE_EXTRA_ROOTS": "", + "OPS_WORKSPACE_ALLOW_ANY_PATH": "", + "AIA_ENABLE_RUN_COMMAND": "1", + }, + clear=False, + ): + clear_workspace_path_access_for_tests() + ws_root = self.root / "data" / "workspace" + ws_root.mkdir(parents=True, exist_ok=True) + (ws_root / "count_directory.py").write_text("print('found-in-sandbox-root')\n", encoding="utf-8") + spec = run_command_tool() + with workspace_path_access_scope(None, None): + r = spec.handler({"command": "python count_directory.py"}) + self.assertTrue(r.get("ok"), r) + self.assertTrue(bool(r.get("script_path_rewritten")), r) + self.assertIn("found-in-sandbox-root", str(r.get("output") or "")) + def test_per_user_extra_roots_from_db(self) -> None: f = self.extra / "u.txt" f.write_text("u", encoding="utf-8")