From cf8aab1f34a8650f473e7bec845b2846d96175d3 Mon Sep 17 00:00:00 2001 From: MattMRE <160173637+mattmre@users.noreply.github.com> Date: Thu, 24 Sep 2026 08:45:08 -0400 Subject: [PATCH 1/2] fix(engine): write pre-cycle vectors back when an upsert fails A failed Qdrant chunk stops the sync and upserts the original vectors for ids already stored. The sedimentation logs no longer say the cycle rolled back, and the collision log stays until every chunk succeeds. --- antigravity_engine.py | 69 ++++++++++----------------- sedimentation.py | 14 ++++-- sedimentation_trainer.py | 61 ++++++++++++++++++----- test_sedimentation_corpus_rollback.py | 66 +++++++++++++++++++++++++ 4 files changed, 151 insertions(+), 59 deletions(-) create mode 100644 test_sedimentation_corpus_rollback.py diff --git a/antigravity_engine.py b/antigravity_engine.py index 1ea33741..ade3678d 100644 --- a/antigravity_engine.py +++ b/antigravity_engine.py @@ -1903,6 +1903,8 @@ def _finish_annealing_cycle(): # --- TRAINING LOOP (wrapped in SafeTrainingContext for F-043) --- self.logger.log_training_start(num_samples=len(training_inputs), learning_rate=learning_rate, epochs=epochs, threshold=threshold) + failed_updates = 0 + total_updates = 0 with SafeTrainingContext( self.checkpoint_manager, self.adapter_path, @@ -2066,7 +2068,8 @@ def _finish_annealing_cycle(): # Batch Update Logic using shared helper (F-031: pass cached payload_map) total_updates, failed_updates = sync_vectors_to_qdrant( self.qdrant, self.collection_name, ordered_ids, - new_vectors_np, chunk_size, self.logger, payload_map + new_vectors_np, chunk_size, self.logger, payload_map, + original_vectors_np=input_tensor.detach().cpu().numpy(), ) # Mark success only if no failed vector updates (F-043) @@ -2077,16 +2080,19 @@ def _finish_annealing_cycle(): f"Training cycle completed successfully. Updated {total_updates} vectors.", vectors_updated=total_updates ) + self.chelation_log.clear() else: self.logger.log_error( "sedimentation_partial_failure", - f"Training completed but {failed_updates} vector updates failed. Rolling back.", + f"Training completed but {failed_updates} vector updates failed. " + "Pre-cycle vectors were written back for the chunks that had already been stored.", vectors_updated=total_updates, vectors_failed=failed_updates ) + if failed_updates: + self.adapter.load(self.adapter_path) self.logger.log_training_complete(final_loss=final_loss, vectors_updated=total_updates, vectors_failed=failed_updates) - self.chelation_log.clear() _finish_annealing_cycle() def run_offline_distillation(self, batch_size: int = 100, learning_rate: float = None, epochs: int = None): @@ -2375,50 +2381,25 @@ def run_offline_distillation(self, batch_size: int = 100, learning_rate: float = self.logger.log_event("offline_distillation_update", "Updating corpus vectors in Qdrant") with torch.no_grad(): + original_vectors_np = input_tensor.detach().cpu().numpy() new_vectors_np = self.adapter(input_tensor).numpy() chunk_size = ChelationConfig.CHUNK_SIZE - total_updates = 0 - failed_updates = 0 - - for i in range(0, len(ordered_ids), chunk_size): - chunk_ids = ordered_ids[i:i+chunk_size] - chunk_vectors = new_vectors_np[i:i+chunk_size] - - try: - existing_points = self.qdrant.retrieve( - collection_name=self.collection_name, - ids=chunk_ids, - with_vectors=False - ) - payload_map = {p.id: p.payload for p in existing_points} - - batch_points = [] - for j, doc_id in enumerate(chunk_ids): - vec = chunk_vectors[j].tolist() - pay = payload_map.get(doc_id, {}) - - batch_points.append(PointStruct( - id=doc_id, - vector=vec, - payload=pay - )) - - if batch_points: - self.qdrant.upsert( - collection_name=self.collection_name, - points=batch_points - ) - total_updates += len(batch_points) - except Exception as e: - self.logger.log_error( - "offline_distillation_update_failed", - f"Failed to update batch {i//chunk_size}", - exception=e, - batch_num=i//chunk_size - ) - failed_updates += len(chunk_ids) - + prior_adapter = None + if self.adapter_path.exists(): + prior_adapter = self.adapter_path.read_bytes() + total_updates, failed_updates = sync_vectors_to_qdrant( + self.qdrant, self.collection_name, ordered_ids, + new_vectors_np, chunk_size, self.logger, None, + original_vectors_np=original_vectors_np, + ) + if failed_updates: + if prior_adapter is None: + if self.adapter_path.exists(): + self.adapter_path.unlink() + else: + self.adapter_path.write_bytes(prior_adapter) + self.adapter.load(self.adapter_path) self.logger.log_training_complete( final_loss=final_loss, vectors_updated=total_updates, diff --git a/sedimentation.py b/sedimentation.py index 80289678..5812c341 100644 --- a/sedimentation.py +++ b/sedimentation.py @@ -126,6 +126,9 @@ def run_hierarchical_sedimentation(self, threshold=3, learning_rate=0.001, epoch # --- Phase 1 & 2: Training wrapped in SafeTrainingContext (F-043) --- print(f"Phase 1: Per-cluster training ({len(clusters)} clusters, {cluster_epochs} epochs)...") + failed_updates = 0 + total_updates = 0 + final_loss = 0.0 with SafeTrainingContext( self.checkpoint_manager, self.engine.adapter_path, @@ -189,7 +192,8 @@ def run_hierarchical_sedimentation(self, threshold=3, learning_rate=0.001, epoch # Use shared helper for Qdrant sync (F-031: pass cached payload_map) total_updates, failed_updates = sync_vectors_to_qdrant( self.engine.qdrant, self.engine.collection_name, ordered_ids, - new_vectors_np, chunk_size, self.logger, payload_map + new_vectors_np, chunk_size, self.logger, payload_map, + original_vectors_np=input_tensor.detach().cpu().numpy(), ) # Mark success only if no failed vector updates (F-043) @@ -200,21 +204,23 @@ def run_hierarchical_sedimentation(self, threshold=3, learning_rate=0.001, epoch f"Hierarchical training completed successfully. Updated {total_updates} vectors.", vectors_updated=total_updates ) + self.engine.chelation_log.clear() else: self.logger.log_error( "hierarchical_sedimentation_partial_failure", - f"Training completed but {failed_updates} vector updates failed. Rolling back.", + f"Training completed but {failed_updates} vector updates failed. " + "Pre-cycle vectors were written back for the chunks that had already been stored.", vectors_updated=total_updates, vectors_failed=failed_updates ) + if failed_updates: + self.engine.adapter.load(self.engine.adapter_path) self.logger.log_training_complete( final_loss=final_loss, vectors_updated=total_updates, vectors_failed=failed_updates, ) - - self.engine.chelation_log.clear() print(f"Hierarchical sedimentation complete. Updated {total_updates} vectors, {failed_updates} failed.") print("--- HIERARCHICAL SLEEP CYCLE COMPLETE ---") diff --git a/sedimentation_trainer.py b/sedimentation_trainer.py index 81c0bfaa..c23b91c9 100644 --- a/sedimentation_trainer.py +++ b/sedimentation_trainer.py @@ -68,7 +68,8 @@ def compute_homeostatic_target(current_vec: np.ndarray, noise_vectors: List[np.n def sync_vectors_to_qdrant(qdrant: Any, collection_name: str, ordered_ids: List, new_vectors_np: np.ndarray, chunk_size: int, logger: Any, - payload_map: dict = None) -> tuple: + payload_map: dict = None, + original_vectors_np: np.ndarray = None) -> tuple: """ Synchronize updated vectors to Qdrant in batches, preserving existing payloads. @@ -90,7 +91,47 @@ def sync_vectors_to_qdrant(qdrant: Any, collection_name: str, ordered_ids: List, """ total_updates = 0 failed_updates = 0 - + written_ids = [] + id_to_index = {doc_id: idx for idx, doc_id in enumerate(ordered_ids)} + + def _restore_written() -> bool: + if not written_ids: + return True + if original_vectors_np is None: + logger.log_error( + "corpus_not_restored", + "Corpus was not restored because the pre-cycle vectors were not provided.", + ) + return False + source_payloads = payload_map + if source_payloads is None: + existing = qdrant.retrieve( + collection_name=collection_name, + ids=written_ids, + with_vectors=False, + ) + source_payloads = {point.id: point.payload for point in existing} + points = [] + for doc_id in written_ids: + idx = id_to_index[doc_id] + pay = source_payloads.get(doc_id, {}) or {} + points.append(PointStruct( + id=doc_id, + vector=original_vectors_np[idx].tolist(), + payload=pay, + )) + try: + if points: + qdrant.upsert(collection_name=collection_name, points=points) + return True + except Exception as restore_error: + logger.log_error( + "corpus_not_restored", + "Corpus was not restored.", + exception=restore_error, + ) + return False + for i in range(0, len(ordered_ids), chunk_size): chunk_ids = ordered_ids[i:i + chunk_size] chunk_vectors = new_vectors_np[i:i + chunk_size] @@ -124,15 +165,13 @@ def sync_vectors_to_qdrant(qdrant: Any, collection_name: str, ordered_ids: List, points=batch_points ) total_updates += len(batch_points) - except ValueError as e: - logger.log_error("database_update", - f"Invalid vector data in batch {i//chunk_size}", - exception=e, batch_num=i//chunk_size) - failed_updates += len(chunk_ids) + written_ids.extend(chunk_ids) except Exception as e: - logger.log_error("database_update", - f"Update batch {i//chunk_size} failed", + logger.log_error("database_update", + f"Update batch {i//chunk_size} failed", exception=e, batch_num=i//chunk_size) - failed_updates += len(chunk_ids) - + failed_updates += len(ordered_ids) - len(written_ids) + _restore_written() + return total_updates, failed_updates + return total_updates, failed_updates diff --git a/test_sedimentation_corpus_rollback.py b/test_sedimentation_corpus_rollback.py new file mode 100644 index 00000000..a2770ad9 --- /dev/null +++ b/test_sedimentation_corpus_rollback.py @@ -0,0 +1,66 @@ +"""Partial Qdrant upserts write the pre-cycle vectors back.""" + +import unittest + +import numpy as np + +from sedimentation_trainer import sync_vectors_to_qdrant + + +class _Point: + def __init__(self, point_id, vector, payload): + self.id = point_id + self.vector = vector + self.payload = payload + + +class _Client: + def __init__(self): + self.points = {} + self.upserts = 0 + + def upsert(self, collection_name, points): + self.upserts += 1 + if any(point.id == "b" for point in points): + raise RuntimeError("second chunk failed") + for point in points: + self.points[point.id] = (list(point.vector), dict(point.payload or {})) + + def retrieve(self, collection_name, ids, with_vectors=False): + return [] + + +class _Logger: + def __init__(self): + self.errors = [] + + def log_error(self, kind, message, **kwargs): + self.errors.append((kind, message)) + + +class TestCorpusRollback(unittest.TestCase): + def test_second_chunk_failure_restores_the_first_id(self): + client = _Client() + logger = _Logger() + originals = np.array([[0.25], [0.5]], dtype=np.float32) + adapted = np.array([[9.0], [9.0]], dtype=np.float32) + total, failed = sync_vectors_to_qdrant( + client, + "docs", + ["a", "b"], + adapted, + 1, + logger, + {"a": {"text": "keep"}, "b": {"text": "later"}}, + original_vectors_np=originals, + ) + self.assertEqual(failed, 1) + self.assertEqual(client.points["a"][0], [0.25]) + self.assertEqual(client.points["a"][1], {"text": "keep"}) + self.assertNotIn("b", client.points) + self.assertFalse(any("Rolling back" in message for _, message in logger.errors)) + self.assertGreater(total, 0) + + +if __name__ == "__main__": + unittest.main() From a9d12250cf6f3f879568614d4373bdbc66d3ac53 Mon Sep 17 00:00:00 2001 From: MattMRE <160173637+mattmre@users.noreply.github.com> Date: Thu, 24 Sep 2026 10:38:59 -0400 Subject: [PATCH 2/2] fix(engine): reload the pre-cycle adapter when a corpus upsert fails --- antigravity_engine.py | 70 +++- sedimentation.py | 18 +- sedimentation_trainer.py | 55 ++-- test_antigravity_engine.py | 2 +- test_recursive_decomposer.py | 2 +- test_sedimentation_corpus_rollback.py | 441 +++++++++++++++++++++++++- test_sedimentation_trainer.py | 33 +- 7 files changed, 568 insertions(+), 53 deletions(-) diff --git a/antigravity_engine.py b/antigravity_engine.py index ade3678d..1f4c9f38 100644 --- a/antigravity_engine.py +++ b/antigravity_engine.py @@ -1,6 +1,9 @@ import copy import hashlib +import os +import tempfile import time +from pathlib import Path import numpy as np from qdrant_client import QdrantClient @@ -20,6 +23,40 @@ from embedding_backend import create_embedding_backend from vector_store import create_vector_store + +def _pre_cycle_adapter_bytes(adapter, adapter_path) -> bytes: + """Bytes of the adapter file from before this cycle's trained save. + + When the file is missing, the current in-memory weights are that + pre-cycle state and are saved once so a later overwrite can be undone. + """ + path = Path(adapter_path) + if not path.exists(): + adapter.save(path) + return path.read_bytes() + + +def _write_adapter_bytes(adapter_path, payload: bytes) -> None: + """Replace adapter_path with payload using a same-directory temp file.""" + destination = Path(adapter_path) + descriptor, temporary_name = tempfile.mkstemp( + prefix="." + destination.name + ".", + suffix=".tmp", + dir=destination.parent, + ) + temporary = Path(temporary_name) + try: + with os.fdopen(descriptor, "wb") as handle: + handle.write(payload) + handle.flush() + os.fsync(handle.fileno()) + os.replace(temporary, destination) + except Exception: + if temporary.exists(): + temporary.unlink() + raise + + class AntigravityEngine: def __init__(self, qdrant_location=":memory:", chelation_p=ChelationConfig.DEFAULT_CHELATION_P, model_name='ollama:nomic-embed-text', use_centering=False, use_quantization=False, training_mode: str = "baseline", teacher_model_name: Optional[str] = None, teacher_models=None, teacher_weight: float = 0.5, store_full_text_payload: Optional[bool] = None): """ @@ -1905,6 +1942,11 @@ def _finish_annealing_cycle(): failed_updates = 0 total_updates = 0 + # create_checkpoint copies the file only when it already exists. + # Saving here when it does not records the pre-cycle weights. Saving + # when it does would replace that file with a later in-memory state. + if not Path(self.adapter_path).exists(): + self.adapter.save(self.adapter_path) with SafeTrainingContext( self.checkpoint_manager, self.adapter_path, @@ -2066,7 +2108,7 @@ def _finish_annealing_cycle(): new_vectors_np = self.adapter(input_tensor).numpy() # Batch Update Logic using shared helper (F-031: pass cached payload_map) - total_updates, failed_updates = sync_vectors_to_qdrant( + total_updates, failed_updates, compensated = sync_vectors_to_qdrant( self.qdrant, self.collection_name, ordered_ids, new_vectors_np, chunk_size, self.logger, payload_map, original_vectors_np=input_tensor.detach().cpu().numpy(), @@ -2082,10 +2124,16 @@ def _finish_annealing_cycle(): ) self.chelation_log.clear() else: + failure_message = ( + f"Training completed but {failed_updates} vector updates failed." + ) + if compensated: + failure_message += ( + " Pre-cycle vectors were written back for the chunks that had already been stored." + ) self.logger.log_error( "sedimentation_partial_failure", - f"Training completed but {failed_updates} vector updates failed. " - "Pre-cycle vectors were written back for the chunks that had already been stored.", + failure_message, vectors_updated=total_updates, vectors_failed=failed_updates ) @@ -2255,6 +2303,9 @@ def run_offline_distillation(self, batch_size: int = 100, learning_rate: float = self.adapter.train() final_loss = 0.0 + # Snapshot before either trained save below. Reading the file after + # those saves would restore the weights this cycle just wrote. + prior_adapter = _pre_cycle_adapter_bytes(self.adapter, self.adapter_path) if getattr(self, '_sedimentation_optimizer_type', 'adam') == "eggroll_es": from evolution_strategies_optimizer import ( @@ -2385,21 +2436,14 @@ def run_offline_distillation(self, batch_size: int = 100, learning_rate: float = new_vectors_np = self.adapter(input_tensor).numpy() chunk_size = ChelationConfig.CHUNK_SIZE - prior_adapter = None - if self.adapter_path.exists(): - prior_adapter = self.adapter_path.read_bytes() - total_updates, failed_updates = sync_vectors_to_qdrant( + total_updates, failed_updates, _compensated = sync_vectors_to_qdrant( self.qdrant, self.collection_name, ordered_ids, new_vectors_np, chunk_size, self.logger, None, original_vectors_np=original_vectors_np, ) if failed_updates: - if prior_adapter is None: - if self.adapter_path.exists(): - self.adapter_path.unlink() - else: - self.adapter_path.write_bytes(prior_adapter) - self.adapter.load(self.adapter_path) + _write_adapter_bytes(self.adapter_path, prior_adapter) + self.adapter.load(self.adapter_path) self.logger.log_training_complete( final_loss=final_loss, vectors_updated=total_updates, diff --git a/sedimentation.py b/sedimentation.py index 5812c341..6a460c8e 100644 --- a/sedimentation.py +++ b/sedimentation.py @@ -12,6 +12,7 @@ import numpy as np import torch import torch.optim as optim +from pathlib import Path from typing import List, Tuple from chelation_logger import get_logger @@ -129,6 +130,11 @@ def run_hierarchical_sedimentation(self, threshold=3, learning_rate=0.001, epoch failed_updates = 0 total_updates = 0 final_loss = 0.0 + # create_checkpoint copies the file only when it already exists. + # Saving here when it does not records the pre-cycle weights. Saving + # when it does would replace that file with a later in-memory state. + if not Path(self.engine.adapter_path).exists(): + self.engine.adapter.save(self.engine.adapter_path) with SafeTrainingContext( self.checkpoint_manager, self.engine.adapter_path, @@ -190,7 +196,7 @@ def run_hierarchical_sedimentation(self, threshold=3, learning_rate=0.001, epoch new_vectors_np = self.engine.adapter(input_tensor).numpy() # Use shared helper for Qdrant sync (F-031: pass cached payload_map) - total_updates, failed_updates = sync_vectors_to_qdrant( + total_updates, failed_updates, compensated = sync_vectors_to_qdrant( self.engine.qdrant, self.engine.collection_name, ordered_ids, new_vectors_np, chunk_size, self.logger, payload_map, original_vectors_np=input_tensor.detach().cpu().numpy(), @@ -206,10 +212,16 @@ def run_hierarchical_sedimentation(self, threshold=3, learning_rate=0.001, epoch ) self.engine.chelation_log.clear() else: + failure_message = ( + f"Training completed but {failed_updates} vector updates failed." + ) + if compensated: + failure_message += ( + " Pre-cycle vectors were written back for the chunks that had already been stored." + ) self.logger.log_error( "hierarchical_sedimentation_partial_failure", - f"Training completed but {failed_updates} vector updates failed. " - "Pre-cycle vectors were written back for the chunks that had already been stored.", + failure_message, vectors_updated=total_updates, vectors_failed=failed_updates ) diff --git a/sedimentation_trainer.py b/sedimentation_trainer.py index c23b91c9..39542ab1 100644 --- a/sedimentation_trainer.py +++ b/sedimentation_trainer.py @@ -87,7 +87,12 @@ def sync_vectors_to_qdrant(qdrant: Any, collection_name: str, ordered_ids: List, skips qdrant.retrieve for payload lookup (F-031 optimization) Returns: - Tuple of (total_updates, failed_updates) counts + Tuple of (total_updates, failed_updates, compensated). + compensated is True only when a chunk failed, at least one id from + this call was already written, and the compensating upsert returned. + It is False when nothing failed, when nothing was written, when the + pre-cycle vectors were not provided, or when the compensation + retrieve or upsert raised. """ total_updates = 0 failed_updates = 0 @@ -103,26 +108,29 @@ def _restore_written() -> bool: "Corpus was not restored because the pre-cycle vectors were not provided.", ) return False - source_payloads = payload_map - if source_payloads is None: - existing = qdrant.retrieve( - collection_name=collection_name, - ids=written_ids, - with_vectors=False, - ) - source_payloads = {point.id: point.payload for point in existing} - points = [] - for doc_id in written_ids: - idx = id_to_index[doc_id] - pay = source_payloads.get(doc_id, {}) or {} - points.append(PointStruct( - id=doc_id, - vector=original_vectors_np[idx].tolist(), - payload=pay, - )) try: - if points: - qdrant.upsert(collection_name=collection_name, points=points) + # Payload retrieve stays in this try. An exception here must not + # escape, or the caller would skip the pre-cycle adapter reload. + source_payloads = payload_map + if source_payloads is None: + existing = qdrant.retrieve( + collection_name=collection_name, + ids=written_ids, + with_vectors=False, + ) + source_payloads = {point.id: point.payload for point in existing} + points = [] + for doc_id in written_ids: + idx = id_to_index[doc_id] + pay = source_payloads.get(doc_id, {}) or {} + points.append(PointStruct( + id=doc_id, + vector=original_vectors_np[idx].tolist(), + payload=pay, + )) + if not points: + return False + qdrant.upsert(collection_name=collection_name, points=points) return True except Exception as restore_error: logger.log_error( @@ -171,7 +179,8 @@ def _restore_written() -> bool: f"Update batch {i//chunk_size} failed", exception=e, batch_num=i//chunk_size) failed_updates += len(ordered_ids) - len(written_ids) - _restore_written() - return total_updates, failed_updates + restored = _restore_written() + compensated = bool(written_ids) and restored + return total_updates, failed_updates, compensated - return total_updates, failed_updates + return total_updates, failed_updates, False diff --git a/test_antigravity_engine.py b/test_antigravity_engine.py index 8087b325..32b15ceb 100644 --- a/test_antigravity_engine.py +++ b/test_antigravity_engine.py @@ -344,7 +344,7 @@ def test_sedimentation_uses_safe_training_context(self): mock_stc_cls.return_value = mock_stc # Mock sync to return success (no failures) - mock_sync.return_value = (5, 0) + mock_sync.return_value = (5, 0, False) engine = self._make_engine() train_param = torch.nn.Parameter(torch.tensor(1.0)) diff --git a/test_recursive_decomposer.py b/test_recursive_decomposer.py index 48f1b451..4e0021a5 100644 --- a/test_recursive_decomposer.py +++ b/test_recursive_decomposer.py @@ -805,7 +805,7 @@ def test_hierarchical_sedimentation_uses_safe_training_context(self): mock_stc_cls.return_value = mock_stc # Mock sync to return success (no failures) - mock_sync.return_value = (10, 0) + mock_sync.return_value = (10, 0, False) # Create mock engine with required attributes mock_engine = MagicMock() diff --git a/test_sedimentation_corpus_rollback.py b/test_sedimentation_corpus_rollback.py index a2770ad9..50d0f96d 100644 --- a/test_sedimentation_corpus_rollback.py +++ b/test_sedimentation_corpus_rollback.py @@ -1,11 +1,19 @@ """Partial Qdrant upserts write the pre-cycle vectors back.""" +import os +import tempfile import unittest +from pathlib import Path import numpy as np from sedimentation_trainer import sync_vectors_to_qdrant +try: + import torch +except ImportError: + torch = None + class _Point: def __init__(self, point_id, vector, payload): @@ -30,6 +38,37 @@ def retrieve(self, collection_name, ids, with_vectors=False): return [] +class _ScriptedClient: + """Succeeds a fixed number of upserts, then raises.""" + + def __init__(self, succeed_upserts, fail_retrieve_after_upserts=None): + self.succeed_upserts = succeed_upserts + self.fail_retrieve_after_upserts = fail_retrieve_after_upserts + self.upserts = 0 + self.points = {} + + def upsert(self, collection_name, points): + self.upserts += 1 + if self.upserts > self.succeed_upserts: + raise RuntimeError("upsert failed") + for point in points: + self.points[point.id] = (list(point.vector), dict(point.payload or {})) + + def retrieve(self, collection_name, ids, with_vectors=False): + if ( + self.fail_retrieve_after_upserts is not None + and self.upserts >= self.fail_retrieve_after_upserts + ): + raise RuntimeError("injected retrieve failure during compensation") + found = [] + for doc_id in ids: + if doc_id not in self.points: + continue + vector, payload = self.points[doc_id] + found.append(_Point(doc_id, vector, payload)) + return found + + class _Logger: def __init__(self): self.errors = [] @@ -37,6 +76,21 @@ def __init__(self): def log_error(self, kind, message, **kwargs): self.errors.append((kind, message)) + def log_event(self, *args, **kwargs): + return None + + def log_training_start(self, *args, **kwargs): + return None + + def log_training_epoch(self, *args, **kwargs): + return None + + def log_training_complete(self, *args, **kwargs): + return None + + def log_checkpoint(self, *args, **kwargs): + return None + class TestCorpusRollback(unittest.TestCase): def test_second_chunk_failure_restores_the_first_id(self): @@ -44,7 +98,7 @@ def test_second_chunk_failure_restores_the_first_id(self): logger = _Logger() originals = np.array([[0.25], [0.5]], dtype=np.float32) adapted = np.array([[9.0], [9.0]], dtype=np.float32) - total, failed = sync_vectors_to_qdrant( + total, failed, compensated = sync_vectors_to_qdrant( client, "docs", ["a", "b"], @@ -55,12 +109,397 @@ def test_second_chunk_failure_restores_the_first_id(self): original_vectors_np=originals, ) self.assertEqual(failed, 1) + self.assertTrue(compensated) self.assertEqual(client.points["a"][0], [0.25]) self.assertEqual(client.points["a"][1], {"text": "keep"}) self.assertNotIn("b", client.points) self.assertFalse(any("Rolling back" in message for _, message in logger.errors)) self.assertGreater(total, 0) + def test_compensating_upsert_leaves_the_adapted_vector(self): + client = _ScriptedClient(succeed_upserts=1) + logger = _Logger() + originals = np.array([[0.25], [0.5]], dtype=np.float32) + adapted = np.array([[9.0], [9.0]], dtype=np.float32) + total, failed, compensated = sync_vectors_to_qdrant( + client, + "docs", + ["a", "b"], + adapted, + 1, + logger, + {"a": {"text": "keep"}, "b": {"text": "later"}}, + original_vectors_np=originals, + ) + self.assertEqual(failed, 1) + self.assertGreater(total, 0) + self.assertFalse(compensated) + self.assertEqual(client.points["a"][0], [9.0]) + self.assertNotIn("b", client.points) + messages = [message for _, message in logger.errors] + self.assertIn("Corpus was not restored.", messages) + self.assertFalse(any("written back" in message for message in messages)) + self.assertFalse(any("Rolling back" in message for message in messages)) + + def test_compensation_retrieve_failure_does_not_escape(self): + client = _ScriptedClient(succeed_upserts=1, fail_retrieve_after_upserts=2) + logger = _Logger() + originals = np.array([[0.25], [0.5]], dtype=np.float32) + adapted = np.array([[9.0], [9.0]], dtype=np.float32) + total, failed, compensated = sync_vectors_to_qdrant( + client, + "docs", + ["a", "b"], + adapted, + 1, + logger, + None, + original_vectors_np=originals, + ) + self.assertEqual(failed, 1) + self.assertGreater(total, 0) + self.assertFalse(compensated) + self.assertEqual(client.points["a"][0], [9.0]) + messages = [message for _, message in logger.errors] + self.assertIn("Corpus was not restored.", messages) + self.assertFalse(any("written back" in message for message in messages)) + + +class _StoredPoint: + def __init__(self, point_id, vector, payload): + self.id = point_id + self.vector = [float(value) for value in vector] + self.payload = dict(payload or {}) + + +class _Corpus: + def __init__(self, rows, adapter_path, fail_on_id="b", fail_compensation=False, + fail_compensation_retrieve=False): + self.points = { + row_id: _StoredPoint(row_id, vector, payload) + for row_id, vector, payload in rows + } + self.adapter_path = adapter_path + self.fail_on_id = fail_on_id + self.fail_compensation = fail_compensation + self.fail_compensation_retrieve = fail_compensation_retrieve + self.second_failed = False + self.trained_adapter_bytes = None + self.first_written = {} + self.upsert_ids = [] + + def scroll(self, collection_name, limit, with_vectors, with_payload, offset): + return list(self.points.values()), None + + def retrieve(self, collection_name, ids, with_vectors=False): + if self.fail_compensation_retrieve and self.second_failed and not with_vectors: + raise RuntimeError("injected retrieve failure during compensation") + return [self.points[doc_id] for doc_id in ids] + + def upsert(self, collection_name, points): + if self.trained_adapter_bytes is None: + self.trained_adapter_bytes = Path(self.adapter_path).read_bytes() + ids = [point.id for point in points] + self.upsert_ids.append(ids) + if self.fail_on_id is not None and any(point.id == self.fail_on_id for point in points): + self.second_failed = True + raise RuntimeError("second chunk failed") + if self.fail_compensation and self.second_failed: + raise RuntimeError("compensation failed") + for point in points: + vector = [float(value) for value in point.vector] + if point.id not in self.first_written: + self.first_written[point.id] = vector + self.points[point.id] = _StoredPoint(point.id, vector, point.payload) + + +def _clone_state(adapter): + return {key: value.detach().cpu().clone() for key, value in adapter.state_dict().items()} + + +def _assert_adapter_state(test, adapter, expected): + current = adapter.state_dict() + test.assertEqual(set(current), set(expected)) + for key, value in expected.items(): + test.assertTrue(torch.equal(current[key].detach().cpu(), value), key) + + +def _assert_file_state(test, path, expected): + loaded = torch.load(path, weights_only=True) + test.assertEqual(set(loaded), set(expected)) + for key, value in expected.items(): + test.assertTrue(torch.equal(loaded[key].detach().cpu(), value), key) + + +class _Teacher: + _projection_enabled = False + _projection = None + + def check_dimension_compatibility(self, student_dim): + return student_dim == 4 + + def generate_distillation_targets(self, texts, current_embeddings, teacher_weight=1.0): + current = np.asarray(current_embeddings, dtype=np.float32) + target = np.zeros_like(current) + target[:, -1] = 1.0 + return target + + +@unittest.skipUnless(torch is not None, "adapter reload tests need torch") +class TestProductionAdapterRollback(unittest.TestCase): + DIM = 4 + + def setUp(self): + from config import ChelationConfig + self._chunk_size = ChelationConfig.CHUNK_SIZE + ChelationConfig.CHUNK_SIZE = 1 + self._tmp = tempfile.TemporaryDirectory() + self.tmp = Path(self._tmp.name) + + def tearDown(self): + from config import ChelationConfig + ChelationConfig.CHUNK_SIZE = self._chunk_size + self._tmp.cleanup() + + def _basis(self, index): + values = [0.0] * self.DIM + values[index] = 1.0 + return values + + def _rows(self): + payload = {"a": {"text": "alpha"}, "b": {"text": "beta"}} + return [ + (doc_id, self._basis(0), payload[doc_id]) + for doc_id in ("a", "b") + ] + + def _new_adapter(self): + from chelation_adapter import ChelationAdapter + torch.manual_seed(0) + return ChelationAdapter(self.DIM) + + def _run_sedimentation(self, client, adapter, path, logger): + from types import SimpleNamespace + from antigravity_engine import AntigravityEngine + from checkpoint_manager import CheckpointManager + noise = np.array(self._basis(1), dtype=np.float32) + chelation_log = { + "a": [noise.copy()], + "b": [noise.copy()], + } + engine = SimpleNamespace( + training_mode="baseline", + teacher_helper=None, + chelation_log=chelation_log, + logger=logger, + qdrant=client, + collection_name="docs", + adapter=adapter, + adapter_path=path, + checkpoint_manager=CheckpointManager(self.tmp / "checkpoints"), + _stability_tracker=None, + _sedimentation_optimizer_type="adam", + _sedimentation_loss_type="mse", + _sedimentation_loss_kwargs={}, + _convergence_enabled=False, + _kalman_lr_enabled=False, + _weight_scheduler=None, + _annealing_controller=None, + ) + AntigravityEngine.run_sedimentation_cycle( + engine, + threshold=1, + learning_rate=1.0, + epochs=5, + noise_injection=0.0, + ) + return chelation_log + + def _assert_failure_keeps_pre_cycle_weights(self, adapter, path, pre_state, trained_bytes): + self.assertNotEqual(path.read_bytes(), trained_bytes) + _assert_file_state(self, path, pre_state) + _assert_adapter_state(self, adapter, pre_state) + + def test_failed_sedimentation_restores_prefix_and_keeps_log(self): + path = self.tmp / "adapter.pt" + adapter = self._new_adapter() + pre_state = _clone_state(adapter) + logger = _Logger() + client = _Corpus(self._rows(), path) + chelation_log = self._run_sedimentation(client, adapter, path, logger) + messages = [message for _, message in logger.errors] + self.assertIn("a", chelation_log) + self.assertIn("b", chelation_log) + self.assertTrue(any("written back" in message for message in messages)) + self.assertFalse(any("Rolling back" in message for message in messages)) + self.assertIn("sedimentation_partial_failure", [kind for kind, _ in logger.errors]) + np.testing.assert_allclose(client.points["a"].vector, self._basis(0), atol=1e-5) + np.testing.assert_allclose(client.points["b"].vector, self._basis(0), atol=1e-5) + self._assert_failure_keeps_pre_cycle_weights( + adapter, path, pre_state, client.trained_adapter_bytes + ) + + def test_compensation_failure_does_not_claim_writeback(self): + path = self.tmp / "adapter.pt" + adapter = self._new_adapter() + pre_state = _clone_state(adapter) + logger = _Logger() + client = _Corpus(self._rows(), path, fail_compensation=True) + chelation_log = self._run_sedimentation(client, adapter, path, logger) + messages = [message for _, message in logger.errors] + kinds = [kind for kind, _ in logger.errors] + self.assertIn("a", chelation_log) + self.assertIn("corpus_not_restored", kinds) + self.assertIn("sedimentation_partial_failure", kinds) + self.assertIn("Corpus was not restored.", messages) + self.assertFalse(any("written back" in message for message in messages)) + self.assertFalse(any("Rolling back" in message for message in messages)) + self.assertEqual(client.points["a"].vector, client.first_written["a"]) + self.assertFalse(np.allclose(client.points["a"].vector, self._basis(0), atol=1e-5)) + self._assert_failure_keeps_pre_cycle_weights( + adapter, path, pre_state, client.trained_adapter_bytes + ) + + def test_successful_sedimentation_clears_chelation_log(self): + path = self.tmp / "adapter.pt" + adapter = self._new_adapter() + pre_state = _clone_state(adapter) + logger = _Logger() + client = _Corpus(self._rows(), path, fail_on_id=None) + chelation_log = self._run_sedimentation(client, adapter, path, logger) + self.assertEqual(chelation_log, {}) + self.assertFalse(logger.errors) + self.assertEqual(client.upsert_ids, [["a"], ["b"]]) + self.assertEqual(path.read_bytes(), client.trained_adapter_bytes) + loaded = torch.load(path, weights_only=True) + self.assertFalse(torch.equal( + loaded["correction_net.0.weight"].detach().cpu(), + pre_state["correction_net.0.weight"], + )) + + def test_existing_adapter_file_is_not_overwritten_before_checkpoint(self): + path = self.tmp / "adapter.pt" + adapter = self._new_adapter() + adapter.save(path) + pre_bytes = path.read_bytes() + with torch.no_grad(): + adapter.correction_net[0].weight.add_(1.0) + logger = _Logger() + client = _Corpus(self._rows(), path, fail_compensation=True) + self._run_sedimentation(client, adapter, path, logger) + self.assertEqual(path.read_bytes(), pre_bytes) + _assert_adapter_state(self, adapter, torch.load(path, weights_only=True)) + self.assertFalse(any("written back" in message for _, message in logger.errors)) + + def test_hierarchical_without_adapter_file_reloads_pre_step_weights(self): + from checkpoint_manager import CheckpointManager + from sedimentation import HierarchicalSedimentationEngine + path = self.tmp / "hierarchical.pt" + adapter = self._new_adapter() + pre_state = _clone_state(adapter) + logger = _Logger() + client = _Corpus(self._rows(), path) + noise = np.array(self._basis(1), dtype=np.float32) + engine = type("Engine", (), {})() + engine.chelation_log = {"a": [noise.copy()], "b": [noise.copy()]} + engine.qdrant = client + engine.collection_name = "docs" + engine.adapter = adapter + engine.adapter_path = path + cwd = os.getcwd() + os.chdir(self.tmp) + try: + hierarchical = HierarchicalSedimentationEngine(engine) + hierarchical.logger = logger + hierarchical.checkpoint_manager = CheckpointManager(self.tmp / "hier-checkpoints") + hierarchical.run_hierarchical_sedimentation( + threshold=1, + learning_rate=1.0, + epochs=4, + ) + finally: + os.chdir(cwd) + self.assertIn("a", engine.chelation_log) + self.assertIn("b", engine.chelation_log) + kinds = [kind for kind, _ in logger.errors] + self.assertIn("hierarchical_sedimentation_partial_failure", kinds) + self.assertFalse(any("written back" in message for _, message in logger.errors)) + self.assertFalse(any("Rolling back" in message for _, message in logger.errors)) + self._assert_failure_keeps_pre_cycle_weights( + adapter, path, pre_state, client.trained_adapter_bytes + ) + + def _run_offline(self, client, adapter, path, logger): + from types import SimpleNamespace + from antigravity_engine import AntigravityEngine + engine = SimpleNamespace( + teacher_helper=_Teacher(), + logger=logger, + training_mode="offline", + vector_size=self.DIM, + qdrant=client, + collection_name="docs", + adapter=adapter, + adapter_path=path, + _sedimentation_optimizer_type="adam", + _sedimentation_loss_type="mse", + _sedimentation_loss_kwargs={}, + _convergence_enabled=False, + _kalman_lr_enabled=False, + _weight_scheduler=None, + ) + AntigravityEngine.run_offline_distillation( + engine, + batch_size=10, + learning_rate=1.0, + epochs=5, + ) + + def test_offline_failed_second_upsert_reloads_existing_file(self): + path = self.tmp / "adapter.pt" + adapter = self._new_adapter() + adapter.save(path) + pre_bytes = path.read_bytes() + pre_state = _clone_state(adapter) + logger = _Logger() + client = _Corpus(self._rows(), path) + self._run_offline(client, adapter, path, logger) + self.assertNotEqual(client.trained_adapter_bytes, pre_bytes) + self.assertEqual(path.read_bytes(), pre_bytes) + _assert_file_state(self, path, pre_state) + _assert_adapter_state(self, adapter, pre_state) + np.testing.assert_allclose(client.points["a"].vector, self._basis(0), atol=1e-5) + + def test_offline_missing_file_reloads_weights_saved_before_training(self): + path = self.tmp / "adapter.pt" + adapter = self._new_adapter() + pre_state = _clone_state(adapter) + logger = _Logger() + client = _Corpus(self._rows(), path) + self.assertFalse(path.exists()) + self._run_offline(client, adapter, path, logger) + self.assertTrue(path.exists()) + self._assert_failure_keeps_pre_cycle_weights( + adapter, path, pre_state, client.trained_adapter_bytes + ) + + def test_offline_compensation_retrieve_failure_reloads_adapter(self): + path = self.tmp / "adapter.pt" + adapter = self._new_adapter() + adapter.save(path) + pre_bytes = path.read_bytes() + pre_state = _clone_state(adapter) + logger = _Logger() + client = _Corpus(self._rows(), path, fail_compensation_retrieve=True) + self._run_offline(client, adapter, path, logger) + messages = [message for _, message in logger.errors] + self.assertIn("Corpus was not restored.", messages) + self.assertFalse(any("written back" in message for message in messages)) + self.assertEqual(path.read_bytes(), pre_bytes) + _assert_adapter_state(self, adapter, pre_state) + self.assertEqual(client.points["a"].vector, client.first_written["a"]) + self.assertFalse(np.allclose(client.points["a"].vector, self._basis(0), atol=1e-5)) + if __name__ == "__main__": unittest.main() diff --git a/test_sedimentation_trainer.py b/test_sedimentation_trainer.py index c2fde62c..2864b7f7 100644 --- a/test_sedimentation_trainer.py +++ b/test_sedimentation_trainer.py @@ -187,13 +187,14 @@ def test_successful_sync(self): ordered_ids = ["doc1"] new_vectors = np.array([[0.1, 0.2, 0.3]]) - total, failed = sync_vectors_to_qdrant( + total, failed, compensated = sync_vectors_to_qdrant( mock_qdrant, "test_collection", ordered_ids, new_vectors, chunk_size=10, logger=mock_logger ) self.assertEqual(total, 1) self.assertEqual(failed, 0) + self.assertFalse(compensated) mock_qdrant.retrieve.assert_called_once() mock_qdrant.upsert.assert_called_once() @@ -221,7 +222,7 @@ def mock_retrieve(collection_name, ids, with_vectors): mock_qdrant.upsert.return_value = None chunk_size = 10 - total, failed = sync_vectors_to_qdrant( + total, failed, compensated = sync_vectors_to_qdrant( mock_qdrant, "test_collection", ordered_ids, new_vectors, chunk_size=chunk_size, logger=mock_logger ) @@ -231,6 +232,7 @@ def mock_retrieve(collection_name, ids, with_vectors): self.assertEqual(mock_qdrant.upsert.call_count, 3) self.assertEqual(total, n_docs) self.assertEqual(failed, 0) + self.assertFalse(compensated) def test_payload_preservation(self): """Test that existing payloads are preserved.""" @@ -272,17 +274,18 @@ def test_value_error_handling(self): ordered_ids = ["doc1", "doc2"] new_vectors = np.array([[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]]) - total, failed = sync_vectors_to_qdrant( + total, failed, compensated = sync_vectors_to_qdrant( mock_qdrant, "test_collection", ordered_ids, new_vectors, chunk_size=10, logger=mock_logger ) self.assertEqual(total, 0) self.assertEqual(failed, 2) + self.assertFalse(compensated) mock_logger.log_error.assert_called_once() call_args = mock_logger.log_error.call_args self.assertEqual(call_args[0][0], "database_update") - self.assertIn("Invalid vector data", call_args[0][1]) + self.assertEqual(call_args[0][1], "Update batch 0 failed") def test_generic_exception_handling(self): """Test that generic exceptions are caught and logged.""" @@ -299,13 +302,14 @@ def test_generic_exception_handling(self): ordered_ids = ["doc1"] new_vectors = np.array([[0.1, 0.2, 0.3]]) - total, failed = sync_vectors_to_qdrant( + total, failed, compensated = sync_vectors_to_qdrant( mock_qdrant, "test_collection", ordered_ids, new_vectors, chunk_size=10, logger=mock_logger ) self.assertEqual(total, 0) self.assertEqual(failed, 1) + self.assertFalse(compensated) mock_logger.log_error.assert_called_once() call_args = mock_logger.log_error.call_args self.assertEqual(call_args[0][0], "database_update") @@ -336,14 +340,17 @@ def side_effect_retrieve(*args, **kwargs): ordered_ids = ["doc1", "doc2"] new_vectors = np.array([[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]]) - total, failed = sync_vectors_to_qdrant( + total, failed, compensated = sync_vectors_to_qdrant( mock_qdrant, "test_collection", ordered_ids, new_vectors, chunk_size=1, logger=mock_logger ) self.assertEqual(total, 1) # First chunk succeeded self.assertEqual(failed, 1) # Second chunk failed - self.assertEqual(mock_logger.log_error.call_count, 1) + self.assertFalse(compensated) + kinds = [call[0][0] for call in mock_logger.log_error.call_args_list] + self.assertIn("database_update", kinds) + self.assertIn("corpus_not_restored", kinds) def test_empty_batch(self): """Test handling of empty batches.""" @@ -353,13 +360,14 @@ def test_empty_batch(self): ordered_ids = [] new_vectors = np.array([]).reshape(0, 384) - total, failed = sync_vectors_to_qdrant( + total, failed, compensated = sync_vectors_to_qdrant( mock_qdrant, "test_collection", ordered_ids, new_vectors, chunk_size=10, logger=mock_logger ) self.assertEqual(total, 0) self.assertEqual(failed, 0) + self.assertFalse(compensated) mock_qdrant.retrieve.assert_not_called() mock_qdrant.upsert.assert_not_called() @@ -382,7 +390,7 @@ def capture_upsert(collection_name, points): upserted_points.extend(points) mock_qdrant.upsert.side_effect = capture_upsert - total, failed = sync_vectors_to_qdrant( + total, failed, compensated = sync_vectors_to_qdrant( mock_qdrant, "test_collection", ordered_ids, new_vectors, chunk_size=10, logger=mock_logger, payload_map=payload_map @@ -395,6 +403,7 @@ def capture_upsert(collection_name, points): self.assertEqual(mock_qdrant.upsert.call_count, 1) self.assertEqual(total, 2) self.assertEqual(failed, 0) + self.assertFalse(compensated) # Verify payloads from provided map were used self.assertEqual(len(upserted_points), 2) @@ -423,7 +432,7 @@ def mock_retrieve(collection_name, ids, with_vectors): new_vectors = np.array([[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]]) # Call WITHOUT payload_map - total, failed = sync_vectors_to_qdrant( + total, failed, compensated = sync_vectors_to_qdrant( mock_qdrant, "test_collection", ordered_ids, new_vectors, chunk_size=10, logger=mock_logger ) @@ -432,6 +441,7 @@ def mock_retrieve(collection_name, ids, with_vectors): mock_qdrant.retrieve.assert_called_once() self.assertEqual(total, 2) self.assertEqual(failed, 0) + self.assertFalse(compensated) def test_payload_map_chunking(self): """Test F-031: payload_map works correctly with multiple chunks.""" @@ -447,7 +457,7 @@ def test_payload_map_chunking(self): mock_qdrant.upsert.return_value = None chunk_size = 10 - total, failed = sync_vectors_to_qdrant( + total, failed, compensated = sync_vectors_to_qdrant( mock_qdrant, "test_collection", ordered_ids, new_vectors, chunk_size=chunk_size, logger=mock_logger, payload_map=payload_map @@ -460,6 +470,7 @@ def test_payload_map_chunking(self): self.assertEqual(mock_qdrant.upsert.call_count, 3) self.assertEqual(total, n_docs) self.assertEqual(failed, 0) + self.assertFalse(compensated) if __name__ == "__main__":