1"""Memory consolidator — applies strategies and summarises aged entries."""
2
3from __future__ import annotations
4
5from datetime import UTC, datetime
6from typing import TYPE_CHECKING
7
8from lexigram.ai.memory.config import ConsolidationConfig
9from lexigram.ai.memory.consolidation.strategies import (
10 AccessFrequencyStrategy,
11 DeduplicationStrategy,
12 RecencyDecayStrategy,
13)
14from lexigram.contracts.ai.memory import (
15 ConsolidationResult,
16 MemoryEntry,
17)
18from lexigram.contracts.core.health import HealthCheckResult, HealthStatus
19from lexigram.logging import (
20 get_logger,
21)
22
23if TYPE_CHECKING:
24 from collections.abc import Awaitable, Callable
25
26logger = get_logger(__name__)
27
28
29class MemoryConsolidator:
30 """Orchestrates consolidation of a batch of MemoryEntry objects.
31
32 Applies deduplication, recency decay pruning, and importance-floor
33 pruning in sequence. Optionally runs a summarisation pass on the
34 remaining aged entries.
35 """
36
37 def __init__(
38 self,
39 config: ConsolidationConfig | None = None,
40 summarise_fn: Callable[[list[MemoryEntry]], Awaitable[MemoryEntry]]
41 | None = None,
42 ) -> None:
43 """Initialise the consolidator.
44
45 Args:
46 config: Consolidation thresholds. Defaults to ``ConsolidationConfig()``.
47 summarise_fn: Optional async callable for summarising aged entry groups.
48 """
49 self._config = config or ConsolidationConfig()
50 self._summarise_fn = summarise_fn
51 self._dedup = DeduplicationStrategy()
52 self._recency = RecencyDecayStrategy(
53 half_life_hours=self._config.age_threshold_hours / 2,
54 )
55 self._importance = AccessFrequencyStrategy(
56 importance_threshold=self._config.importance_prune_threshold,
57 )
58
59 async def consolidate(self, entries: list[MemoryEntry]) -> ConsolidationResult:
60 """Consolidate *entries* via deduplication, decay, and importance pruning.
61
62 Args:
63 entries: Entries to process.
64
65 Returns:
66 ConsolidationResult with counts of processed, consolidated, pruned,
67 and extracted entities.
68 """
69 start = datetime.now(UTC)
70 n = len(entries)
71
72 # Step 1: Deduplication
73 unique, dupes = self._dedup.deduplicate(entries)
74 pruned_count = len(dupes)
75
76 # Step 2: Recency decay pruning
77 after_recency, recency_pruned = self._recency.filter(unique)
78 pruned_count += len(recency_pruned)
79
80 # Step 3: Importance pruning
81 final_kept, imp_pruned = self._importance.filter(after_recency)
82 pruned_count += len(imp_pruned)
83
84 # Step 4: Optional summarisation of surviving large batches
85 consolidated_count = 0
86 if self._summarise_fn and len(final_kept) > self._config.batch_size:
87 batches = [
88 final_kept[i : i + self._config.batch_size]
89 for i in range(0, len(final_kept), self._config.batch_size)
90 ]
91 consolidated_count = len(batches)
92 for batch in batches:
93 await self._summarise_fn(batch)
94
95 elapsed = (datetime.now(UTC) - start).total_seconds() * 1000
96 result = ConsolidationResult(
97 entries_processed=n,
98 entries_consolidated=consolidated_count,
99 entries_pruned=pruned_count,
100 entities_extracted=0,
101 duration_ms=elapsed,
102 )
103 logger.info(
104 "consolidation_complete",
105 processed=n,
106 pruned=pruned_count,
107 consolidated=consolidated_count,
108 duration_ms=elapsed,
109 )
110 return result
111
112 async def health_check(self, timeout: float = 5.0) -> HealthCheckResult:
113 """Report consolidator health.
114
115 Args:
116 timeout: Maximum seconds for the health check.
117
118 Returns:
119 HealthCheckResult indicating HEALTHY status.
120 """
121 return HealthCheckResult(
122 component="memory_consolidator",
123 status=HealthStatus.HEALTHY,
124 )
125
126
127__all__ = ["MemoryConsolidator"]