From 56711aa32130d4128e4030c60df2329cd111e725 Mon Sep 17 00:00:00 2001 From: X-iZhang Date: Thu, 12 Feb 2026 00:47:28 +0000 Subject: [PATCH] feat: enhance session management to isolate agent data and prevent cross-agent visibility --- EvoScientist/cli/commands.py | 2 +- EvoScientist/sessions.py | 35 ++++++++++++++++++++++++++--------- README.md | 28 +++++++++++++++------------- tests/test_sessions.py | 17 +++++++++++++++++ 4 files changed, 59 insertions(+), 23 deletions(-) diff --git a/EvoScientist/cli/commands.py b/EvoScientist/cli/commands.py index b1c40cb..dcc1ec0 100644 --- a/EvoScientist/cli/commands.py +++ b/EvoScientist/cli/commands.py @@ -416,7 +416,7 @@ def _main_callback( console.print("[dim]Loading agent...[/dim]") agent = _load_agent(workspace_dir=workspace_dir, checkpointer=checkpointer) tid = thread_id or generate_thread_id() - cmd_run(agent, prompt, thread_id=tid, show_thinking=show_thinking, workspace_dir=workspace_dir) + cmd_run(agent, prompt, thread_id=tid, show_thinking=show_thinking, workspace_dir=workspace_dir, model=config.model) import nest_asyncio # type: ignore[import-untyped] nest_asyncio.apply() diff --git a/EvoScientist/sessions.py b/EvoScientist/sessions.py index 55006c6..18c0933 100644 --- a/EvoScientist/sessions.py +++ b/EvoScientist/sessions.py @@ -229,18 +229,22 @@ async def get_most_recent() -> str | None: async def thread_exists(thread_id: str) -> bool: - """Return ``True`` if *thread_id* has at least one checkpoint.""" + """Return ``True`` if *thread_id* has at least one EvoScientist checkpoint.""" db_path = str(get_db_path()) async with aiosqlite.connect(db_path, timeout=30.0) as conn: if not await _table_exists(conn, "checkpoints"): return False - query = "SELECT 1 FROM checkpoints WHERE thread_id = ? LIMIT 1" - async with conn.execute(query, (thread_id,)) as cur: + query = """ + SELECT 1 FROM checkpoints + WHERE thread_id = ? AND json_extract(metadata, '$.agent_name') = ? + LIMIT 1 + """ + async with conn.execute(query, (thread_id, AGENT_NAME)) as cur: return (await cur.fetchone()) is not None async def find_similar_threads(thread_id: str, limit: int = 5) -> list[str]: - """Find thread IDs that start with *thread_id* (prefix match).""" + """Find EvoScientist thread IDs that start with *thread_id* (prefix match).""" db_path = str(get_db_path()) async with aiosqlite.connect(db_path, timeout=30.0) as conn: if not await _table_exists(conn, "checkpoints"): @@ -249,22 +253,24 @@ async def find_similar_threads(thread_id: str, limit: int = 5) -> list[str]: SELECT DISTINCT thread_id FROM checkpoints WHERE thread_id LIKE ? + AND json_extract(metadata, '$.agent_name') = ? ORDER BY thread_id LIMIT ? """ - async with conn.execute(query, (thread_id + "%", limit)) as cur: + async with conn.execute(query, (thread_id + "%", AGENT_NAME, limit)) as cur: rows = await cur.fetchall() return [r[0] for r in rows] async def delete_thread(thread_id: str) -> bool: - """Delete all checkpoints (and writes) for *thread_id*.""" + """Delete all EvoScientist checkpoints (and writes) for *thread_id*.""" db_path = str(get_db_path()) async with aiosqlite.connect(db_path, timeout=30.0) as conn: if not await _table_exists(conn, "checkpoints"): return False cur = await conn.execute( - "DELETE FROM checkpoints WHERE thread_id = ?", (thread_id,) + "DELETE FROM checkpoints WHERE thread_id = ? AND json_extract(metadata, '$.agent_name') = ?", + (thread_id, AGENT_NAME), ) deleted = cur.rowcount > 0 if await _table_exists(conn, "writes"): @@ -288,10 +294,11 @@ async def get_thread_metadata(thread_id: str) -> dict | None: json_extract(metadata, '$.updated_at') as updated_at FROM checkpoints WHERE thread_id = ? + AND json_extract(metadata, '$.agent_name') = ? ORDER BY checkpoint_id DESC LIMIT 1 """ - async with conn.execute(query, (thread_id,)) as cur: + async with conn.execute(query, (thread_id, AGENT_NAME)) as cur: row = await cur.fetchone() if not row: return None @@ -305,11 +312,21 @@ async def get_thread_metadata(thread_id: str) -> dict | None: async def get_thread_messages(thread_id: str) -> list: """Return the list of LangChain message objects for *thread_id*. + Only returns messages for EvoScientist threads. Returns an empty list if the thread has no checkpoints. """ db_path = str(get_db_path()) - serde = JsonPlusSerializer() async with aiosqlite.connect(db_path, timeout=30.0) as conn: if not await _table_exists(conn, "checkpoints"): return [] + # Verify this thread belongs to EvoScientist before loading messages + check = """ + SELECT 1 FROM checkpoints + WHERE thread_id = ? AND json_extract(metadata, '$.agent_name') = ? + LIMIT 1 + """ + async with conn.execute(check, (thread_id, AGENT_NAME)) as cur: + if not await cur.fetchone(): + return [] + serde = JsonPlusSerializer() return await _load_checkpoint_messages(conn, thread_id, serde) diff --git a/README.md b/README.md index aa2b89c..ba90e00 100644 --- a/README.md +++ b/README.md @@ -11,8 +11,8 @@ [![PyPI](https://img.shields.io/badge/PyPI-EvoScientist%20v0.0.1-3da9fc?style=for-the-badge&logo=python&logoColor=3da9fc)](https://pypi.org/project/EvoScientist/) [![Project Page](https://img.shields.io/badge/Project-Page-ff8e3c?style=for-the-badge&logo=googlelens&logoColor=ff8e3c)]() -[![arXiv](https://img.shields.io/badge/arXiv-xxxx.xxxx-b31b1b?style=for-the-badge&logo=arxiv&logoColor=b31b1b)]() [![License](https://img.shields.io/badge/License-MIT-green?style=for-the-badge)]() + @@ -39,7 +39,7 @@ - [📊 Evaluation](#-evaluation) - [📝 Citation](#-citation) - [📚 Acknowledgments](#-acknowledgments) -- [📦 Codebase Contributors](#-codebase-contributors) +- [📦 EvoScientist Team](#-evoscientist-team) - [📜 License](#-license) ## 🤖 Supported Models @@ -439,11 +439,20 @@ This project builds upon the following outstanding open-source works: We thank the authors for their valuable contributions to the open-source community. -## 📦 Codebase Contributors +## 📦 EvoScientist Team + -
+ + Xi Zhang +
+ Xi Zhang† +
+
Yougang Lyu - - Xi Zhang -
- Xi Zhang -
-
-For any enquiries or collaboration opportunities, please contact: [**youganglyu@gmail.com**](mailto:youganglyu@gmail.com) +† Project Leader + +For any enquiries or collaboration opportunities, please contact: [**EvoScientist.ai@gmail.com**](mailto:evoscientist.ai@gmail.com) ## 📜 License diff --git a/tests/test_sessions.py b/tests/test_sessions.py index 7d6db51..45fc5fe 100644 --- a/tests/test_sessions.py +++ b/tests/test_sessions.py @@ -212,6 +212,23 @@ class TestThreadFunctions(unittest.TestCase): def test_delete_nonexistent(self): self.assertFalse(_run(delete_thread("nope1234"))) + # -- Agent isolation: OtherAgent data should never be visible -- + + def test_thread_exists_ignores_other_agent(self): + self.assertFalse(_run(thread_exists("zzz99999"))) + + def test_find_similar_ignores_other_agent(self): + similar = _run(find_similar_threads("zzz")) + self.assertEqual(len(similar), 0) + + def test_get_metadata_ignores_other_agent(self): + meta = _run(get_thread_metadata("zzz99999")) + self.assertIsNone(meta) + + def test_delete_ignores_other_agent(self): + # Should not delete OtherAgent's data + self.assertFalse(_run(delete_thread("zzz99999"))) + if __name__ == "__main__": unittest.main()