feat: enhance session management to isolate agent data and prevent cross-agent visibility

This commit is contained in:
X-iZhang
2026-02-12 00:47:28 +00:00
parent 10266c660f
commit 56711aa321
4 changed files with 59 additions and 23 deletions
+1 -1
View File
@@ -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()
+26 -9
View File
@@ -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)
+15 -13
View File
@@ -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)]()
<!-- [![arXiv](https://img.shields.io/badge/arXiv-xxxx.xxxx-b31b1b?style=for-the-badge&logo=arxiv&logoColor=b31b1b)]() -->
<!-- [![Gradio Demo](https://img.shields.io/badge/Gradio-Online_Demo-FFCC00?style=for-the-badge&logo=gradio&logoColor=yellow&labelColor=grey)]()
[![Evaluation Split](https://img.shields.io/badge/HF-Test_Dataset-AECBFA?style=for-the-badge&logo=huggingface&logoColor=FFCC00&labelColor=grey)]() -->
@@ -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
<table>
<tbody>
<tr>
<td align="center">
<a href="https://x-izhang.github.io/">
<img src="https://x-izhang.github.io/author/xi-zhang/avatar_hu13660783057866068725.jpg"
width="100" height="100"
style="object-fit: cover; border-radius: 20%;" alt="Xi Zhang"/>
<br />
<sub><b>Xi Zhang</b><sup>†</sup></sub>
</a>
</td>
<td align="center">
<a href="https://youganglyu.github.io/">
<img src="https://youganglyu.github.io/images/profile.png"
@@ -453,15 +462,6 @@ We thank the authors for their valuable contributions to the open-source communi
<sub><b>Yougang Lyu</b></sub>
</a>
</td>
<td align="center">
<a href="https://x-izhang.github.io/">
<img src="https://x-izhang.github.io/author/xi-zhang/avatar_hu13660783057866068725.jpg"
width="100" height="100"
style="object-fit: cover; border-radius: 20%;" alt="Xi Zhang"/>
<br />
<sub><b>Xi Zhang</b></sub>
</a>
</td>
<td align="center">
<a href="https://din0s.me/">
<img src="https://din0s.me/images/pk.jpg"
@@ -475,7 +475,9 @@ We thank the authors for their valuable contributions to the open-source communi
</tbody>
</table>
For any enquiries or collaboration opportunities, please contact: [**youganglyu@gmail.com**](mailto:youganglyu@gmail.com)
<sup>†</sup> Project Leader
For any enquiries or collaboration opportunities, please contact: [**EvoScientist.ai@gmail.com**](mailto:evoscientist.ai@gmail.com)
## 📜 License
+17
View File
@@ -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()