feat: enhance session management to isolate agent data and prevent cross-agent visibility
This commit is contained in:
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -11,8 +11,8 @@
|
||||
|
||||
[](https://pypi.org/project/EvoScientist/)
|
||||
[]()
|
||||
[]()
|
||||
[]()
|
||||
<!-- []() -->
|
||||
<!-- []()
|
||||
[]() -->
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user