diff --git a/frontends/desktop_bridge.py b/frontends/desktop_bridge.py index a42388d8f..c79dd531e 100644 --- a/frontends/desktop_bridge.py +++ b/frontends/desktop_bridge.py @@ -197,7 +197,13 @@ def _session_dict(self, s: "Session") -> dict: "llm_history": llm_hist} def _session_file(self, sid: str) -> Path: - return self._sessions_dir / f"{sid}.json" + if not isinstance(sid, str) or not sid or "/" in sid or '\\' in sid: + raise ValueError("invalid session id") + root = self._sessions_dir.resolve() + target = (root / f"{sid}.json").resolve() + if target.parent != root: + raise ValueError(f"invalid session id: {sid!r}") + return target def _persist_session(self, s: "Session"): """Write a single session file. Cost is O(one session), independent of how many @@ -206,9 +212,10 @@ def _persist_session(self, s: "Session"): self._sessions_dir.mkdir(parents=True, exist_ok=True) with self.lock: data = self._session_dict(s) - tmp = self._sessions_dir / f"{s.id}.json.tmp" + target = self._session_file(s.id) + tmp = target.with_suffix(target.suffix + ".tmp") tmp.write_text(json.dumps(data, ensure_ascii=False, default=str), encoding="utf-8") - os.replace(tmp, self._session_file(s.id)) # atomic swap + os.replace(tmp, target) # atomic swap except Exception as e: print(f"[bridge] persist session {s.id} failed: {e}", file=sys.stderr) @@ -228,8 +235,10 @@ def _persist(self): self._persist_session(s) def _session_from_item(self, item: dict) -> "Session": + sid = item["id"] + self._session_file(sid) msgs = item.get("messages", []) - return Session(id=item["id"], title=item.get("title", "New chat"), + return Session(id=sid, title=item.get("title", "New chat"), cwd=item.get("cwd", self.ga_root), created_at=item.get("created_at", time.time()), updated_at=item.get("updated_at", time.time()), @@ -238,7 +247,7 @@ def _session_from_item(self, item: dict) -> "Session": pinned=item.get("pinned", False), untitled=item.get("untitled", True), plan_scan_baseline=_load_plan_baseline(item, msgs), - plan_path=_sanitize_desktop_plan_path(item["id"], item.get("plan_path") or ""), + plan_path=_sanitize_desktop_plan_path(sid, item.get("plan_path") or ""), status="idle", agent=None, llm_history=item.get("llm_history"), llm_no=item.get("llm_no")) @@ -324,6 +333,11 @@ def import_sessions(self, source_dir: str) -> dict: if not sid or sid in self.sessions: skipped += 1 continue + try: + self._session_file(sid) + except ValueError: + skipped += 1 + continue sess = self._session_from_item(item) self.sessions[sid] = sess new_sessions.append(sess) diff --git a/tests/test_desktop_session_ids.py b/tests/test_desktop_session_ids.py new file mode 100644 index 000000000..a2242bae7 --- /dev/null +++ b/tests/test_desktop_session_ids.py @@ -0,0 +1,135 @@ +import importlib.util +import json +import os +import shutil +import sys +import tempfile +import unittest +from pathlib import Path + + +ROOT = Path(__file__).resolve().parents[1] +FRONTENDS = ROOT / "frontends" +_TEST_GA_ROOT = tempfile.TemporaryDirectory() +unittest.addModuleCleanup(_TEST_GA_ROOT.cleanup) +_TEST_GA_PATH = Path(_TEST_GA_ROOT.name) +(_TEST_GA_PATH / "agentmain.py").touch() + + +def _load_module(name: str, path: Path): + spec = importlib.util.spec_from_file_location(name, path) + if spec is None or spec.loader is None: + raise ImportError(f"failed to load {path}") + module = importlib.util.module_from_spec(spec) + sys.modules[name] = module + spec.loader.exec_module(module) + return module + + +_load_module("plan_state", FRONTENDS / "plan_state.py") +_old_ga_root = os.environ.get("GA_ROOT") +_old_argv = sys.argv[:] +os.environ["GA_ROOT"] = str(_TEST_GA_PATH) +sys.argv = [sys.argv[0]] +try: + bridge = _load_module("desktop_bridge_session_id_test", FRONTENDS / "desktop_bridge.py") +finally: + sys.argv = _old_argv + if _old_ga_root is None: + os.environ.pop("GA_ROOT", None) + else: + os.environ["GA_ROOT"] = _old_ga_root + + +class DesktopSessionIdPersistenceTests(unittest.TestCase): + def setUp(self): + shutil.rmtree(_TEST_GA_PATH / "temp", ignore_errors=True) + for path in _TEST_GA_PATH.glob("*.json"): + path.unlink() + self.manager = bridge.AgentManager() + self.manager.sessions.clear() + self.manager.active_session_id = None + + @staticmethod + def _write_session(source: Path, sid: str): + sessions = source / "temp" / "desktop_sessions" + sessions.mkdir(parents=True, exist_ok=True) + (sessions / "import.json").write_text( + json.dumps({"id": sid, "messages": [], "msg_seq": 0}), + encoding="utf-8", + ) + + def test_import_rejects_traversal_id_without_writing_outside_store(self): + with tempfile.TemporaryDirectory() as tmp: + source = Path(tmp) / "source" + self._write_session(source, "../../escape") + escaped = _TEST_GA_PATH / "escape.json" + + result = self.manager.import_sessions(str(source)) + + self.assertEqual(result["sessionsAdded"], 0) + self.assertEqual(result["sessionsSkipped"], 1) + self.assertNotIn("../../escape", self.manager.sessions) + self.assertFalse(escaped.exists()) + + def test_import_rejects_absolute_id_without_writing_outside_store(self): + sid = str(_TEST_GA_PATH / "absolute-escape") + escaped = Path(f"{sid}.json") + with tempfile.TemporaryDirectory() as tmp: + source = Path(tmp) / "source" + self._write_session(source, sid) + + result = self.manager.import_sessions(str(source)) + + self.assertEqual(result["sessionsAdded"], 0) + self.assertEqual(result["sessionsSkipped"], 1) + self.assertNotIn(sid, self.manager.sessions) + self.assertFalse(escaped.exists()) + + def test_import_rejects_normalized_alias_that_collides_with_existing_session(self): + safe_id = "sess-safe123" + alias_id = f"nested/../{safe_id}" + with tempfile.TemporaryDirectory() as tmp: + source = Path(tmp) / "source" + self._write_session(source, safe_id) + first = self.manager.import_sessions(str(source)) + self._write_session(source, alias_id) + second = self.manager.import_sessions(str(source)) + + persisted = json.loads( + (self.manager._sessions_dir / f"{safe_id}.json").read_text(encoding="utf-8") + ) + self.assertEqual(first["sessionsAdded"], 1) + self.assertEqual(second["sessionsAdded"], 0) + self.assertEqual(second["sessionsSkipped"], 1) + self.assertNotIn(alias_id, self.manager.sessions) + self.assertEqual(persisted["id"], safe_id) + + def test_import_rejects_windows_style_path_separator(self): + sid = r"nested\..\sess-safe123" + with tempfile.TemporaryDirectory() as tmp: + source = Path(tmp) / "source" + self._write_session(source, sid) + + result = self.manager.import_sessions(str(source)) + + self.assertEqual(result["sessionsAdded"], 0) + self.assertEqual(result["sessionsSkipped"], 1) + self.assertNotIn(sid, self.manager.sessions) + + def test_import_keeps_valid_session_ids(self): + sid = "sess-safe456" + with tempfile.TemporaryDirectory() as tmp: + source = Path(tmp) / "source" + self._write_session(source, sid) + + result = self.manager.import_sessions(str(source)) + + self.assertEqual(result["sessionsAdded"], 1) + self.assertEqual(result["sessionsSkipped"], 0) + self.assertIn(sid, self.manager.sessions) + self.assertTrue((self.manager._sessions_dir / f"{sid}.json").is_file()) + + +if __name__ == "__main__": + unittest.main()