From 52719319353dc64a386549e2619f0512a03e44e0 Mon Sep 17 00:00:00 2001 From: X-iZhang Date: Thu, 12 Feb 2026 18:30:27 +0000 Subject: [PATCH] feat(tests): add tests for _merge_memory backslash safety and _step_mcp_servers npx handling --- EvoScientist/memory.py | 6 ++-- EvoScientist/onboard.py | 12 ++++++- EvoScientist/sessions.py | 14 ++++++-- tests/test_memory_merge.py | 73 ++++++++++++++++++++++++++++++++++++++ tests/test_onboard.py | 44 +++++++++++++++++++++++ tests/test_sessions.py | 53 +++++++++++++++++++++++++++ 6 files changed, 195 insertions(+), 7 deletions(-) create mode 100644 tests/test_memory_merge.py diff --git a/EvoScientist/memory.py b/EvoScientist/memory.py index bc11f9e..13ef8a3 100644 --- a/EvoScientist/memory.py +++ b/EvoScientist/memory.py @@ -302,8 +302,7 @@ def _merge_memory(existing_md: str, extracted: dict[str, Any]) -> str: if value and value != "null": # Replace the line "- **Label**: ..." with new value pattern = rf"(- \*\*{label}\*\*: ).*" - replacement = rf"\g<1>{value}" - result = re.sub(pattern, replacement, result) + result = re.sub(pattern, lambda m: m.group(1) + value, result) # --- Research Preferences --- prefs = extracted.get("research_preferences") @@ -320,8 +319,7 @@ def _merge_memory(existing_md: str, extracted: dict[str, Any]) -> str: value = prefs.get(key) if value and value != "null": pattern = rf"(- \*\*{label}\*\*: ).*" - replacement = rf"\g<1>{value}" - result = re.sub(pattern, replacement, result) + result = re.sub(pattern, lambda m: m.group(1) + value, result) # --- Experiment History (append) --- exp = extracted.get("experiment_conclusion") diff --git a/EvoScientist/onboard.py b/EvoScientist/onboard.py index 7fdf71d..6cbef5b 100644 --- a/EvoScientist/onboard.py +++ b/EvoScientist/onboard.py @@ -992,7 +992,17 @@ def _step_mcp_servers() -> list[str]: if srv["name"] in selected ) if needs_npx: - _ensure_npx("some MCP servers require Node.js") + if not _ensure_npx("some MCP servers require Node.js"): + npx_servers = { + srv["name"] + for srv in _RECOMMENDED_MCP_SERVERS + if srv["name"] in selected and srv.get("command") == "npx" + } + selected = [s for s in selected if s not in npx_servers] + if npx_servers: + console.print(f" [yellow]\u26a0 Skipping {', '.join(sorted(npx_servers))} (npx not available)[/yellow]") + if not selected: + return [] installed = [] for name in selected: diff --git a/EvoScientist/sessions.py b/EvoScientist/sessions.py index 18c0933..b8ccf73 100644 --- a/EvoScientist/sessions.py +++ b/EvoScientist/sessions.py @@ -268,13 +268,23 @@ async def delete_thread(thread_id: str) -> bool: async with aiosqlite.connect(db_path, timeout=30.0) as conn: if not await _table_exists(conn, "checkpoints"): return False + # Delete writes FIRST — the subquery needs checkpoints to still exist + if await _table_exists(conn, "writes"): + await conn.execute( + """DELETE FROM writes + WHERE thread_id = ? + AND checkpoint_id IN ( + SELECT checkpoint_id FROM checkpoints + WHERE thread_id = ? + AND json_extract(metadata, '$.agent_name') = ? + )""", + (thread_id, thread_id, AGENT_NAME), + ) cur = await conn.execute( "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"): - await conn.execute("DELETE FROM writes WHERE thread_id = ?", (thread_id,)) await conn.commit() return deleted diff --git a/tests/test_memory_merge.py b/tests/test_memory_merge.py new file mode 100644 index 0000000..bb95095 --- /dev/null +++ b/tests/test_memory_merge.py @@ -0,0 +1,73 @@ +"""Tests for _merge_memory — backslash-safe regex replacement.""" + +import pytest + +from EvoScientist.memory import _merge_memory, DEFAULT_MEMORY_TEMPLATE + + +class TestMergeMemoryBackslashSafety: + """Ensure values containing regex-special sequences survive _merge_memory.""" + + def test_backslash_n_preserved(self): + """A value containing literal '\\n' must not become a newline.""" + extracted = { + "user_profile": {"name": r"C:\new_user"}, + } + result = _merge_memory(DEFAULT_MEMORY_TEMPLATE, extracted) + assert r"C:\new_user" in result + # The replacement must not introduce an actual newline inside the Name line + for line in result.splitlines(): + if "**Name**" in line: + assert r"C:\new_user" in line + break + else: + pytest.fail("Name line not found") + + def test_backreference_preserved(self): + r"""A value containing '\\1' must not be treated as a backreference.""" + extracted = { + "user_profile": {"role": r"A\1B"}, + } + result = _merge_memory(DEFAULT_MEMORY_TEMPLATE, extracted) + assert r"A\1B" in result + for line in result.splitlines(): + if "**Role**" in line: + assert r"A\1B" in line + break + else: + pytest.fail("Role line not found") + + def test_windows_path_preserved(self): + r"""A Windows-style path must survive without corruption.""" + extracted = { + "research_preferences": { + "preferred_frameworks": r"C:\path\to\file", + }, + } + result = _merge_memory(DEFAULT_MEMORY_TEMPLATE, extracted) + assert r"C:\path\to\file" in result + + def test_multiple_backslash_fields(self): + """Multiple fields with backslashes all survive.""" + extracted = { + "user_profile": { + "name": r"user\name", + "institution": r"MIT\Lab\42", + }, + "research_preferences": { + "hardware": r"GPU\0", + }, + } + result = _merge_memory(DEFAULT_MEMORY_TEMPLATE, extracted) + assert r"user\name" in result + assert r"MIT\Lab\42" in result + assert r"GPU\0" in result + + def test_plain_value_still_works(self): + """Sanity check: normal values without backslashes work fine.""" + extracted = { + "user_profile": {"name": "Alice", "role": "Researcher"}, + } + result = _merge_memory(DEFAULT_MEMORY_TEMPLATE, extracted) + assert "- **Name**: Alice" in result + assert "- **Role**: Researcher" in result diff --git a/tests/test_onboard.py b/tests/test_onboard.py index 95f4b91..66ddbd1 100644 --- a/tests/test_onboard.py +++ b/tests/test_onboard.py @@ -584,6 +584,50 @@ class TestStepChannels: _step_channels(config) +class TestStepMcpServersNpxFailure: + def test_npx_failure_skips_npx_servers(self): + """When _ensure_npx returns False, npx-dependent servers must be skipped.""" + from EvoScientist.onboard import _step_mcp_servers, _RECOMMENDED_MCP_SERVERS + + # Pick server names: one npx-based, one non-npx (URL-based) + npx_name = next(s["name"] for s in _RECOMMENDED_MCP_SERVERS if s.get("command") == "npx") + url_name = next(s["name"] for s in _RECOMMENDED_MCP_SERVERS if "url" in s) + + with mock.patch("EvoScientist.onboard._checkbox_ask", return_value=[npx_name, url_name]), \ + mock.patch("EvoScientist.onboard._ensure_npx", return_value=False), \ + mock.patch("EvoScientist.onboard._check_npx", return_value=False), \ + mock.patch("EvoScientist.mcp.client._load_user_config", return_value={}), \ + mock.patch("EvoScientist.mcp.client.add_mcp_server") as mock_add, \ + mock.patch("EvoScientist.onboard.console"): + result = _step_mcp_servers() + + # The npx server must NOT have been added + added_names = [call.args[0] for call in mock_add.call_args_list] + assert npx_name not in added_names + # The URL server should still be added + assert url_name in added_names + assert url_name in result + assert npx_name not in result + + def test_npx_failure_returns_empty_when_all_npx(self): + """When all selected servers are npx-based and npx fails, return [].""" + from EvoScientist.onboard import _step_mcp_servers, _RECOMMENDED_MCP_SERVERS + + npx_names = [s["name"] for s in _RECOMMENDED_MCP_SERVERS if s.get("command") == "npx"] + assert len(npx_names) >= 1, "Test requires at least one npx server" + + with mock.patch("EvoScientist.onboard._checkbox_ask", return_value=npx_names), \ + mock.patch("EvoScientist.onboard._ensure_npx", return_value=False), \ + mock.patch("EvoScientist.onboard._check_npx", return_value=False), \ + mock.patch("EvoScientist.mcp.client._load_user_config", return_value={}), \ + mock.patch("EvoScientist.mcp.client.add_mcp_server") as mock_add, \ + mock.patch("EvoScientist.onboard.console"): + result = _step_mcp_servers() + + assert result == [] + mock_add.assert_not_called() + + class TestStepParameters: def test_returns_parameters(self): """Test parameters step returns all values.""" diff --git a/tests/test_sessions.py b/tests/test_sessions.py index 45fc5fe..7533fc3 100644 --- a/tests/test_sessions.py +++ b/tests/test_sessions.py @@ -229,6 +229,59 @@ class TestThreadFunctions(unittest.TestCase): # Should not delete OtherAgent's data self.assertFalse(_run(delete_thread("zzz99999"))) + def test_delete_thread_preserves_other_agent_writes(self): + """Deleting a shared thread_id must only remove writes linked to + EvoScientist checkpoints, leaving OtherAgent's writes intact.""" + + shared_tid = "shared01" + + async def _insert(): + import aiosqlite + async with aiosqlite.connect(self._db_path) as conn: + # EvoScientist checkpoint + write + evo_meta = json.dumps({"agent_name": AGENT_NAME, "updated_at": "2025-02-01T00:00:00+00:00"}) + await conn.execute( + "INSERT INTO checkpoints (thread_id, checkpoint_ns, checkpoint_id, metadata) VALUES (?, '', ?, ?)", + (shared_tid, "cp_evo_shared", evo_meta), + ) + await conn.execute( + "INSERT INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, blob) " + "VALUES (?, '', ?, 't1', 0, 'ch', 'str', X'AA')", + (shared_tid, "cp_evo_shared"), + ) + + # OtherAgent checkpoint + write on the SAME thread_id + other_meta = json.dumps({"agent_name": "OtherAgent", "updated_at": "2025-02-01T00:00:00+00:00"}) + await conn.execute( + "INSERT INTO checkpoints (thread_id, checkpoint_ns, checkpoint_id, metadata) VALUES (?, '', ?, ?)", + (shared_tid, "cp_other_shared", other_meta), + ) + await conn.execute( + "INSERT INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, blob) " + "VALUES (?, '', ?, 't2', 0, 'ch', 'str', X'BB')", + (shared_tid, "cp_other_shared"), + ) + await conn.commit() + + _run(_insert()) + + # Delete — should only affect EvoScientist's data + _run(delete_thread(shared_tid)) + + # Verify OtherAgent's writes survive + async def _check(): + import aiosqlite + async with aiosqlite.connect(self._db_path) as conn: + async with conn.execute( + "SELECT checkpoint_id FROM writes WHERE thread_id = ?", (shared_tid,) + ) as cur: + rows = await cur.fetchall() + return [r[0] for r in rows] + + remaining = _run(_check()) + self.assertIn("cp_other_shared", remaining) + self.assertNotIn("cp_evo_shared", remaining) + if __name__ == "__main__": unittest.main()