feat(tests): add tests for _merge_memory backslash safety and _step_mcp_servers npx handling

This commit is contained in:
X-iZhang
2026-02-12 18:30:27 +00:00
parent a61333d4c5
commit 5271931935
6 changed files with 195 additions and 7 deletions
+2 -4
View File
@@ -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")
+11 -1
View File
@@ -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:
+12 -2
View File
@@ -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
+73
View File
@@ -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
+44
View File
@@ -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."""
+53
View File
@@ -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()