feat(tests): add tests for _merge_memory backslash safety and _step_mcp_servers npx handling
This commit is contained in:
@@ -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
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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."""
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user