Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
46 changes: 35 additions & 11 deletions src/surreal_memory/engine/consolidation.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,10 @@
# finished one.
_CHANGE_LOG_COLLAPSE_CAP = 200_000

# A semantic-link manifest can be large, so checkpoint it at bounded apply
# groups instead of rewriting it after every synapse.
_SEMANTIC_LINK_APPLY_CHECKPOINT_BATCH = 50


def _encode_prune_synapse_cursor(synapse: Synapse) -> str:
"""Serialize the stable (created_at, id) prune cursor into progress state."""
Expand Down Expand Up @@ -10692,10 +10696,34 @@ async def checkpoint_discovery(
raw_skipped = metrics.get("skipped_existing")
skipped = int(raw_skipped) if isinstance(raw_skipped, (int, float)) else 0
failures = int(counters.get("semantic_link_failures", 0))
applied_since_checkpoint = 0

async def checkpoint_apply() -> None:
nonlocal applied_since_checkpoint
if applied_since_checkpoint == 0:
return
await self._checkpoint_progress(
"semantic_link_apply",
cursor=cursor,
pending=[serialized],
counters={
"semantic_synapses_created": created,
"semantic_synapses_skipped": skipped,
"semantic_link_failures": failures,
},
)
applied_since_checkpoint = 0

for synapse in synapses:
if cursor is not None and synapse.id <= cursor:
continue
await self._check_progress_budget()
try:
await self._check_progress_budget()
except ConsolidationPausedError:
# The per-row budget guard remains in place, while any writes
# since the last durable cursor are committed before pausing.
await checkpoint_apply()
raise
existing = await self._storage.get_synapse(synapse.id)
if existing is not None:
if (
Expand Down Expand Up @@ -10724,16 +10752,12 @@ async def checkpoint_discovery(
"Semantic synapse write failed (not a duplicate)", exc_info=True
)
cursor = synapse.id
await self._checkpoint_progress(
"semantic_link_apply",
cursor=cursor,
pending=[serialized],
counters={
"semantic_synapses_created": created,
"semantic_synapses_skipped": skipped,
"semantic_link_failures": failures,
},
)
applied_since_checkpoint += 1
if applied_since_checkpoint >= _SEMANTIC_LINK_APPLY_CHECKPOINT_BATCH:
await checkpoint_apply()

# Persist the final partial group before marking the stage complete.
await checkpoint_apply()

await self._checkpoint_progress(
"semantic_link_complete",
Expand Down
154 changes: 153 additions & 1 deletion tests/unit/test_consolidation_mid_strategies_resume.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@
ConsolidationStrategy,
)
from surreal_memory.engine.consolidation_progress import (
ConsolidationLeaseLostError,
ConsolidationPausedError,
ConsolidationProgressError,
)
Expand Down Expand Up @@ -674,6 +675,27 @@ def _semantic_result() -> SemanticDiscoveryResult:
)


def _semantic_batch_result(count: int) -> SemanticDiscoveryResult:
synapses = [
Synapse.create(
f"semantic-source-{index:04d}",
f"semantic-target-{index:04d}",
SynapseType.SIMILAR_TO,
weight=0.54,
metadata={"_semantic_discovery": True},
synapse_id=f"semantic-edge-{index:04d}",
)
for index in range(count)
]
return SemanticDiscoveryResult(
neurons_embedded=count * 2,
pairs_evaluated=count,
synapses_created=count,
eligible_total=count * 2,
synapses=synapses,
)


async def _semantic_discovery(*_args: Any, **_kwargs: Any) -> SemanticDiscoveryResult:
return _semantic_result()

Expand Down Expand Up @@ -714,6 +736,132 @@ async def test_semantic_link_replays_saved_synapse_after_interruption_without_du
assert progress.strategy_state("semantic_link")["pending"] == []


@pytest.mark.asyncio
async def test_semantic_link_batches_apply_checkpoint_and_resumes_partial_group(
monkeypatch: pytest.MonkeyPatch,
) -> None:
result = _semantic_batch_result(123)
storage = _Storage()
progress = _Progress("semantic_link")
monkeypatch.setattr(
"surreal_memory.engine.semantic_discovery.discover_semantic_synapses",
lambda *_args, **_kwargs: _semantic_discovery_result(result),
)

engine = _engine(storage, ConsolidationStrategy.SEMANTIC_LINK, progress)
original_budget_check = engine._check_progress_budget
pause_once = True

async def pause_after_63_writes() -> None:
nonlocal pause_once
if pause_once and len(storage.added_synapses) >= 63:
pause_once = False
raise ConsolidationPausedError("simulated mid-group budget pause")
await original_budget_check()

engine._check_progress_budget = pause_after_63_writes # type: ignore[method-assign]
with pytest.raises(ConsolidationPausedError, match="mid-group budget pause"):
await engine._semantic_link(ConsolidationReport(), dry_run=False)

apply_writes = [write for write in progress.writes if write["phase"] == "semantic_link_apply"]
assert len(apply_writes) == 2
assert apply_writes[-1]["cursor"] == "semantic-edge-0062"
assert apply_writes[-1]["counters"]["semantic_synapses_created"] == 63
checkpoint_sizes = [
len(json.dumps(write["pending"], separators=(",", ":")).encode("utf-8"))
for write in apply_writes
]
assert checkpoint_sizes[0] == checkpoint_sizes[1]
assert max(checkpoint_sizes) < len(result.synapses) * 1000

report = ConsolidationReport()
await _engine(storage, ConsolidationStrategy.SEMANTIC_LINK, progress)._semantic_link(
report, dry_run=False
)

all_apply_writes = [
write for write in progress.writes if write["phase"] == "semantic_link_apply"
]
assert [write["cursor"] for write in all_apply_writes] == [
"semantic-edge-0049",
"semantic-edge-0062",
"semantic-edge-0112",
"semantic-edge-0122",
]
assert len(storage.added_synapses) == 123
assert len(set(storage.added_synapses)) == 123
assert report.semantic_synapses_created == 123


async def _semantic_discovery_result(
result: SemanticDiscoveryResult,
) -> SemanticDiscoveryResult:
return result


@pytest.mark.asyncio
async def test_semantic_link_lease_loss_stops_after_uncheckpointed_apply_group(
monkeypatch: pytest.MonkeyPatch,
) -> None:
result = _semantic_batch_result(50)
storage = _Storage()
progress = _Progress("semantic_link")
monkeypatch.setattr(
"surreal_memory.engine.semantic_discovery.discover_semantic_synapses",
lambda *_args, **_kwargs: _semantic_discovery_result(result),
)
durable_checkpoint = progress.checkpoint

async def lose_lease_on_apply(strategy: str, phase: str, **kwargs: Any) -> None:
if phase == "semantic_link_apply":
raise ConsolidationLeaseLostError("simulated lease loss")
await durable_checkpoint(strategy, phase, **kwargs)

monkeypatch.setattr(progress, "checkpoint", lose_lease_on_apply)
with pytest.raises(ConsolidationLeaseLostError, match="lease loss"):
await _engine(storage, ConsolidationStrategy.SEMANTIC_LINK, progress)._semantic_link(
ConsolidationReport(), dry_run=False
)

assert len(storage.added_synapses) == 50
assert progress.strategy_state("semantic_link")["phase"] == "semantic_link_pending"
assert [write["phase"] for write in progress.writes] == ["semantic_link_pending"]


@pytest.mark.asyncio
async def test_semantic_link_rejects_changed_row_after_apply_checkpoint(
monkeypatch: pytest.MonkeyPatch,
) -> None:
result = _semantic_batch_result(1)
storage = _Storage()
progress = _Progress("semantic_link")
monkeypatch.setattr(
"surreal_memory.engine.semantic_discovery.discover_semantic_synapses",
lambda *_args, **_kwargs: _semantic_discovery_result(result),
)
durable_checkpoint = progress.checkpoint

async def fail_before_apply_checkpoint(strategy: str, phase: str, **kwargs: Any) -> None:
if phase == "semantic_link_apply":
raise OSError("simulated checkpoint failure")
await durable_checkpoint(strategy, phase, **kwargs)

monkeypatch.setattr(progress, "checkpoint", fail_before_apply_checkpoint)
with pytest.raises(OSError, match="checkpoint failure"):
await _engine(storage, ConsolidationStrategy.SEMANTIC_LINK, progress)._semantic_link(
ConsolidationReport(), dry_run=False
)

saved = result.synapses[0]
storage.synapses[saved.id] = dc_replace(saved, weight=saved.weight / 2)
monkeypatch.setattr(progress, "checkpoint", durable_checkpoint)
with pytest.raises(ConsolidationProgressError, match="changed after its checkpoint"):
await _engine(storage, ConsolidationStrategy.SEMANTIC_LINK, progress)._semantic_link(
ConsolidationReport(), dry_run=False
)
assert progress.strategy_state("semantic_link")["phase"] == "semantic_link_pending"


@pytest.mark.asyncio
async def test_semantic_link_resume_accepts_store_public_id_normalization_without_duplicate(
monkeypatch: pytest.MonkeyPatch,
Expand Down Expand Up @@ -854,7 +1002,11 @@ async def test_semantic_link_resumes_after_first_durable_discovery_page(
await _engine(storage, ConsolidationStrategy.SEMANTIC_LINK, progress)._semantic_link(
ConsolidationReport(), dry_run=False
)
assert len(storage.added_synapses) == 1
# The final partial apply group is fully durable before its pause is raised.
assert sorted(storage.added_synapses) == expected_synapse_ids
apply_checkpoint = progress.strategy_state("semantic_link")
assert apply_checkpoint["cursor"] == expected_synapse_ids[-1]
assert apply_checkpoint["counters"]["semantic_synapses_created"] == len(expected_synapse_ids)

await _engine(storage, ConsolidationStrategy.SEMANTIC_LINK, progress)._semantic_link(
ConsolidationReport(), dry_run=False
Expand Down
Loading