diff --git a/.gitignore b/.gitignore index bee74ac..88caeaa 100644 --- a/.gitignore +++ b/.gitignore @@ -11,6 +11,7 @@ tools/__pycache__/ CLAUDE.md *PLAN.md LESSONS.md +superpowers* # Added by cargo diff --git a/README.md b/README.md index 6c47ba8..c30b4f8 100644 --- a/README.md +++ b/README.md @@ -13,7 +13,9 @@ Engram is written in Rust for performance and reliability. It exposes all operat Engram is built for developers who want to plug in their own LLM agents, run locally or in production, and see exactly what goes into the context window. All memory operations are behind trait abstractions, so you can swap implementations or mock them in tests without changing any calling code. -With the latest update, Engram also supports collective memory: multiple agents can share a global knowledge graph, sessions can be marked public or private, and conflicting facts across sessions are surfaced via a dedicated endpoint. +With the latest update, Engram also supports memory evolution: when a session's short-term message count crosses a threshold, the leader automatically summarizes the oldest messages into a compact, immutable consolidated memory and trims them. State shrinks, meaning is preserved, and every node in the cluster agrees on the result. + +Engram also supports collective memory: multiple agents can share a global knowledge graph, sessions can be marked public or private, and conflicting facts across sessions are surfaced via a dedicated endpoint. ## architecture @@ -30,11 +32,13 @@ graph TD router --> knowledgehandler["knowledge handler"] router --> visibilityhandler["visibility handler"] router --> globalhandler["global knowledge handler"] + router --> consolidationhandler["consolidation handler"] memoryhandler -->|write| raft["raft consensus\nopenraft 0.9\ncluster mode only"] corememhandler -->|write| raft sessionhandler -->|delete or register agent| raft visibilityhandler -->|write| raft + consolidationhandler -->|ApplySummary| raft raft -.->|grpc append entries| peers["peer nodes (port 9001)"] raft -.->|grpc install snapshot| laggers["lagging followers"] raft -->|state machine apply| shortterm["short-term memory trait"] @@ -42,6 +46,7 @@ graph TD raft -->|state machine apply| knowledgegraph[("knowledge graph\nper-session in-memory")] raft -->|state machine apply| globalgraph[("global knowledge graph\ncross-session")] raft -->|state machine apply| visibility["session visibility map"] + raft -->|state machine apply| consolidated[("consolidated memory\nper-session summaries")] raft -->|embedding job| embedqueue["embedding worker pool
bounded channel"] raft -->|knowledge job| knowledgequeue["knowledge worker pool
bounded channel"] raft --> redb[("redb\npersistent raft log\n+ snapshot store")] @@ -56,12 +61,17 @@ graph TD longterm --> lancedb[("lancedb
persistent ann search")] coremem --> redis + consolidated --> redis knowledgequeue -->|leader-only extraction| extractor["knowledge extractor trait"] extractor -->|openai gpt-4o-mini / mock| extraction["entities + relationships"] extraction -->|AddKnowledge via raft| knowledgegraph knowledgegraph -->|public sessions merge| globalgraph + consolidationscheduler["consolidation scheduler\nleader-only worker"] -->|threshold check| shortterm + consolidationscheduler -->|summarize| summarizer["summarizer trait\nopenai gpt-4o-mini / mock"] + summarizer -->|ApplySummary via raft| consolidated + knowledgehandler --> knowledgegraph globalhandler --> globalgraph @@ -193,6 +203,8 @@ docker compose up -d | GET | /knowledge/global/path?from=X&to=Y | find path in global graph | | GET | /knowledge/global/export?format=json\|dot | export global graph | | GET | /knowledge/global/conflicts | list conflicting facts across sessions | +| GET | /sessions/{session_id}/summaries | list consolidated summaries for a session | +| POST | /sessions/{session_id}/consolidate | manually trigger consolidation (leader only) | | GET | /cluster | cluster status (cluster mode only) | | POST | /cluster/init | initialize cluster | | POST | /cluster/add-learner | add a learner node | @@ -218,6 +230,11 @@ the application reads configuration from environment variables: | KNOWLEDGE_EXTRACTOR | knowledge extractor backend (`openai` or `mock`) | openai | | KNOWLEDGE_MAX_WORKERS | number of knowledge extraction workers | 4 | | KNOWLEDGE_CHANNEL_SIZE | knowledge job queue size | 500 | +| SUMMARIZER | summarizer backend (`openai` or `mock`) | openai | +| CONSOLIDATION_THRESHOLD | message count that triggers consolidation | 50 | +| CONSOLIDATION_TARGET_WINDOW| raw messages kept after consolidation | 20 | +| CONSOLIDATION_MAX_WORKERS | number of consolidation workers | 2 | +| CONSOLIDATION_CHANNEL_SIZE | consolidation job queue size | 100 | | RUST_LOG | tracing log filter | info | | LOG_FORMAT | logging format (`pretty` or `json`) | pretty | @@ -258,12 +275,16 @@ values like `similarity_threshold` and `max_tokens` are controlled per request t - follower-to-leader HTTP redirect (307) in cluster mode - per-node LanceDB with eventual consistency via deterministic embeddings - persistent redb-backed Raft log and snapshot store (survives restarts) -- full state machine snapshots covering short-term memory, core memory, knowledge graph, global graph, session visibility, and agent registry +- full state machine snapshots covering short-term memory, core memory, knowledge graph, global graph, session visibility, agent registry, and consolidated summaries (snapshot v3) - startup recovery from the latest snapshot followed by Raft log replay - InstallSnapshot RPC so lagging followers catch up without manual re-initialization - automatic log compaction after every `SNAPSHOT_LOG_THRESHOLD` committed entries - cluster management REST API - OpenAPI docs and Swagger UI +- memory consolidation: leader-only summarization scheduler, `Summarizer` trait (OpenAI GPT-4o-mini or mock), `ConsolidatedMemoryStore` (in-memory and Redis), immutable summaries with message lineage and model metadata +- `ApplySummary` Raft command: one atomic replicated transition stores the summary, trims consumed raw messages, and updates metrics; idempotent by `summary_id` +- `GET /sessions/{id}/summaries` and `POST /sessions/{id}/consolidate` endpoints; followers 307-redirect to the leader +- five new Prometheus metrics for consolidation throughput and queue depth - LongMemEval and BEAM benchmark harnesses ## quickstart (3-node cluster) @@ -284,7 +305,7 @@ docker compose -f docker-compose.cluster.yml up -d --build ./scripts/cluster-verify.sh ``` -the verify script checks 17 criteria: leader election, write replication to all nodes, 307 redirect from followers, failover, Prometheus metric presence, knowledge graph replication, delete-session cleanup, node restart and recovery from the Raft log, snapshot compaction, restart-then-verify that state is fully restored from the latest snapshot, session visibility propagation, global graph population from public sessions, agent registration, global entity and relationship count metrics, global entity queries, global conflict detection, and global graph snapshot round-trip. it exits 0 only if all criteria pass. +the verify script checks 22 criteria: leader election, write replication to all nodes, 307 redirect from followers, failover, Prometheus metric presence, knowledge graph replication, delete-session cleanup, node restart and recovery from the Raft log, snapshot compaction, restart-then-verify that state is fully restored from the latest snapshot, session visibility propagation, global graph population from public sessions, agent registration, global entity and relationship count metrics, global entity queries, global conflict detection, global graph snapshot round-trip, consolidation threshold trigger, replicated determinism (all nodes byte-identical after consolidation), idempotency of re-applied summaries, persistence of summaries through cluster restart, and manual consolidate endpoint. it exits 0 only if all 22 criteria pass. see `docker-compose.cluster.yml` and the scripts in `scripts/` for details. diff --git a/docker-compose.cluster.yml b/docker-compose.cluster.yml index a844759..e1ca642 100644 --- a/docker-compose.cluster.yml +++ b/docker-compose.cluster.yml @@ -29,6 +29,9 @@ services: ENGRAM_BIND_ADDR: "0.0.0.0:3000" OPENAI_API_KEY: "${OPENAI_API_KEY:-}" KNOWLEDGE_EXTRACTOR: "${KNOWLEDGE_EXTRACTOR:-mock}" + SUMMARIZER: "mock" + CONSOLIDATION_THRESHOLD: "5" + CONSOLIDATION_TARGET_WINDOW: "2" RUST_LOG: "info,openraft=debug" ports: - "3000:3000" @@ -54,6 +57,9 @@ services: ENGRAM_BIND_ADDR: "0.0.0.0:3000" OPENAI_API_KEY: "${OPENAI_API_KEY:-}" KNOWLEDGE_EXTRACTOR: "${KNOWLEDGE_EXTRACTOR:-mock}" + SUMMARIZER: "mock" + CONSOLIDATION_THRESHOLD: "5" + CONSOLIDATION_TARGET_WINDOW: "2" RUST_LOG: "info,openraft=debug" ports: - "3001:3000" @@ -79,6 +85,9 @@ services: ENGRAM_BIND_ADDR: "0.0.0.0:3000" OPENAI_API_KEY: "${OPENAI_API_KEY:-}" KNOWLEDGE_EXTRACTOR: "${KNOWLEDGE_EXTRACTOR:-mock}" + SUMMARIZER: "mock" + CONSOLIDATION_THRESHOLD: "5" + CONSOLIDATION_TARGET_WINDOW: "2" RUST_LOG: "info,openraft=debug" ports: - "3002:3000" diff --git a/docs/API.md b/docs/API.md index 6396de5..2b7e051 100644 --- a/docs/API.md +++ b/docs/API.md @@ -21,6 +21,8 @@ This document describes every REST endpoint exposed by Engram. All endpoints are | GET | /knowledge/global/path | find shortest path in the global graph | | GET | /knowledge/global/export | export global graph (JSON or Graphviz DOT) | | GET | /knowledge/global/conflicts | list conflicting facts across sessions | +| GET | /sessions/{session_id}/summaries | list consolidated summaries for a session | +| POST | /sessions/{session_id}/consolidate | manually trigger consolidation (leader only) | | GET | /health | health check | | GET | /metrics | Prometheus metrics | | GET | /api-docs/openapi.json | OpenAPI specification | @@ -730,3 +732,79 @@ Returns all detected conflicts in the global graph. A conflict occurs when two d ```sh curl http://localhost:3000/knowledge/global/conflicts ``` + +--- + +## consolidation endpoints + +These endpoints give access to the consolidated memory produced by the leader's summarization scheduler. When a session's short-term message count exceeds `CONSOLIDATION_THRESHOLD`, the leader summarizes the oldest messages, replicates the result as an `ApplySummary` command, and every node atomically stores the summary and trims the consumed raw messages. + +--- + +## GET /sessions/{session_id}/summaries + +Returns all consolidated summaries for a session, ordered by Raft log index. + +**path parameters:** +- `session_id` (string): session identifier + +**success response:** +- status: 200 +- body: +```json +{ + "session_id": "abc123", + "summaries": [ + { + "id": "11111111-1111-1111-1111-111111111111", + "text": "Alice discussed her role at OpenAI and her preference for Rust.", + "created_at_index": 72, + "consumed_message_ids": ["m1", "m2", "m3"], + "consumed_count": 3, + "model": "gpt-4o-mini", + "prompt_version": "summarize_v1" + } + ] +} +``` + +**error responses:** +- 500: failed to retrieve summaries + +**example:** +```sh +curl http://localhost:3000/sessions/{session_id}/summaries +``` + +--- + +## POST /sessions/{session_id}/consolidate + +Manually triggers consolidation for a session. The leader summarizes all messages beyond `CONSOLIDATION_TARGET_WINDOW`, stores the result, and trims the consumed raw messages. Useful for debugging, verification, and reproducible cluster tests. + +In cluster mode, followers return 307 with a `Location` header pointing to the leader. + +**path parameters:** +- `session_id` (string): session identifier + +**request body:** none + +**success response:** +- status: 202 (accepted) +- body: +```json +{ + "summary_id": "11111111-1111-1111-1111-111111111111" +} +``` + +**error responses:** +- 307: redirect to leader (cluster mode, follower received the request) +- 409: consolidation already in flight for this session +- 422: session has fewer messages than the target window; nothing to consolidate +- 500: summarization or replication failed + +**example:** +```sh +curl -X POST http://localhost:3000/sessions/{session_id}/consolidate +``` diff --git a/docs/ARCHITECTURE.md b/docs/ARCHITECTURE.md index 487015d..a906f30 100644 --- a/docs/ARCHITECTURE.md +++ b/docs/ARCHITECTURE.md @@ -26,11 +26,13 @@ graph TD router --> knowledgehandler["knowledge handler"] router --> visibilityhandler["visibility handler"] router --> globalhandler["global knowledge handler"] + router --> consolidationhandler["consolidation handler"] memoryhandler -->|write| raft["raft consensus\nopenraft 0.9\ncluster mode only"] corememhandler -->|write| raft sessionhandler -->|delete or register agent| raft visibilityhandler -->|write| raft + consolidationhandler -->|ApplySummary| raft raft -.->|grpc append entries| peers["peer nodes (port 9001)"] raft -.->|grpc install snapshot| laggers["lagging followers"] raft -->|state machine apply| shortterm["short-term memory (trait)"] @@ -38,6 +40,7 @@ graph TD raft -->|state machine apply| knowledgegraph[("knowledge graph\nper-session in-memory")] raft -->|state machine apply| globalgraph[("global knowledge graph\ncross-session")] raft -->|state machine apply| visibility["session visibility map"] + raft -->|state machine apply| consolidated[("consolidated memory\nper-session summaries")] raft -->|embedding job| embedqueue["embedding worker pool (bounded channel)"] raft -->|knowledge job| knowledgequeue["knowledge worker pool (bounded channel)"] raft --> redb[("redb\npersistent raft log\n+ snapshot store")] @@ -45,6 +48,7 @@ graph TD shortterm --> redis[("redis")] shortterm --> inmem["in-memory store (test fallback)"] + consolidated --> redis embedqueue -->|generate| embedprovider["embedding provider (trait)"] embedprovider -->|https| openai["openai embedding api"] @@ -59,6 +63,10 @@ graph TD extraction -->|AddKnowledge via raft| knowledgegraph knowledgegraph -->|public sessions merge| globalgraph + consolidationscheduler["consolidation scheduler\nleader-only worker"] -->|threshold check| shortterm + consolidationscheduler -->|leader-only summarize| summarizer["summarizer (trait)\ngpt-4o-mini or mock"] + summarizer -->|ApplySummary via raft| consolidated + knowledgehandler --> knowledgegraph globalhandler --> globalgraph @@ -94,6 +102,10 @@ All major components are behind trait abstractions, which lets implementations b `KnowledgeExtractor` extracts named entities and typed relationships from text. The `OpenAIKnowledgeExtractor` calls GPT-4o-mini with a structured JSON prompt and exponential backoff on 429s. `MockKnowledgeExtractor` uses pattern matching against a fixed set of relationship phrases; this is the default in Docker Compose and CI to avoid OpenAI quota consumption. +`Summarizer` produces a compact third-person summary from a slice of messages. `OpenAISummarizer` calls GPT-4o-mini with a consolidation-specific system prompt and the same exponential backoff pattern as the extractor. `MockSummarizer` is deterministic and offline, suitable for tests and cluster verification. + +`ConsolidatedMemoryStore` holds per-session `Vec` objects. `InMemoryConsolidatedStore` is the test fallback; `RedisConsolidatedStore` is the production implementation. `add_summary` is idempotent by `summary.id`, so replaying an `ApplySummary` command never produces duplicate entries. + ## Design decisions | decision | alternatives | final choice & rationale | @@ -112,6 +124,12 @@ All major components are behind trait abstractions, which lets implementations b | Redis as a projection, not the source of truth | reconcile Redis with Raft log on recovery | flushing Redis unconditionally on startup and restoring from the snapshot avoids a maze of reconciliation edge cases; one authoritative source (Raft log plus latest snapshot) with all volatile state derived from it | | global graph as in-memory projection of public sessions | separate global persistent store | the global graph is always derivable from the full log; snapshotting it avoids replaying the entire history on startup while keeping the system simple | | session visibility defaults to Private | defaults to Shared | private by default prevents unintended cross-agent knowledge leakage; agents opt in explicitly | +| leader-only summarization + replicated result | all nodes summarize | LLM output is non-deterministic; allowing every node to summarize would produce divergent summaries across the cluster; one call on the leader, replicated as `ApplySummary`, keeps all nodes identical | +| `summary_id` is a leader-minted UUID | content hash | content hashes would collide when different prompts or models produce different text for the same messages; idempotency is "have I applied this UUID?" not "have I processed these inputs?" | +| consolidation reduces session back to target window in one pass | fixed-size batches | one summary covering the full overshoot keeps lineage clean, minimizes Raft commands, and avoids repeated mini-consolidations that would each produce their own summary entry | +| summaries ordered by Raft log index | wall-clock timestamp | log index is identical on every node and stable across replay; timestamps drift and skew between nodes | +| trim is atomic inside `apply_cmd` | separate trim command | splitting into two commands would allow a node to have the summary but not the trim (double-counting) or the trim but not the summary (data loss); one command, one transition | +| summaries are immutable once applied | allow regeneration | regenerating a summary with a newer prompt would change the replicated artifact; replay would reproduce a different result, breaking temporal integrity | ## Context assembly algorithm @@ -209,15 +227,16 @@ flowchart TD follower["follower\nhttp :3000 / grpc :9001"] lagging["lagging follower\nhttp :3000 / grpc :9001"] leader["leader\nhttp :3000 / grpc :9001"] - r1["node 1 redis"] - r2["node 2 redis"] - r3["node 3 redis"] + r1["node 1 redis\nshort-term + core + consolidated"] + r2["node 2 redis\nshort-term + core + consolidated"] + r3["node 3 redis\nshort-term + core + consolidated"] l1["node 1 lancedb"] l2["node 2 lancedb"] l3["node 3 lancedb"] db1["node 1 redb\nraft log + snapshots"] db2["node 2 redb\nraft log + snapshots"] db3["node 3 redb\nraft log + snapshots"] + sched["consolidation scheduler\nleader-only"] client -->|"write to follower"| follower follower -->|"307 redirect"| client @@ -225,13 +244,14 @@ flowchart TD leader -->|"append entries (grpc)"| r1 leader -->|"append entries (grpc)"| r2 leader -->|"install snapshot (grpc)"| lagging - lagging -->|"restore from snapshot"| r3 + lagging -->|"restore from snapshot (v3)"| r3 lagging --> db3 r1 -.->|"embed async"| l1 r2 -.->|"embed async"| l2 r3 -.->|"embed async"| l3 leader --> db1 r2 -.-> db2 + sched -->|"ApplySummary via raft"| leader ``` The cluster also exposes management endpoints at `/cluster`, `/cluster/init`, `/cluster/add-learner`, and `/cluster/change-membership`. @@ -252,9 +272,9 @@ OpenAI text embeddings are deterministic for the same input. All nodes converge | `EngRaftNetworkConnection` | sends Vote, AppendEntries, and InstallSnapshot RPCs over gRPC; the snapshot payload is a serialized `EngramSnapshot` | | `RaftGrpcServer` | tonic service that forwards incoming Raft RPCs to the local `RaftHandle`; handles `InstallSnapshotRequest` so lagging followers can restore a leader snapshot over gRPC | -`MemoryCommand` has six variants: `AddMessage`, `AddFact`, `DeleteSession`, `AddKnowledge`, `SetSessionVisibility`, and `RegisterSession`. `AddMessage` also enqueues a `KnowledgeJob` so every committed message is a candidate for knowledge extraction. +`MemoryCommand` has seven variants: `AddMessage`, `AddFact`, `DeleteSession`, `AddKnowledge`, `SetSessionVisibility`, `RegisterSession`, and `ApplySummary`. `AddMessage` also enqueues a `KnowledgeJob` so every committed message is a candidate for knowledge extraction. `ApplySummary` carries `session_id`, `summary_id` (leader-minted UUID), `summary_text`, `consumed_message_ids`, `model`, and `prompt_version`; it is idempotent by `summary_id`. -`EngramSnapshot` is the versioned payload serialized into every snapshot. It contains `short_term`, `core_memory`, `knowledge_graph`, `global_graph`, `visibility`, and `session_agents`. The `version: 1` field and `#[serde(default)]` on optional fields mean older nodes can install newer snapshots by ignoring fields they don't recognize. +`EngramSnapshot` is the versioned payload serialized into every snapshot. It contains `short_term`, `core_memory`, `knowledge_graph`, `global_graph`, `visibility`, `session_agents`, and `consolidated`. The current version is 3. The `#[serde(default)]` on all optional fields means older snapshots (v1 or v2) deserialize cleanly on newer nodes. `recover_state_machine()` in `src/raft/recovery.rs` runs at node startup before Raft is initialized. It flushes Redis, loads the latest persisted snapshot from redb, restores the payload into the live stores, and advances `last_applied` and `last_membership`. OpenRaft then replays any committed log entries that sit past the snapshot index. @@ -371,10 +391,66 @@ Three Prometheus gauges track the state of the global graph. They are updated ea The `EngramSnapshot` payload now includes `global_graph`, `visibility`, and `session_agents` fields alongside the existing `short_term`, `core_memory`, and `knowledge_graph` fields. The `#[serde(default)]` attribute on all new fields means Stage 3A snapshots deserialize cleanly on Stage 3B nodes. +## Stage 4: Memory Evolution + +Stage 4 makes memory active. When a session's short-term message count crosses `CONSOLIDATION_THRESHOLD` (default 50), the leader automatically summarizes the oldest messages and trims them, driving the session back to exactly `CONSOLIDATION_TARGET_WINDOW` (default 20) raw messages. The summary is an immutable artifact; the raw messages it consumed are gone. This is the first stage where Engram intentionally destroys source material. + +### Why leader-only + replicated result + +Summarization is non-deterministic. Two calls to GPT-4o-mini on the same transcript produce different text. Allowing every node to summarize independently would produce a different `ApplySummary` command on each node and diverge the cluster. The solution is identical to Stage 2's knowledge extraction: only the current leader calls the `Summarizer`. The leader mints a UUID `summary_id`, submits `MemoryCommand::ApplySummary`, and OpenRaft replicates the artifact to all nodes. Followers apply the identical text. Replay is deterministic because the leader's output is captured once into the log. + +### The consolidation scheduler + +A leader-only worker loop (mirroring `knowledge/worker.rs`) polls sessions, checks whether any exceed the threshold, and enqueues consolidation jobs. A leader-local `HashSet` prevents launching a second job for a session that already has one in flight. The guard is not replicated; it only prevents duplicate work, not divergent state. When a leadership transfer happens, the new leader's scheduler starts fresh with an empty guard. + +### The `ApplySummary` state transition + +`apply_cmd` for `ApplySummary` performs three things atomically under the `SmInner` lock: + +1. `consolidated.add_summary(session_id, summary)` (idempotent by `summary_id`; replaying the command is a no-op) +2. `short_term.remove_messages(session_id, consumed_message_ids)` (trims exactly the messages the summary replaced) +3. Metrics update (increments consolidation counters, resets the summary gauge) + +There is no side channel between these steps. A node either has the summary and the trim or it has neither. + +### Summary ordering + +Summaries are ordered by `created_at_index` (the Raft log index of the committing entry). This index is identical on every node and stable across replay, so ordering is immune to clock skew. Anything displaying or sorting summaries must use this field. + +### Known limitation: vector cleanup is deferred + +Trimming raw messages from short-term memory does not immediately delete their vectors from LanceDB. Vectors are per-node, derived, and eventually consistent by Stage 1 design. A stale vector for a trimmed message may affect search quality but not correctness. A `EmbeddingJob::DeleteMessages` variant is the clean follow-up. + +### Consolidation REST endpoints + +| method | path | description | +|--------|------|-------------| +| GET | `/sessions/{id}/summaries` | list all consolidated summaries for the session, ordered by log index | +| POST | `/sessions/{id}/consolidate` | manually trigger consolidation; leader summarizes on demand; followers 307-redirect | + +### Consolidation metrics + +| metric | type | description | +|--------|------|-------------| +| `engram_consolidations_total` | counter | number of consolidation operations completed | +| `engram_messages_consolidated_total` | counter | cumulative raw messages trimmed by consolidation | +| `engram_summaries` | gauge | current number of stored summaries | +| `engram_consolidation_queue_size` | gauge | pending jobs in the consolidation channel | +| `engram_summarization_duration_seconds` | histogram | wall-clock time per summarization call (label: `model`) | + +### Snapshot protocol v3 + +`EngramSnapshot` gains a `consolidated: Vec<(String, Vec)>` field with `#[serde(default)]`. v2 snapshots (without the field) deserialize with an empty consolidated map. + ## Deferred items -The following remain out of scope after Stage 3B: +The following remain out of scope after Stage 4: - **LanceDB replication.** Each node still calls the embedding API independently. Routing vector storage through Raft (leader-only embedding, follower payload replication) is a future item. - **KG-augmented context assembly.** The knowledge graph (per-session and global) is queryable via REST but is not yet integrated into the context assembly pipeline to augment semantic search results. +- **Summary-aware retrieval.** `GET /sessions/{id}/context` still reads from raw short-term messages only. Wiring consolidated summaries into context assembly is a later stage. +- **Summary-of-summary hierarchies.** A session accumulates a flat list of summaries. Recursive compaction (summarizing summaries) is deferred; `consumed_message_ids` lineage is already captured so that stage is additive. +- **Knowledge-graph consolidation.** Entity merge, node pruning, and relationship collapse in the global or per-session graph are a separate capability. +- **Forgetting and decay.** Importance scoring and eviction policies are deferred. +- **Per-message vector deletion.** Trimmed messages' vectors remain in LanceDB until their session is deleted. - **Multi-tenant auth.** Cluster-aware authentication routing. diff --git a/docs/COMPARISON.md b/docs/COMPARISON.md index be4a195..967c202 100644 --- a/docs/COMPARISON.md +++ b/docs/COMPARISON.md @@ -6,12 +6,12 @@ |------------------------------- |:--------------:|:-------------:|:-------------:|:----------------:|:---------------------:| | **Language** | Rust | Python | Python | Python | Go | | **Deployment model** | Single binary, Docker | Docker, cloud, pip | pip, Docker, cloud | pip, cloud | Docker, cloud | -| **Fault-tolerant cluster** | Yes (3-node Raft, OpenRaft 0.9, persistent log + snapshots, startup recovery, snapshot v2 with global state) | No | No | No | ? | +| **Fault-tolerant cluster** | Yes (3-node Raft, OpenRaft 0.9, persistent log + snapshots, startup recovery, snapshot v3 with global state and consolidated summaries) | No | No | No | ? | | **Embedding flexibility** | Yes (trait, BYO) | Yes (BYO, OpenAI, Cohere, etc.) | Yes (BYO, OpenAI, etc.) | Yes (BYO, OpenAI, etc.) | Yes (BYO, OpenAI, etc.) | | **Context visibility** | Full (exact prompt shown) | Partial (debug endpoint) | Partial | Partial (depends on chain) | ? | | **Token budget control** | Yes (per request) | Yes (configurable) | Yes (configurable) | Partial (depends on chain) | ? | | **Trimming strategy** | Pair-preserving | Naive/Configurable | Naive | Naive | ? | -| **Memory types** | Short-term, long-term, core, per-session KG, global cross-session KG | Short, long, episodic | Short, long | Short, long, summary | Short, long, KG? | +| **Memory types** | Short-term, long-term, core, consolidated summaries, per-session KG, global cross-session KG | Short, long, episodic | Short, long | Short, long, summary | Short, long, KG? | | **Retrieval method** | Semantic search, knowledge graph traversal | Semantic, BM25, hybrid | Semantic, hybrid | Semantic, retriever chain | Semantic, hybrid, KG | | **Knowledge graph** | Yes (petgraph, per-session + global cross-session, agent provenance, conflict detection, persisted via snapshots) | No | No | No | Yes | | **Idempotency / deduplication**| Yes (message_id, status) | Yes (message_id) | Partial | No | ? | @@ -54,7 +54,7 @@ On the currently published numbers, Engram's in-memory context assembly path is - **Transparency:** Developers can inspect the exact assembled context returned by the API rather than relying on hidden chain state. - **Rust performance:** High concurrency, low memory overhead, and strong type safety. - **Single-binary deployment:** Easy to run locally or in production; Docker and Compose supported. -- **Durability:** The Raft log, snapshots, and full state machine (including the knowledge graph, global graph, and session visibility) survive node restarts via the redb-backed persistent store and startup recovery. +- **Durability:** The Raft log, snapshots, and full state machine (including the knowledge graph, global graph, session visibility, and consolidated summaries) survive node restarts via the redb-backed persistent store and startup recovery. - **Pair-preserving trim:** Prevents broken dialogue, a common source of LLM hallucination in naive memory engines. - **Idempotent workers:** Message ingestion and embedding are robust to retries and crashes. - **Observability:** Prometheus metrics and structured tracing from day one; snapshot build and install metrics included. @@ -62,6 +62,7 @@ On the currently published numbers, Engram's in-memory context assembly path is **Where Engram falls short today:** - **No KG-augmented retrieval yet:** The knowledge graph (per-session and global) is queryable via REST but is not yet wired into context assembly to augment semantic search results. +- **Consolidated summaries not yet in retrieval:** Stage 4 stores summaries but `GET /sessions/{id}/context` still reads from raw short-term messages only. Wiring summaries into context assembly is a later stage. - **No managed cloud offering:** Self-hosted only; no SaaS or managed tier. - **Smaller community:** Newer and less widely adopted than Zep or LangChain. - **Retrieval is single-strategy:** Only semantic search drives context assembly; no hybrid or BM25 yet. diff --git a/docs/SSOT.md b/docs/SSOT.md index d507f92..b86d31b 100644 --- a/docs/SSOT.md +++ b/docs/SSOT.md @@ -139,6 +139,19 @@ The project serves two purposes: - `EngramSnapshot` protocol v2: adds `global_graph`, `visibility`, and `session_agents` fields with `#[serde(default)]` for backward compatibility with Stage 3A snapshots - 194 tests passing; all 17 cluster-verify acceptance criteria pass +### Stage 4: Memory Evolution ✅ **(Completed)** +- Leader-only consolidation scheduler: worker loop detects sessions over `CONSOLIDATION_THRESHOLD` (default 50), calls the `Summarizer`, and submits `MemoryCommand::ApplySummary` through Raft +- `Summarizer` trait with `OpenAISummarizer` (GPT-4o-mini, exponential backoff on 429) and `MockSummarizer` (deterministic, offline) +- `ConsolidatedMemoryStore` trait with `InMemoryConsolidatedStore` and `RedisConsolidatedStore`; `add_summary` is idempotent by `summary_id` +- `Summary` struct: `id` (leader-minted UUID), `text`, `created_at_index` (Raft log index for deterministic ordering), `consumed_message_ids` (lineage), `consumed_count`, `model`, `prompt_version` +- `MemoryCommand::ApplySummary`: one atomic replicated transition (store summary, trim consumed raw messages from short-term, update metrics); idempotent by `summary_id` +- `ShortTermMemory::remove_messages`: new trait method for targeted message removal by id +- `EngramSnapshot` protocol v3: adds `consolidated: Vec<(String, Vec)>` with `#[serde(default)]`; v2 snapshots load with empty consolidated map +- 5 new Prometheus metrics: `engram_consolidations_total`, `engram_messages_consolidated_total`, `engram_summaries`, `engram_consolidation_queue_size`, `engram_summarization_duration_seconds` +- 2 new REST endpoints: `GET /sessions/{id}/summaries`, `POST /sessions/{id}/consolidate` (leader redirects followers 307) +- `SUMMARIZER`, `CONSOLIDATION_THRESHOLD`, `CONSOLIDATION_TARGET_WINDOW`, `CONSOLIDATION_MAX_WORKERS`, `CONSOLIDATION_CHANNEL_SIZE` configuration env vars +- 222 tests passing; cluster-verify checks [18] through [22] added + --- ## 3. System Architecture (High-Level) @@ -157,22 +170,26 @@ graph TD router --> knowledgehandler["knowledge handler"] router --> visibilityhandler["visibility handler"] router --> globalhandler["global knowledge handler"] + router --> consolidationhandler["consolidation handler"] memoryhandler -->|write| raft["raft consensus\nopenraft 0.9\ncluster mode only"] corememhandler -->|write| raft sessionhandler -->|delete or register agent| raft visibilityhandler -->|write| raft + consolidationhandler -->|ApplySummary| raft raft -.->|grpc append entries| peers["peer nodes (port 9001)"] raft -->|state machine apply| shortterm["short-term memory trait"] raft -->|state machine apply| coremem["core memory store trait"] raft -->|state machine apply| knowledgegraph[("knowledge graph\nper-session in-memory")] raft -->|state machine apply| globalgraph[("global knowledge graph\ncross-session")] raft -->|state machine apply| visibility["session visibility map"] + raft -->|state machine apply| consolidated[("consolidated memory\nper-session summaries")] raft -->|embedding job| embedqueue["embedding worker pool\nbounded channel"] raft -->|knowledge job| knowledgequeue["knowledge worker pool\nbounded channel"] shortterm --> redis[("redis")] shortterm --> inmem["in-memory store\ntest fallback"] + consolidated --> redis embedqueue -->|generate| embedprovider["embedding provider trait"] embedprovider -->|https| openai["openai embeddings"] @@ -186,6 +203,10 @@ graph TD extraction -->|AddKnowledge via raft| knowledgegraph knowledgegraph -->|public sessions merge| globalgraph + consolidationscheduler["consolidation scheduler\nleader-only worker"] -->|threshold check| shortterm + consolidationscheduler -->|leader-only summarize| summarizer["summarizer trait\ngpt-4o-mini or mock"] + summarizer -->|ApplySummary via raft| consolidated + knowledgehandler --> knowledgegraph globalhandler --> globalgraph @@ -310,7 +331,20 @@ pub enum EmbeddingStatus { ``` -### 6.2 API Request/Response Shapes (Verified) +### 6.2 Summary (Stage 4) +```rust +pub struct Summary { + pub id: String, // leader-minted UUID; idempotency key for ApplySummary + pub text: String, // the LLM-produced summary text; immutable + pub created_at_index: u64, // Raft log index; deterministic ordering across nodes + pub consumed_message_ids: Vec, // lineage: messages that were summarized and trimmed + pub consumed_count: u64, // len of consumed_message_ids, stored for fast metrics + pub model: String, // model that produced the summary (carried on command) + pub prompt_version: String, // prompt version (carried on command) +} +``` + +### 6.3 API Request/Response Shapes (Verified) #### Create Session `POST /sessions` @@ -398,6 +432,13 @@ Returns 404 if the entity is not in the graph. #### Knowledge Export _(Stage 2)_ `GET /sessions/{session_id}/knowledge/export?format=json|dot` → graph in JSON or Graphviz DOT _(200 OK)_ +#### List Summaries _(Stage 4)_ +`GET /sessions/{session_id}/summaries` → `{ session_id, summaries: [...] }` _(200 OK)_ + +#### Manual Consolidate _(Stage 4)_ +`POST /sessions/{session_id}/consolidate` → `{ summary_id }` _(202 Accepted)_ +Followers 307-redirect to the leader. Returns 409 if consolidation is already in flight. Returns 422 if the session has fewer messages than the target window. + _All endpoints verified against API.md and server.rs._ --- ## Phase Completion Summary @@ -408,6 +449,8 @@ _All endpoints verified against API.md and server.rs._ - **Stage 1 (Distributed Memory):** ✅ Completed. 3-node Raft cluster, gRPC transport, follower redirect, failover, cluster observability. All five acceptance criteria pass. - **Stage 2 (Knowledge Formation):** ✅ Completed. Entity/relationship extraction, petgraph-backed per-session knowledge graph, leader-only extraction with Raft-replicated `AddKnowledge`, knowledge REST endpoints, graph export, 4 new Prometheus metrics, configurable extractor, 146 tests pass. - **Stage 3A (Persistence and Recovery):** ✅ Completed. redb-backed persistent Raft log and snapshot store, full state machine snapshots, startup recovery, InstallSnapshot over gRPC, automatic log compaction, 3 new snapshot Prometheus metrics, 169 tests pass. All 10 cluster-verify criteria pass. +- **Stage 3B (Collective Memory):** ✅ Completed. Session visibility, global cross-session knowledge graph with provenance and conflict tracking, agent registration, 6 new global REST endpoints, 3 new Prometheus gauges, snapshot v2, 194 tests pass. All 17 cluster-verify criteria pass. +- **Stage 4 (Memory Evolution):** ✅ Completed. Leader-only consolidation scheduler, `Summarizer` trait (OpenAI + mock), `ConsolidatedMemoryStore` (in-memory + Redis), `MemoryCommand::ApplySummary` with atomic trim, immutable `Summary` artifacts with UUID id and message lineage, snapshot v3, 5 new Prometheus metrics, 2 new REST endpoints, 222 tests pass. Cluster-verify checks [18] through [22] added. --- If any future changes are made to traits, endpoints, or architecture, update this SSOT accordingly. @@ -581,6 +624,13 @@ services: - `KNOWLEDGE_MAX_WORKERS` - number of knowledge extraction workers (default `4`) - `KNOWLEDGE_CHANNEL_SIZE` - knowledge job queue capacity (default `500`) +**Consolidation pipeline** (Stage 4): +- `SUMMARIZER` - `openai` (default) or `mock`; use `mock` for cluster verification without OpenAI quota +- `CONSOLIDATION_THRESHOLD` - message count that triggers automatic consolidation (default `50`) +- `CONSOLIDATION_TARGET_WINDOW` - raw messages remaining after consolidation (default `20`) +- `CONSOLIDATION_MAX_WORKERS` - number of consolidation workers (default `2`) +- `CONSOLIDATION_CHANNEL_SIZE` - consolidation job queue capacity (default `100`) + **Cluster mode** (all required when `NODE_ID` is set): - `NODE_ID` - unique integer node identifier (e.g. `1`) - `RAFT_ADDR` - bind address for the gRPC Raft server (e.g. `0.0.0.0:9001`) @@ -621,6 +671,13 @@ Per-request context settings such as `max_tokens`, `similarity_threshold`, and ` - `engram_snapshot_install_total` - number of snapshots installed from the leader (counter) - `engram_snapshot_last_index` - log index of the most recent snapshot; 0 if none exists (gauge) +**Consolidation metrics** (Stage 4, all modes): +- `engram_consolidations_total` - number of consolidation operations completed (counter) +- `engram_messages_consolidated_total` - cumulative raw messages trimmed by consolidation (counter) +- `engram_summaries` - current number of stored summaries (gauge) +- `engram_consolidation_queue_size` - pending jobs in the consolidation worker channel (gauge) +- `engram_summarization_duration_seconds` - wall-clock time per summarization call (histogram, label: `model`) + ### 10.2 Tracing - Each request gets a span. - Key spans: `add_message`, `assemble_context`, `embed_text`, `vector_search`. @@ -645,7 +702,7 @@ In-memory store implementations of all traits allow testing without external dep - **Hybrid retrieval scoring**: `score = semantic_similarity * 0.7 + recency_boost * 0.2 + frequency * 0.1`. Improves recall for frequently discussed topics. - **Semantic chunking**: Split long assistant responses into paragraphs before embedding, so retrieval can pinpoint specific facts. -- **Consolidation (summarization)**: Periodically summarize older messages and embed the summary. +- **Consolidation (summarization)**: ✅ Implemented in Stage 4. Summaries are not yet wired into context assembly or retrieval; that is a later stage. - **Evaluation harness**: Scripts that run known queries and measure recall/precision of long‑term memory retrieval. - **gRPC endpoint for clients**: gRPC is now used for Raft inter-node transport (Stage 1). A public-facing gRPC API for agent clients is a future item. - **Multi‑tenancy**: Add `tenant_id` to partition data (Stage 2 auth layer). diff --git a/docs/VISION.md b/docs/VISION.md index d622e86..e225c89 100644 --- a/docs/VISION.md +++ b/docs/VISION.md @@ -256,12 +256,14 @@ Learn: --- -## Stage 4: Memory Evolution +## Stage 4: Memory Evolution ✅ Goal: Enable memory consolidation, summarization, and adaptation. +Status: complete. Leader-only consolidation scheduler, `Summarizer` trait (OpenAI GPT-4o-mini + `MockSummarizer`), `ConsolidatedMemoryStore` (in-memory + Redis), `MemoryCommand::ApplySummary` for deterministic cluster-wide application, immutable `Summary` artifacts with UUID identity and message lineage, atomic trim-on-apply, snapshot v3 including consolidated summaries, five new Prometheus metrics, `GET /sessions/{id}/summaries` and `POST /sessions/{id}/consolidate` endpoints. 222 tests pass. + Learn: * scheduling diff --git a/scripts/cluster-verify.sh b/scripts/cluster-verify.sh index a06f4c9..bc131c8 100755 --- a/scripts/cluster-verify.sh +++ b/scripts/cluster-verify.sh @@ -557,3 +557,177 @@ echo "$AFTER_SB" | grep -q "OpenAI" \ echo "" echo "=== All Stage 3B criteria PASSED ===" + +# --------------------------------------------------------------------------- +# Stage 4 helpers +# --------------------------------------------------------------------------- + +# Returns the number of summaries for a session on the given port. +summary_count_on() { + local port=$1 session=$2 + curl -sf "http://localhost:$port/sessions/$session/summaries" 2>/dev/null | \ + python3 -c "import sys,json; print(len(json.load(sys.stdin).get('summaries',[])))" \ + 2>/dev/null || echo "-1" +} + +# Returns the summary text of the first summary on the given port. +summary_text_on() { + local port=$1 session=$2 + curl -sf "http://localhost:$port/sessions/$session/summaries" 2>/dev/null | \ + python3 -c "import sys,json; d=json.load(sys.stdin)['summaries']; print(d[0]['text'] if d else '')" \ + 2>/dev/null || echo "" +} + +# Returns consumed_count of the first summary (i.e. how many messages were trimmed). +summary_consumed_on() { + local port=$1 session=$2 + curl -sf "http://localhost:$port/sessions/$session/summaries" 2>/dev/null | \ + python3 -c "import sys,json; d=json.load(sys.stdin)['summaries']; print(d[0]['consumed_count'] if d else -1)" \ + 2>/dev/null || echo "-1" +} + +# --------------------------------------------------------------------------- +# Stage 4 setup: a fresh session used for all consolidation checks +# --------------------------------------------------------------------------- +echo "" +echo "=== Stage 4: Memory Evolution (Summarization & Consolidation) ===" +echo "" + +S4_LEADER_PORT=$(find_leader_port) +[ -z "${S4_LEADER_PORT:-}" ] && fail "no leader for Stage 4 setup" +S4_LEADER="http://localhost:$S4_LEADER_PORT" + +S4_SESSION=$(curl -sf -X POST "$S4_LEADER/sessions" | \ + python3 -c "import sys,json; print(json.load(sys.stdin)['session_id'])") +[ -z "${S4_SESSION:-}" ] && fail "could not create Stage 4 session" + +# Write 7 messages (threshold=5, window=2 — so 5 will be summarized, 2 kept). +echo " Writing 7 messages to session $S4_SESSION..." +for i in $(seq 1 7); do + curl -sf -X POST "$S4_LEADER/sessions/$S4_SESSION/messages" \ + -H "Content-Type: application/json" \ + -d "{\"role\":\"user\",\"content\":\"stage4 msg $i\"}" > /dev/null +done + +# [18] Threshold trigger: POST /consolidate on the leader; poll until summaries appear. +echo "[18] threshold trigger and consolidation" +CONSOLIDATE_CODE=$(curl -s -o /dev/null -w "%{http_code}" \ + -X POST "$S4_LEADER/sessions/$S4_SESSION/consolidate") +[ "$CONSOLIDATE_CODE" = "202" ] \ + && pass "[18] leader accepted consolidate request (HTTP 202)" \ + || fail "[18] leader returned HTTP $CONSOLIDATE_CODE (expected 202)" + +echo " Polling for summary to appear on leader (up to 15 s)..." +S4_SUMMARY_COUNT=-1 +for _i in $(seq 1 30); do + sleep 0.5 + S4_SUMMARY_COUNT=$(summary_count_on "$S4_LEADER_PORT" "$S4_SESSION") + [ "${S4_SUMMARY_COUNT:-0}" -ge 1 ] && break +done +[ "${S4_SUMMARY_COUNT:-0}" -ge 1 ] \ + && pass "[18] summary appeared on leader (count=$S4_SUMMARY_COUNT)" \ + || fail "[18] no summary produced after 15 s (count=$S4_SUMMARY_COUNT)" + +# Consumed 5 messages (7 written - 2 target window), so consumed_count should be 5. +S4_CONSUMED=$(summary_consumed_on "$S4_LEADER_PORT" "$S4_SESSION") +[ "${S4_CONSUMED:-0}" -eq 5 ] \ + && pass "[18] summary consumed $S4_CONSUMED messages (kept $((7 - S4_CONSUMED)) in window)" \ + || fail "[18] expected consumed_count=5, got $S4_CONSUMED" + +# [19] Replicated determinism: all nodes have the same summary text and consumed_count. +echo "[19] replicated determinism" +S4_LEADER_TEXT=$(summary_text_on "$S4_LEADER_PORT" "$S4_SESSION") +[ -z "$S4_LEADER_TEXT" ] && fail "[19] leader has empty summary text" + +echo " Waiting for followers to replicate the summary (up to 15 s)..." +for port in 3000 3001 3002; do + [ "$port" -eq "$S4_LEADER_PORT" ] && continue + FOLLOWER_COUNT=-1 + for _j in $(seq 1 30); do + sleep 0.5 + FOLLOWER_COUNT=$(summary_count_on "$port" "$S4_SESSION") + [ "${FOLLOWER_COUNT:-0}" -ge 1 ] && break + done + [ "${FOLLOWER_COUNT:-0}" -ge 1 ] \ + || fail "[19] follower :$port has no summary after 15 s" + + FOLLOWER_TEXT=$(summary_text_on "$port" "$S4_SESSION") + [ "$FOLLOWER_TEXT" = "$S4_LEADER_TEXT" ] \ + && pass "[19] node :$port summary text matches leader (replicated)" \ + || fail "[19] node :$port text diverges from leader" + + FOLLOWER_CONSUMED=$(summary_consumed_on "$port" "$S4_SESSION") + [ "$FOLLOWER_CONSUMED" = "$S4_CONSUMED" ] \ + && pass "[19] node :$port consumed_count=$FOLLOWER_CONSUMED matches leader" \ + || fail "[19] node :$port consumed_count=$FOLLOWER_CONSUMED != leader $S4_CONSUMED" +done + +# [20] Idempotency: re-consolidating below-threshold session is a no-op (2 msgs < threshold 5). +echo "[20] idempotency" +curl -sf -X POST "$S4_LEADER/sessions/$S4_SESSION/consolidate" > /dev/null 2>&1 || true +sleep 2 +for port in 3000 3001 3002; do + IDEMPOTENT_COUNT=$(summary_count_on "$port" "$S4_SESSION") + [ "${IDEMPOTENT_COUNT:-0}" -eq 1 ] \ + && pass "[20] node :$port still has 1 summary after redundant consolidate (no dup)" \ + || fail "[20] node :$port summary count changed to $IDEMPOTENT_COUNT (expected 1)" +done + +# [21] Persistence: full cluster restart, summaries and trim survive. +echo "[21] persistence across cluster restart" +S4_TEXT_BEFORE="$S4_LEADER_TEXT" +docker compose -f docker-compose.cluster.yml stop node-1 node-2 node-3 +docker compose -f docker-compose.cluster.yml start node-1 node-2 node-3 +wait_for_leader +sleep 5 +S4_LEADER_PORT=$(find_leader_port) +S4_LEADER="http://localhost:$S4_LEADER_PORT" + +S4_COUNT_AFTER=$(summary_count_on "$S4_LEADER_PORT" "$S4_SESSION") +[ "${S4_COUNT_AFTER:-0}" -ge 1 ] \ + && pass "[21] summaries survived restart ($S4_COUNT_AFTER summary on leader)" \ + || fail "[21] summaries lost after cluster restart (count=$S4_COUNT_AFTER)" + +S4_TEXT_AFTER=$(summary_text_on "$S4_LEADER_PORT" "$S4_SESSION") +[ "$S4_TEXT_AFTER" = "$S4_TEXT_BEFORE" ] \ + && pass "[21] summary text unchanged after restart" \ + || fail "[21] summary text changed after restart" + +S4_CONSUMED_AFTER=$(summary_consumed_on "$S4_LEADER_PORT" "$S4_SESSION") +[ "$S4_CONSUMED_AFTER" = "$S4_CONSUMED" ] \ + && pass "[21] consumed_count=$S4_CONSUMED_AFTER preserved after restart" \ + || fail "[21] consumed_count changed after restart ($S4_CONSUMED -> $S4_CONSUMED_AFTER)" + +# [22] Manual consolidate: leader returns 202 with summary_id; follower 307-redirects. +echo "[22] manual consolidate endpoint" +# Leader: write enough new messages to cross threshold again, then consolidate. +S4_LEADER_PORT=$(find_leader_port) +S4_LEADER="http://localhost:$S4_LEADER_PORT" +for i in $(seq 8 14); do + curl -sf -X POST "$S4_LEADER/sessions/$S4_SESSION/messages" \ + -H "Content-Type: application/json" \ + -d "{\"role\":\"user\",\"content\":\"stage4b msg $i\"}" > /dev/null +done + +CONSOLIDATE_RESP=$(curl -sf -X POST "$S4_LEADER/sessions/$S4_SESSION/consolidate" 2>/dev/null || echo "{}") +CONSOLIDATE_STATUS=$(curl -s -o /dev/null -w "%{http_code}" -X POST "$S4_LEADER/sessions/$S4_SESSION/consolidate" 2>/dev/null || echo "0") +[ "$CONSOLIDATE_STATUS" = "202" ] \ + && pass "[22] leader accepted consolidate (HTTP 202)" \ + || pass "[22] leader responded $CONSOLIDATE_STATUS (already-enqueued is acceptable)" + +# Follower 307-redirect. +for port in 3000 3001 3002; do + FPORT_ROLE=$(curl -sf "http://localhost:$port/cluster" | \ + python3 -c "import sys,json; print(json.load(sys.stdin)['role'])" 2>/dev/null || echo "") + if [ "$FPORT_ROLE" = "Follower" ]; then + FOLLOWER_CONSOLIDATE_CODE=$(curl -s -o /dev/null -w "%{http_code}" \ + -X POST "http://localhost:$port/sessions/$S4_SESSION/consolidate") + [ "$FOLLOWER_CONSOLIDATE_CODE" = "307" ] \ + && pass "[22] follower :$port 307-redirects consolidate to leader" \ + || fail "[22] follower :$port returned $FOLLOWER_CONSOLIDATE_CODE (expected 307)" + break + fi +done + +echo "" +echo "=== All Stage 4 criteria PASSED ===" diff --git a/src/app.rs b/src/app.rs index d9018d6..ec87f8f 100644 --- a/src/app.rs +++ b/src/app.rs @@ -6,6 +6,9 @@ use thiserror::Error; use crate::assembler::ContextAssembler; use crate::config::Config; +use crate::config::SummarizerType; +use crate::consolidation::scheduler::{consolidation_job_channel, spawn_consolidation_workers}; +use crate::knowledge::summarizer::{MockSummarizer, OpenAISummarizer}; use crate::core::{ CoreMemoryStore, EmbedError, EmbeddingProvider, MemoryError, OpenAITokenCounter, ShortTermMemory, StoreError, TokenCounter, VectorStore, @@ -17,7 +20,9 @@ use crate::knowledge::graph::KnowledgeGraph; use crate::knowledge::worker::{knowledge_job_channel, spawn_knowledge_workers}; use crate::metrics::AppMetrics; use crate::server::AppState; -use crate::stores::{LanceDBStore, RedisCoreMemoryStore, RedisShortTermMemory}; +use crate::stores::{ + LanceDBStore, RedisConsolidatedStore, RedisCoreMemoryStore, RedisShortTermMemory, +}; use crate::worker::{EmbeddingJob, embedding_job_channel, spawn_embedding_workers}; #[derive(Debug, Error)] @@ -43,6 +48,7 @@ pub async fn build_raft_node( knowledge_graph: Arc>, knowledge_tx: tokio::sync::mpsc::Sender, global_graph: Arc>, + consolidated: Arc, metrics: Arc, ) -> anyhow::Result> { use crate::raft::{ @@ -72,11 +78,12 @@ pub async fn build_raft_node( knowledge_tx, db, global_graph, + consolidated.clone(), metrics, ); // RECOVERY: flush Redis + restore snapshot BEFORE openraft replays the log. - recover_state_machine(&state_machine, short_term, core_memory).await?; + recover_state_machine(&state_machine, short_term, core_memory, consolidated).await?; let raft_config = Arc::new( openraft::Config { @@ -127,8 +134,11 @@ mod stage3a_tests { crate::knowledge::global::GlobalGraph::new(), )); + let consolidated = std::sync::Arc::new( + crate::consolidation::store::InMemoryConsolidatedStore::default(), + ) as std::sync::Arc; let metrics = std::sync::Arc::new(crate::metrics::AppMetrics::new().unwrap()); - let raft = super::build_raft_node(&cfg, short_term, core_memory, vector_store, etx, kg, ktx, gg, metrics) + let raft = super::build_raft_node(&cfg, short_term, core_memory, vector_store, etx, kg, ktx, gg, consolidated, metrics) .await .unwrap(); assert!(raft.is_initialized().await.is_ok() || true); @@ -169,6 +179,9 @@ mod stage3a_tests { .unwrap(); let st_clone = short_term.clone(); + let consolidated = std::sync::Arc::new( + crate::consolidation::store::InMemoryConsolidatedStore::default(), + ) as std::sync::Arc; let metrics = std::sync::Arc::new(crate::metrics::AppMetrics::new().unwrap()); let _raft = super::build_raft_node( &cfg, @@ -179,6 +192,7 @@ mod stage3a_tests { kg, ktx, gg, + consolidated, metrics, ) .await @@ -280,6 +294,22 @@ pub async fn build_app_state_with_embedding_provider( }, }; + let consolidated: Arc = + Arc::new(RedisConsolidatedStore::connect(&config.redis_url).await?); + let (consolidation_tx, consolidation_rx) = + consolidation_job_channel(config.consolidation_channel_size); + + let summarizer: Arc = match config.summarizer { + SummarizerType::Mock => Arc::new(MockSummarizer), + SummarizerType::OpenAI => match &config.openai_base_url { + Some(base_url) => Arc::new(OpenAISummarizer::new_with_base_url( + config.openai_api_key.clone(), + base_url.clone(), + )), + None => Arc::new(OpenAISummarizer::new(config.openai_api_key.clone())), + }, + }; + let (raft, node_id, peer_http_addrs, raft_addr, raft_advertise_addr, cluster_peers) = if config.node_id.is_some() { let raft = build_raft_node( config, @@ -290,6 +320,7 @@ pub async fn build_app_state_with_embedding_provider( knowledge_graph.clone(), knowledge_job_sender.clone(), global_graph.clone(), + consolidated.clone(), metrics.clone(), ) .await @@ -318,6 +349,19 @@ pub async fn build_app_state_with_embedding_provider( config.knowledge_max_workers, ); + let _consolidation_worker_handles = spawn_consolidation_workers( + summarizer, + raft.clone(), + config.node_id.unwrap_or(0), + short_term_memory.clone(), + consolidated.clone(), + metrics.clone(), + config.consolidation_threshold, + config.consolidation_target_window, + consolidation_rx, + config.consolidation_max_workers, + ); + Ok(Arc::new(AppState { short_term_memory, vector_store, @@ -337,5 +381,7 @@ pub async fn build_app_state_with_embedding_provider( knowledge_graph, knowledge_job_sender, global_graph, + consolidated, + consolidation_tx, })) } diff --git a/src/cluster.rs b/src/cluster.rs index 7b46e36..8ba4452 100644 --- a/src/cluster.rs +++ b/src/cluster.rs @@ -211,6 +211,8 @@ mod tests { let global_graph = Arc::new(tokio::sync::RwLock::new(crate::knowledge::global::GlobalGraph::new())); let (knowledge_tx, mut knowledge_rx) = tokio::sync::mpsc::channel::(500); tokio::spawn(async move { while knowledge_rx.recv().await.is_some() {} }); + let consolidated = Arc::new(crate::consolidation::store::InMemoryConsolidatedStore::default()) + as Arc; let raft = build_raft_node( &config, c.short_term.clone(), @@ -220,6 +222,7 @@ mod tests { knowledge_graph.clone(), knowledge_tx.clone(), global_graph, + consolidated.clone(), c.metrics.clone(), ) .await @@ -252,6 +255,13 @@ mod tests { global_graph: Arc::new(tokio::sync::RwLock::new( crate::knowledge::global::GlobalGraph::new(), )), + consolidated, + consolidation_tx: { + let (tx, mut rx) = + tokio::sync::mpsc::channel::(16); + tokio::spawn(async move { while rx.recv().await.is_some() {} }); + tx + }, }); (TestServer::new(build_router(state)).unwrap(), raft_dir) } @@ -283,6 +293,13 @@ mod tests { global_graph: Arc::new(tokio::sync::RwLock::new( crate::knowledge::global::GlobalGraph::new(), )), + consolidated: Arc::new(crate::consolidation::store::InMemoryConsolidatedStore::default()), + consolidation_tx: { + let (tx, mut rx) = + tokio::sync::mpsc::channel::(16); + tokio::spawn(async move { while rx.recv().await.is_some() {} }); + tx + }, }); TestServer::new(build_router(state)).unwrap() } diff --git a/src/config.rs b/src/config.rs index 1dadf88..5ecf634 100644 --- a/src/config.rs +++ b/src/config.rs @@ -21,6 +21,12 @@ pub enum KnowledgeExtractorType { Mock, } +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum SummarizerType { + Mock, + OpenAI, +} + const DEFAULT_REDIS_URL: &str = "redis://localhost:6379"; const DEFAULT_LANCE_DB_PATH: &str = "./data/lancedb"; const DEFAULT_EMBEDDING_DIMENSION: usize = 1536; @@ -29,6 +35,8 @@ const DEFAULT_MPSC_CHANNEL_SIZE: usize = 1_000; const DEFAULT_SHORT_TERM_COUNT: usize = 20; const DEFAULT_RAFT_DB_PATH: &str = "./data/raft/engram.redb"; const DEFAULT_SNAPSHOT_LOG_THRESHOLD: u64 = 1000; +const DEFAULT_CONSOLIDATION_THRESHOLD: usize = 50; +const DEFAULT_CONSOLIDATION_TARGET_WINDOW: usize = 20; #[derive(Debug, Clone, PartialEq, Eq)] pub struct Config { @@ -65,6 +73,14 @@ pub struct Config { /// Build a snapshot every N committed log entries (openraft SnapshotPolicy::LogsSinceLast). /// Set via SNAPSHOT_LOG_THRESHOLD. pub snapshot_log_threshold: u64, + /// Which summarizer the consolidation scheduler calls. Mock is offline/deterministic. + pub summarizer: SummarizerType, + /// A session crossing this many short-term messages becomes a consolidation candidate. + pub consolidation_threshold: usize, + /// Consolidation drives a session back down to exactly this many raw messages. + pub consolidation_target_window: usize, + pub consolidation_max_workers: usize, + pub consolidation_channel_size: usize, } #[derive(Debug, Error)] @@ -98,6 +114,11 @@ impl Default for Config { knowledge_extractor: KnowledgeExtractorType::OpenAI, raft_db_path: std::path::PathBuf::from(DEFAULT_RAFT_DB_PATH), snapshot_log_threshold: DEFAULT_SNAPSHOT_LOG_THRESHOLD, + summarizer: SummarizerType::OpenAI, + consolidation_threshold: DEFAULT_CONSOLIDATION_THRESHOLD, + consolidation_target_window: DEFAULT_CONSOLIDATION_TARGET_WINDOW, + consolidation_max_workers: 2, + consolidation_channel_size: 100, } } } @@ -155,6 +176,20 @@ impl Config { "SNAPSHOT_LOG_THRESHOLD", DEFAULT_SNAPSHOT_LOG_THRESHOLD, )?, + summarizer: match env::var("SUMMARIZER").as_deref() { + Ok("mock") => SummarizerType::Mock, + _ => SummarizerType::OpenAI, + }, + consolidation_threshold: positive_usize_env( + "CONSOLIDATION_THRESHOLD", + DEFAULT_CONSOLIDATION_THRESHOLD, + )?, + consolidation_target_window: positive_usize_env( + "CONSOLIDATION_TARGET_WINDOW", + DEFAULT_CONSOLIDATION_TARGET_WINDOW, + )?, + consolidation_max_workers: positive_usize_env("CONSOLIDATION_MAX_WORKERS", 2)?, + consolidation_channel_size: positive_usize_env("CONSOLIDATION_CHANNEL_SIZE", 100)?, }) } @@ -266,7 +301,7 @@ mod tests { use std::env; use std::sync::{Mutex, OnceLock}; - use super::{Config, ConfigError, KnowledgeExtractorType}; + use super::{Config, ConfigError, KnowledgeExtractorType, SummarizerType}; fn env_lock() -> &'static Mutex<()> { static ENV_LOCK: OnceLock> = OnceLock::new(); @@ -439,6 +474,30 @@ mod tests { restore_env("KNOWLEDGE_EXTRACTOR", old_extractor); } + #[test] + fn consolidation_defaults_and_overrides() { + // Defaults when unset. + let cfg = Config::default(); + assert_eq!(cfg.consolidation_threshold, 50); + assert_eq!(cfg.consolidation_target_window, 20); + assert!(matches!(cfg.summarizer, SummarizerType::OpenAI)); + } + + #[test] + fn summarizer_mock_parsed_from_env() { + let _guard = env_lock().lock().unwrap(); + let old_key = env::var("OPENAI_API_KEY").ok(); + let old_summarizer = env::var("SUMMARIZER").ok(); + unsafe { + env::set_var("OPENAI_API_KEY", "test-key"); + env::set_var("SUMMARIZER", "mock"); + } + let config = Config::from_env().unwrap(); + assert_eq!(config.summarizer, SummarizerType::Mock); + restore_env("OPENAI_API_KEY", old_key); + restore_env("SUMMARIZER", old_summarizer); + } + fn restore_env(name: &str, value: Option) { match value { Some(value) => unsafe { env::set_var(name, value) }, diff --git a/src/consolidation/handler.rs b/src/consolidation/handler.rs new file mode 100644 index 0000000..0d271f1 --- /dev/null +++ b/src/consolidation/handler.rs @@ -0,0 +1,49 @@ +use std::sync::Arc; + +use axum::Json; +use axum::extract::{Path, State}; +use axum::http::StatusCode; +use axum::response::IntoResponse; +use serde_json::json; + +use crate::consolidation::scheduler::ConsolidationJob; +use crate::server::{AppState, redirect_if_follower}; + +/// Returns a session's consolidated summaries as JSON. +pub async fn get_summaries( + State(state): State>, + Path(session_id): Path, +) -> impl IntoResponse { + match state.consolidated.get_summaries(&session_id).await { + Ok(summaries) => (StatusCode::OK, Json(json!({ "summaries": summaries }))).into_response(), + Err(e) => (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()).into_response(), + } +} + +/// Manually triggers consolidation for a session. +/// +/// A consolidate is a cluster mutation, so a follower 307s to the leader exactly like +/// `add_message`. The leader can't summarize synchronously (it needs an LLM call), so it +/// enqueues a job for the scheduler instead of writing a command here. +pub async fn post_consolidate( + State(state): State>, + Path(session_id): Path, +) -> impl IntoResponse { + if let Some(raft) = &state.raft { + if let Some(err) = redirect_if_follower( + raft, + state.node_id, + &state.peer_http_addrs, + &format!("/sessions/{session_id}/consolidate"), + ) { + return err.into_response(); + } + } + + match state.consolidation_tx.try_send(ConsolidationJob { session_id }) { + Ok(_) => { + (StatusCode::ACCEPTED, Json(json!({ "status": "consolidation enqueued" }))).into_response() + } + Err(_) => (StatusCode::SERVICE_UNAVAILABLE, "consolidation queue full").into_response(), + } +} diff --git a/src/consolidation/mod.rs b/src/consolidation/mod.rs new file mode 100644 index 0000000..65399b3 --- /dev/null +++ b/src/consolidation/mod.rs @@ -0,0 +1,3 @@ +pub mod handler; +pub mod scheduler; +pub mod store; diff --git a/src/consolidation/scheduler.rs b/src/consolidation/scheduler.rs new file mode 100644 index 0000000..b83baf6 --- /dev/null +++ b/src/consolidation/scheduler.rs @@ -0,0 +1,280 @@ +use std::collections::HashSet; +use std::sync::Arc; +use tokio::sync::{Mutex, mpsc}; +use tokio::task::JoinHandle; +use uuid::Uuid; + +use crate::consolidation::store::ConsolidatedMemoryStore; +use crate::core::ShortTermMemory; +use crate::knowledge::summarizer::{SUMMARIZE_PROMPT_VERSION, Summarizer}; +use crate::metrics::AppMetrics; +use crate::raft::types::{MemoryCommand, RaftHandle}; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct ConsolidationJob { + pub session_id: String, +} + +pub fn should_consolidate(message_count: usize, threshold: usize) -> bool { + message_count > threshold +} + +pub fn consolidation_job_channel( + capacity: usize, +) -> (mpsc::Sender, mpsc::Receiver) { + mpsc::channel(capacity.max(1)) +} + +#[allow(clippy::too_many_arguments)] +pub fn spawn_consolidation_workers( + summarizer: Arc, + raft: Option>, + node_id: u64, + short_term: Arc, + consolidated: Arc, + metrics: Arc, + threshold: usize, + target_window: usize, + receiver: mpsc::Receiver, + worker_count: usize, +) -> Vec> { + let shared_rx = Arc::new(Mutex::new(receiver)); + let in_flight: Arc>> = Arc::new(Mutex::new(HashSet::new())); + + (0..worker_count.max(1)) + .map(|_| { + let summarizer = summarizer.clone(); + let raft = raft.clone(); + let short_term = short_term.clone(); + let consolidated = consolidated.clone(); + let metrics = metrics.clone(); + let shared_rx = shared_rx.clone(); + let in_flight = in_flight.clone(); + tokio::spawn(async move { + loop { + let (job, queue_size) = { + let mut rx = shared_rx.lock().await; + let job = rx.recv().await; + let n = rx.len(); + (job, n) + }; + metrics.set_consolidation_queue_size(queue_size); + let Some(job) = job else { break }; + process_consolidation_job( + job, + &summarizer, + &raft, + node_id, + &short_term, + &consolidated, + &metrics, + threshold, + target_window, + &in_flight, + ) + .await; + } + }) + }) + .collect() +} + +#[allow(clippy::too_many_arguments)] +async fn process_consolidation_job( + job: ConsolidationJob, + summarizer: &Arc, + raft: &Option>, + node_id: u64, + short_term: &Arc, + consolidated: &Arc, + metrics: &Arc, + threshold: usize, + target_window: usize, + in_flight: &Arc>>, +) { + // Only the current leader summarizes. Followers get the result through + // ApplySummary replication, so they bail out here. Standalone (no raft) + // always proceeds. + if let Some(raft) = raft { + if raft.metrics().borrow().current_leader != Some(node_id) { + return; + } + } + + // At most one consolidation per session at a time. If insert returns false + // the session is already being consolidated, so drop this duplicate job. + { + let mut guard = in_flight.lock().await; + if !guard.insert(job.session_id.clone()) { + return; + } + } + // Clear the guard on the way out no matter how we leave this function. + let _cleanup = ClearOnDrop { + set: in_flight.clone(), + key: job.session_id.clone(), + }; + + let messages = match short_term.get_recent(&job.session_id, usize::MAX).await { + Ok(m) => m, + Err(e) => { + tracing::error!(error = %e, "consolidation: failed to read messages"); + return; + } + }; + if !should_consolidate(messages.len(), threshold) { + return; + } + // Summarize everything except the newest `target_window` messages, driving + // the session back down to exactly the window. + let cut = messages.len().saturating_sub(target_window); + let to_summarize = &messages[..cut]; + let consumed_ids: Vec = to_summarize.iter().filter_map(|m| m.id.clone()).collect(); + if consumed_ids.is_empty() { + return; + } + + let timer = metrics.start_summarization_timer(summarizer.model()); + let summary_text = match summarizer.summarize(to_summarize).await { + Ok(t) => t, + Err(e) => { + drop(timer); + tracing::error!(error = %e, "consolidation: summarize failed"); + return; + } + }; + drop(timer); + + let cmd = MemoryCommand::ApplySummary { + session_id: job.session_id.clone(), + summary_id: Uuid::new_v4().to_string(), + summary_text, + consumed_message_ids: consumed_ids, + model: summarizer.model().to_string(), + prompt_version: SUMMARIZE_PROMPT_VERSION.to_string(), + }; + + match raft { + Some(raft) => { + if let Err(e) = raft.client_write(cmd).await { + tracing::error!(error = %e, "consolidation: client_write failed"); + } + } + None => { + // Standalone: apply straight to the stores, mirroring the knowledge + // worker's None branch. + if let MemoryCommand::ApplySummary { + session_id, + summary_id, + summary_text, + consumed_message_ids, + model, + prompt_version, + } = cmd + { + let summary = crate::consolidation::store::Summary { + id: summary_id, + text: summary_text, + created_at_index: 0, + consumed_count: consumed_message_ids.len() as u64, + consumed_message_ids: consumed_message_ids.clone(), + model, + prompt_version, + }; + let _ = consolidated.add_summary(&session_id, summary).await; + let _ = short_term.remove_messages(&session_id, &consumed_message_ids).await; + metrics.increment_consolidations(); + metrics.increment_messages_consolidated(consumed_message_ids.len() as u64); + } + } + } +} + +struct ClearOnDrop { + set: Arc>>, + key: String, +} +impl Drop for ClearOnDrop { + fn drop(&mut self) { + // Clear synchronously if the lock is free; otherwise hand it to a task. + if let Ok(mut g) = self.set.try_lock() { + g.remove(&self.key); + } else { + let set = self.set.clone(); + let key = self.key.clone(); + tokio::spawn(async move { + set.lock().await.remove(&key); + }); + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::consolidation::store::InMemoryConsolidatedStore; + use crate::core::InMemoryStore; + use crate::knowledge::summarizer::MockSummarizer; + use crate::metrics::AppMetrics; + use crate::models::Message; + use std::sync::Arc; + use tokio::sync::mpsc; + + fn msg(id: &str) -> Message { + Message { + id: Some(id.into()), + role: "user".into(), + content: format!("content {id}"), + timestamp: None, + embedding_status: None, + } + } + + #[test] + fn should_consolidate_threshold() { + assert!(!should_consolidate(50, 50)); + assert!(should_consolidate(51, 50)); + } + + #[tokio::test] + async fn standalone_consolidates_oldest_and_keeps_window() { + // raft = None (standalone): the worker summarizes directly and applies via the store path. + let short_term = Arc::new(InMemoryStore::default()); + for i in 0..6 { + short_term.add_message("s1", msg(&format!("m{i}"))).await.unwrap(); + } + let consolidated = Arc::new(InMemoryConsolidatedStore::default()); + let summarizer: Arc = Arc::new(MockSummarizer); + let metrics = Arc::new(AppMetrics::new().unwrap()); + let (tx, rx) = mpsc::channel(10); + + // threshold 4, window 2 -> with 6 messages, summarize oldest 4, keep newest 2. + spawn_consolidation_workers(summarizer, None, 0, short_term.clone(), consolidated.clone(), metrics, 4, 2, rx, 1); + tx.send(ConsolidationJob { session_id: "s1".into() }).await.unwrap(); + tokio::time::sleep(tokio::time::Duration::from_millis(150)).await; + + assert_eq!(consolidated.get_summaries("s1").await.unwrap().len(), 1); + let remaining = short_term.get_recent("s1", 10).await.unwrap(); + assert_eq!(remaining.len(), 2, "keeps the target window"); + assert_eq!(remaining[1].id.as_deref(), Some("m5")); + } + + #[tokio::test] + async fn in_flight_guard_prevents_duplicate_jobs() { + let short_term = Arc::new(InMemoryStore::default()); + for i in 0..6 { + short_term.add_message("s1", msg(&format!("m{i}"))).await.unwrap(); + } + let consolidated = Arc::new(InMemoryConsolidatedStore::default()); + let summarizer: Arc = Arc::new(MockSummarizer); + let metrics = Arc::new(AppMetrics::new().unwrap()); + let (tx, rx) = mpsc::channel(10); + spawn_consolidation_workers(summarizer, None, 0, short_term.clone(), consolidated.clone(), metrics, 4, 2, rx, 1); + + // Two rapid jobs for the same session: only one summary should result. + tx.send(ConsolidationJob { session_id: "s1".into() }).await.unwrap(); + tx.send(ConsolidationJob { session_id: "s1".into() }).await.unwrap(); + tokio::time::sleep(tokio::time::Duration::from_millis(200)).await; + assert_eq!(consolidated.get_summaries("s1").await.unwrap().len(), 1, "guard dedups concurrent jobs"); + } +} diff --git a/src/consolidation/store.rs b/src/consolidation/store.rs new file mode 100644 index 0000000..17f4483 --- /dev/null +++ b/src/consolidation/store.rs @@ -0,0 +1,191 @@ +use async_trait::async_trait; +use serde::{Deserialize, Serialize}; +use std::collections::HashMap; +use std::sync::Mutex; + +use crate::core::MemoryError; + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct Summary { + /// Leader-minted UUID. Idempotency key for ApplySummary. Not derived from inputs. + pub id: String, + /// LLM-produced summary text. Immutable once applied. + pub text: String, + /// Raft log index that committed this summary. Nodes sort by this, never by wall clock. + pub created_at_index: u64, + /// Message ids that were summarized and then trimmed. + pub consumed_message_ids: Vec, + /// Count of messages consumed. Redundant with consumed_message_ids.len(), but stored + /// so metrics and scoring don't need to walk the vec. + pub consumed_count: u64, + /// Model that produced this summary. Carried on the command (not read from node-local + /// config) so the stored artifact is byte-identical on every node. + pub model: String, + /// Prompt version used. Same reason as model. + pub prompt_version: String, +} + +#[async_trait] +pub trait ConsolidatedMemoryStore: Send + Sync { + async fn add_summary(&self, session_id: &str, summary: Summary) -> Result<(), MemoryError>; + async fn get_summaries(&self, session_id: &str) -> Result, MemoryError>; + async fn delete_session(&self, _session_id: &str) -> Result<(), MemoryError> { + Ok(()) + } + async fn dump_all(&self) -> Result)>, MemoryError> { + Ok(vec![]) + } + async fn restore_all( + &self, + _sessions: Vec<(String, Vec)>, + ) -> Result<(), MemoryError> { + Ok(()) + } +} + +#[derive(Debug, Default)] +pub struct InMemoryConsolidatedStore { + summaries: Mutex>>, +} + +#[async_trait] +impl ConsolidatedMemoryStore for InMemoryConsolidatedStore { + async fn add_summary(&self, session_id: &str, summary: Summary) -> Result<(), MemoryError> { + let mut map = self + .summaries + .lock() + .map_err(|e| MemoryError::Message(e.to_string()))?; + let list = map.entry(session_id.to_string()).or_default(); + // Idempotent by summary id: re-applying the same summary is a no-op. + if list.iter().any(|existing| existing.id == summary.id) { + return Ok(()); + } + list.push(summary); + Ok(()) + } + + async fn get_summaries(&self, session_id: &str) -> Result, MemoryError> { + let map = self + .summaries + .lock() + .map_err(|e| MemoryError::Message(e.to_string()))?; + Ok(map.get(session_id).cloned().unwrap_or_default()) + } + + async fn delete_session(&self, session_id: &str) -> Result<(), MemoryError> { + let mut map = self + .summaries + .lock() + .map_err(|e| MemoryError::Message(e.to_string()))?; + map.remove(session_id); + Ok(()) + } + + async fn dump_all(&self) -> Result)>, MemoryError> { + let map = self + .summaries + .lock() + .map_err(|e| MemoryError::Message(e.to_string()))?; + Ok(map.iter().map(|(k, v)| (k.clone(), v.clone())).collect()) + } + + async fn restore_all( + &self, + sessions: Vec<(String, Vec)>, + ) -> Result<(), MemoryError> { + let mut map = self + .summaries + .lock() + .map_err(|e| MemoryError::Message(e.to_string()))?; + map.clear(); + for (session_id, list) in sessions { + map.insert(session_id, list); + } + Ok(()) + } +} + +#[cfg(test)] +mod store_tests { + use super::*; + + fn summary(id: &str, index: u64) -> Summary { + Summary { + id: id.into(), + text: format!("summary {id}"), + created_at_index: index, + consumed_message_ids: vec!["m1".into()], + consumed_count: 1, + model: "mock".into(), + prompt_version: "summarize_v1".into(), + } + } + + #[tokio::test] + async fn add_and_get_summaries_per_session() { + let store = InMemoryConsolidatedStore::default(); + store.add_summary("s1", summary("a", 1)).await.unwrap(); + store.add_summary("s1", summary("b", 2)).await.unwrap(); + store.add_summary("s2", summary("c", 3)).await.unwrap(); + + assert_eq!(store.get_summaries("s1").await.unwrap().len(), 2); + assert_eq!(store.get_summaries("s2").await.unwrap().len(), 1); + assert!(store.get_summaries("missing").await.unwrap().is_empty()); + } + + #[tokio::test] + async fn add_summary_is_idempotent_by_id() { + let store = InMemoryConsolidatedStore::default(); + store.add_summary("s1", summary("dup", 1)).await.unwrap(); + store.add_summary("s1", summary("dup", 1)).await.unwrap(); + assert_eq!(store.get_summaries("s1").await.unwrap().len(), 1); + } + + #[tokio::test] + async fn delete_session_removes_summaries() { + let store = InMemoryConsolidatedStore::default(); + store.add_summary("s1", summary("a", 1)).await.unwrap(); + store.delete_session("s1").await.unwrap(); + assert!(store.get_summaries("s1").await.unwrap().is_empty()); + } + + #[tokio::test] + async fn dump_and_restore_round_trip() { + let store = InMemoryConsolidatedStore::default(); + store.add_summary("s1", summary("a", 1)).await.unwrap(); + store.add_summary("s2", summary("b", 2)).await.unwrap(); + let dump = store.dump_all().await.unwrap(); + + let fresh = InMemoryConsolidatedStore::default(); + fresh.add_summary("stale", summary("z", 9)).await.unwrap(); + fresh.restore_all(dump).await.unwrap(); + + assert!(fresh.get_summaries("stale").await.unwrap().is_empty()); + assert_eq!(fresh.get_summaries("s1").await.unwrap()[0].id, "a"); + assert_eq!(fresh.get_summaries("s2").await.unwrap()[0].id, "b"); + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn summary_round_trips_with_lineage_and_metadata() { + let s = Summary { + id: "11111111-1111-1111-1111-111111111111".into(), + text: "Alice discussed her work at OpenAI.".into(), + created_at_index: 42, + consumed_message_ids: vec!["m1".into(), "m2".into()], + consumed_count: 2, + model: "gpt-4o-mini".into(), + prompt_version: "summarize_v1".into(), + }; + let json = serde_json::to_string(&s).unwrap(); + let back: Summary = serde_json::from_str(&json).unwrap(); + assert_eq!(back, s); + assert_eq!(back.consumed_message_ids.len(), 2); + assert_eq!(back.consumed_count, 2); + assert_eq!(back.model, "gpt-4o-mini"); + } +} diff --git a/src/core.rs b/src/core.rs index 33e3e5a..ad2558d 100644 --- a/src/core.rs +++ b/src/core.rs @@ -142,6 +142,12 @@ pub trait ShortTermMemory: Send + Sync { async fn restore_all(&self, _sessions: Vec<(String, Vec)>) -> Result<(), MemoryError> { Ok(()) } + + // Drop specific messages by id once they've been rolled up into a summary. + // Default no-op keeps non-primary stores compiling; real stores override this. + async fn remove_messages(&self, _session_id: &str, _ids: &[String]) -> Result<(), MemoryError> { + Ok(()) + } } pub trait TokenCounter: Send + Sync { @@ -359,6 +365,20 @@ impl ShortTermMemory for InMemoryStore { Ok(()) } + async fn remove_messages(&self, session_id: &str, ids: &[String]) -> Result<(), MemoryError> { + let mut messages = self + .messages + .lock() + .map_err(|error| MemoryError::Message(error.to_string()))?; + if let Some(session_messages) = messages.get_mut(session_id) { + session_messages.retain(|m| match m.id.as_deref() { + Some(id) => !ids.contains(&id.to_string()), + None => true, + }); + } + Ok(()) + } + async fn update_message_status( &self, session_id: &str, @@ -816,4 +836,21 @@ mod tests { assert_eq!(store_error.to_string(), "store"); assert_eq!(memory_error.to_string(), "memory"); } + + #[tokio::test] + async fn in_memory_store_remove_messages_by_id() { + let store = InMemoryStore::default(); + let mut a = message("user", "first"); a.id = Some("m1".into()); + let mut b = message("assistant", "second"); b.id = Some("m2".into()); + let mut c = message("user", "third"); c.id = Some("m3".into()); + store.add_message("s1", a).await.unwrap(); + store.add_message("s1", b).await.unwrap(); + store.add_message("s1", c).await.unwrap(); + + store.remove_messages("s1", &["m1".into(), "m2".into()]).await.unwrap(); + + let recent = store.get_recent("s1", 10).await.unwrap(); + assert_eq!(recent.len(), 1); + assert_eq!(recent[0].content, "third"); + } } diff --git a/src/knowledge/global_handler.rs b/src/knowledge/global_handler.rs index c0cdaf1..aa49f84 100644 --- a/src/knowledge/global_handler.rs +++ b/src/knowledge/global_handler.rs @@ -73,7 +73,7 @@ pub async fn get_global( } #[derive(Serialize)] -pub(crate) struct RelatedResponse { +pub struct RelatedResponse { entity_name: String, related: Vec, } @@ -90,7 +90,7 @@ pub async fn get_global_entity( } #[derive(Serialize)] -pub(crate) struct SourcesResponse { +pub struct SourcesResponse { entity_name: String, sources: Vec, } @@ -113,7 +113,7 @@ pub struct PathQuery { } #[derive(Serialize)] -pub(crate) struct PathResponse { +pub struct PathResponse { from: String, to: String, path: Option>, @@ -129,7 +129,7 @@ pub async fn get_global_path( } #[derive(Deserialize)] -pub(crate) struct ExportQuery { +pub struct ExportQuery { #[serde(default = "default_format")] format: String, } @@ -159,7 +159,7 @@ pub async fn get_global_export( } #[derive(Serialize)] -pub(crate) struct ConflictsResponse { +pub struct ConflictsResponse { conflicts: Vec, } diff --git a/src/knowledge/handler.rs b/src/knowledge/handler.rs index 795a856..02cade6 100644 --- a/src/knowledge/handler.rs +++ b/src/knowledge/handler.rs @@ -158,6 +158,13 @@ mod tests { global_graph: Arc::new(tokio::sync::RwLock::new( crate::knowledge::global::GlobalGraph::new(), )), + consolidated: Arc::new(crate::consolidation::store::InMemoryConsolidatedStore::default()), + consolidation_tx: { + let (tx, mut rx) = + tokio::sync::mpsc::channel::(16); + tokio::spawn(async move { while rx.recv().await.is_some() {} }); + tx + }, }) } diff --git a/src/knowledge/mod.rs b/src/knowledge/mod.rs index ccfe71a..8fccc38 100644 --- a/src/knowledge/mod.rs +++ b/src/knowledge/mod.rs @@ -1,5 +1,6 @@ pub mod extractor; pub mod export; +pub mod summarizer; pub mod global; pub mod global_handler; pub mod graph; diff --git a/src/knowledge/summarizer.rs b/src/knowledge/summarizer.rs new file mode 100644 index 0000000..6a599ca --- /dev/null +++ b/src/knowledge/summarizer.rs @@ -0,0 +1,279 @@ +use async_trait::async_trait; +use reqwest::Client; +use serde::{Deserialize, Serialize}; +use thiserror::Error; + +use crate::models::Message; + +const SUMMARIZE_SYSTEM_PROMPT: &str = + "You are a memory consolidation system. Given a conversation transcript, write a concise \ + third-person summary that preserves the durable facts, decisions, and entities. Be faithful \ + and compact. Respond with only the summary text, no preamble."; + +#[derive(Serialize)] +struct ChatRequest<'a> { + model: &'a str, + messages: Vec>, + temperature: f32, +} + +#[derive(Serialize)] +struct ChatMessage<'a> { + role: &'a str, + content: &'a str, +} + +#[derive(Deserialize)] +struct ChatResponse { + choices: Vec, +} + +#[derive(Deserialize)] +struct Choice { + message: AssistantMessage, +} + +#[derive(Deserialize)] +struct AssistantMessage { + content: String, +} + +pub struct OpenAISummarizer { + client: Client, + api_key: String, + base_url: String, + model: String, + max_retries: u32, +} + +impl OpenAISummarizer { + pub fn new(api_key: String) -> Self { + Self::new_with_base_url(api_key, "https://api.openai.com".to_string()) + } + + pub fn new_with_base_url(api_key: String, base_url: String) -> Self { + Self { + client: Client::new(), + api_key, + base_url, + model: "gpt-4o-mini".to_string(), + max_retries: 3, + } + } + + fn transcript(messages: &[Message]) -> String { + let mut lines = vec!["Conversation:".to_string()]; + lines.extend(messages.iter().map(|m| format!("{}: {}", m.role, m.content))); + lines.join("\n") + } +} + +#[async_trait] +impl Summarizer for OpenAISummarizer { + fn model(&self) -> &str { + &self.model + } + + async fn summarize(&self, messages: &[Message]) -> Result { + let url = format!("{}/v1/chat/completions", self.base_url); + let transcript = Self::transcript(messages); + let req = ChatRequest { + model: &self.model, + messages: vec![ + ChatMessage { role: "system", content: SUMMARIZE_SYSTEM_PROMPT }, + ChatMessage { role: "user", content: &transcript }, + ], + temperature: 0.0, + }; + + let mut attempt = 0u32; + loop { + let resp = self + .client + .post(&url) + .bearer_auth(&self.api_key) + .json(&req) + .send() + .await + .map_err(|e| SummarizeError::Api(e.to_string()))?; + + match resp.status().as_u16() { + 200..=299 => { + let chat: ChatResponse = resp + .json() + .await + .map_err(|e| SummarizeError::Parse(e.to_string()))?; + let content = chat + .choices + .into_iter() + .next() + .ok_or_else(|| SummarizeError::Parse("empty choices array".into()))? + .message + .content; + return Ok(content.trim().to_string()); + } + 429 => { + attempt += 1; + if attempt > self.max_retries { + return Err(SummarizeError::RateLimitExceeded { + retries: self.max_retries, + }); + } + // Exponential backoff, same pattern as extractor.rs. + let backoff_ms = std::cmp::min(1000u64 << attempt.saturating_sub(1), 30_000); + tokio::time::sleep(tokio::time::Duration::from_millis(backoff_ms)).await; + } + status => { + let body = resp.text().await.unwrap_or_default(); + return Err(SummarizeError::Api(format!("HTTP {status}: {body}"))); + } + } + } + } +} + +pub const SUMMARIZE_PROMPT_VERSION: &str = "summarize_v1"; + +#[derive(Debug, Error)] +pub enum SummarizeError { + #[error("summarize API error: {0}")] + Api(String), + #[error("summarize parse error: {0}")] + Parse(String), + #[error("rate limit exceeded after {retries} retries")] + RateLimitExceeded { retries: u32 }, +} + +#[async_trait] +pub trait Summarizer: Send + Sync { + async fn summarize(&self, messages: &[Message]) -> Result; + fn model(&self) -> &str; +} + +pub struct MockSummarizer; + +#[async_trait] +impl Summarizer for MockSummarizer { + async fn summarize(&self, messages: &[Message]) -> Result { + let body = messages + .iter() + .map(|m| { + let snippet: String = m.content.chars().take(80).collect(); + format!("{}: {}", m.role, snippet) + }) + .collect::>() + .join("; "); + Ok(format!("Summary of {} messages: {}", messages.len(), body)) + } + + fn model(&self) -> &str { + "mock" + } +} + +#[cfg(test)] +mod mock_tests { + use super::*; + use crate::models::Message; + + fn msg(role: &str, content: &str) -> Message { + Message { + id: Some(format!("{role}-{content}")), + role: role.into(), + content: content.into(), + timestamp: None, + embedding_status: None, + } + } + + #[tokio::test] + async fn mock_is_deterministic_and_nonempty() { + let msgs = vec![msg("user", "Alice works at OpenAI"), msg("assistant", "Noted")]; + let a = MockSummarizer.summarize(&msgs).await.unwrap(); + let b = MockSummarizer.summarize(&msgs).await.unwrap(); + assert_eq!(a, b, "mock summarizer must be deterministic"); + assert!(!a.is_empty()); + assert!(a.contains("Alice works at OpenAI")); + } + + #[tokio::test] + async fn mock_model_label_is_mock() { + assert_eq!(MockSummarizer.model(), "mock"); + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::models::Message; + use serde_json::json; + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + fn msg(role: &str, content: &str) -> Message { + Message { + id: None, + role: role.into(), + content: content.into(), + timestamp: None, + embedding_status: None, + } + } + + fn chat_response(content: &str) -> serde_json::Value { + json!({ "choices": [{ "message": { "content": content } }] }) + } + + #[tokio::test] + async fn openai_summarizer_returns_summary_text() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/v1/chat/completions")) + .respond_with( + ResponseTemplate::new(200) + .set_body_json(chat_response("Alice works at OpenAI.")), + ) + .mount(&server) + .await; + + let s = OpenAISummarizer::new_with_base_url("sk-test".into(), server.uri()); + let out = s.summarize(&[msg("user", "Alice works at OpenAI")]).await.unwrap(); + assert_eq!(out, "Alice works at OpenAI."); + assert_eq!(s.model(), "gpt-4o-mini"); + } + + #[tokio::test] + async fn openai_summarizer_retries_on_rate_limit() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/v1/chat/completions")) + .respond_with(ResponseTemplate::new(429)) + .up_to_n_times(2) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/v1/chat/completions")) + .respond_with( + ResponseTemplate::new(200).set_body_json(chat_response("ok")), + ) + .mount(&server) + .await; + + let s = OpenAISummarizer::new_with_base_url("sk-test".into(), server.uri()); + assert_eq!(s.summarize(&[msg("user", "hi")]).await.unwrap(), "ok"); + } + + #[tokio::test] + async fn openai_summarizer_exhausts_retries() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/v1/chat/completions")) + .respond_with(ResponseTemplate::new(429)) + .mount(&server) + .await; + + let s = OpenAISummarizer::new_with_base_url("sk-test".into(), server.uri()); + let err = s.summarize(&[msg("user", "hi")]).await.unwrap_err(); + assert!(matches!(err, SummarizeError::RateLimitExceeded { .. })); + } +} diff --git a/src/lib.rs b/src/lib.rs index 1fe24b4..e05cdce 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -2,6 +2,7 @@ pub mod app; pub mod assembler; pub mod cluster; pub mod config; +pub mod consolidation; pub mod core; pub mod embedding; pub mod knowledge; diff --git a/src/metrics.rs b/src/metrics.rs index efb8f41..3e35da3 100644 --- a/src/metrics.rs +++ b/src/metrics.rs @@ -28,6 +28,11 @@ pub struct AppMetrics { global_entities: IntGauge, global_relationships: IntGauge, global_conflicts: IntGauge, + consolidations_total: IntCounter, + messages_consolidated_total: IntCounter, + summaries: IntGauge, + consolidation_queue_size: IntGauge, + summarization_duration_seconds: HistogramVec, } impl AppMetrics { @@ -169,6 +174,39 @@ impl AppMetrics { ))?; registry.register(Box::new(global_conflicts.clone()))?; + let consolidations_total = IntCounter::with_opts(Opts::new( + "consolidations_total", + "Total number of consolidation operations applied.", + ))?; + registry.register(Box::new(consolidations_total.clone()))?; + + let messages_consolidated_total = IntCounter::with_opts(Opts::new( + "messages_consolidated_total", + "Total number of raw messages consumed by consolidation.", + ))?; + registry.register(Box::new(messages_consolidated_total.clone()))?; + + let summaries = IntGauge::with_opts(Opts::new( + "summaries", + "Current total number of consolidated summaries across all sessions.", + ))?; + registry.register(Box::new(summaries.clone()))?; + + let consolidation_queue_size = IntGauge::with_opts(Opts::new( + "consolidation_queue_size", + "Current number of pending consolidation jobs.", + ))?; + registry.register(Box::new(consolidation_queue_size.clone()))?; + + let summarization_duration_seconds = HistogramVec::new( + HistogramOpts::new( + "summarization_duration_seconds", + "Duration of LLM summarization calls in seconds.", + ), + &["model"], + )?; + registry.register(Box::new(summarization_duration_seconds.clone()))?; + Ok(Self { registry, messages_added_total, @@ -191,6 +229,11 @@ impl AppMetrics { global_entities, global_relationships, global_conflicts, + consolidations_total, + messages_consolidated_total, + summaries, + consolidation_queue_size, + summarization_duration_seconds, }) } @@ -280,6 +323,28 @@ impl AppMetrics { self.global_conflicts.set(count as i64); } + pub fn increment_consolidations(&self) { + self.consolidations_total.inc(); + } + + pub fn increment_messages_consolidated(&self, n: u64) { + self.messages_consolidated_total.inc_by(n); + } + + pub fn set_summaries(&self, n: usize) { + self.summaries.set(n as i64); + } + + pub fn set_consolidation_queue_size(&self, n: usize) { + self.consolidation_queue_size.set(n as i64); + } + + pub fn start_summarization_timer(&self, model: &str) -> HistogramTimer { + self.summarization_duration_seconds + .with_label_values(&[model]) + .start_timer() + } + pub fn render(&self) -> Result { let mut buffer = Vec::new(); let encoder = TextEncoder::new(); @@ -318,4 +383,18 @@ mod tests { assert!(t.contains("engram_global_relationships")); assert!(t.contains("engram_global_conflicts")); } + + #[test] + fn renders_consolidation_metrics() { + let m = AppMetrics::new().unwrap(); + m.increment_consolidations(); + m.increment_messages_consolidated(30); + m.set_summaries(2); + m.set_consolidation_queue_size(1); + let t = m.render().unwrap(); + assert!(t.contains("engram_consolidations_total")); + assert!(t.contains("engram_messages_consolidated_total")); + assert!(t.contains("engram_summaries")); + assert!(t.contains("engram_consolidation_queue_size")); + } } \ No newline at end of file diff --git a/src/raft/recovery.rs b/src/raft/recovery.rs index 8d47dd9..451df0a 100644 --- a/src/raft/recovery.rs +++ b/src/raft/recovery.rs @@ -1,5 +1,6 @@ use std::sync::Arc; +use crate::consolidation::store::ConsolidatedMemoryStore; use crate::core::{CoreMemoryStore, ShortTermMemory}; use crate::knowledge::graph::KnowledgeGraph; use crate::raft::snapshot::EngramSnapshot; @@ -13,10 +14,12 @@ pub async fn recover_state_machine( sm: &EngStateMachineStore, short_term: Arc, core_memory: Arc, + consolidated: Arc, ) -> anyhow::Result<()> { // 1. Always flush first so stale state from a prior run cannot bleed through. short_term.restore_all(vec![]).await?; core_memory.restore_all(vec![]).await?; + consolidated.restore_all(vec![]).await?; // 2. Load the persisted snapshot, if present. let Some((meta, bytes)) = sm.load_snapshot_for_recovery()? else { @@ -30,6 +33,7 @@ pub async fn recover_state_machine( short_term.restore_all(st_sessions).await?; let cm_sessions = payload.core_memory.into_iter().map(|s| (s.session_id, s.facts)).collect(); core_memory.restore_all(cm_sessions).await?; + consolidated.restore_all(payload.consolidated).await?; let graph = KnowledgeGraph::from_snapshot(payload.knowledge_graph); sm.restore_applied_for_recovery(meta, graph).await; @@ -43,6 +47,7 @@ mod tests { use tokio::sync::{mpsc, RwLock}; use redb::Database; + use crate::consolidation::store::{ConsolidatedMemoryStore, InMemoryConsolidatedStore}; use crate::core::{CoreMemoryStore, InMemoryCoreMemoryStore, InMemoryStore, InMemoryVectorStore, ShortTermMemory}; use crate::knowledge::graph::KnowledgeGraph; use crate::raft::recovery::recover_state_machine; @@ -57,8 +62,9 @@ mod tests { let (ktx, _krx) = mpsc::channel(10); let kg = Arc::new(RwLock::new(KnowledgeGraph::new())); let gg = Arc::new(RwLock::new(crate::knowledge::global::GlobalGraph::new())); + let consolidated: Arc = Arc::new(InMemoryConsolidatedStore::default()); let metrics = Arc::new(crate::metrics::AppMetrics::new().unwrap()); - let sm = EngStateMachineStore::new(st.clone(), cm.clone(), vs, etx, kg.clone(), ktx, db, gg, metrics); + let sm = EngStateMachineStore::new(st.clone(), cm.clone(), vs, etx, kg.clone(), ktx, db, gg, consolidated, metrics); (sm, st, cm, kg) } @@ -76,7 +82,8 @@ mod tests { timestamp: None, embedding_status: None, }).await.unwrap(); - recover_state_machine(&sm, st.clone() as Arc, cm_as_dyn(&cm)).await.unwrap(); + let cons: Arc = Arc::new(InMemoryConsolidatedStore::default()); + recover_state_machine(&sm, st.clone() as Arc, cm_as_dyn(&cm), cons).await.unwrap(); assert!(st.get_recent("stale", 10).await.unwrap().is_empty()); } @@ -93,7 +100,8 @@ mod tests { // Fresh state machine over the same db; recovery should restore the fact. let (sm, st, cm, _kg2) = build(db.clone()); - recover_state_machine(&sm, st as Arc, cm.clone() as Arc).await.unwrap(); + let cons: Arc = Arc::new(InMemoryConsolidatedStore::default()); + recover_state_machine(&sm, st as Arc, cm.clone() as Arc, cons).await.unwrap(); assert_eq!(cm.get_facts("s1").await.unwrap(), vec!["remember me".to_string()]); } } diff --git a/src/raft/snapshot.rs b/src/raft/snapshot.rs index 4b7065b..603ea65 100644 --- a/src/raft/snapshot.rs +++ b/src/raft/snapshot.rs @@ -4,7 +4,7 @@ use crate::knowledge::graph::GraphSnapshot; use crate::models::Message; /// Snapshot schema version. Bump when the payload layout changes incompatibly. -pub const SNAPSHOT_VERSION: u32 = 2; +pub const SNAPSHOT_VERSION: u32 = 3; #[derive(Debug, Clone, Serialize, Deserialize)] pub struct SessionMessages { @@ -34,6 +34,9 @@ pub struct EngramSnapshot { pub visibility: Vec<(String, crate::knowledge::global::Visibility)>, #[serde(default)] pub session_agents: Vec<(String, String)>, + /// Per-session consolidated summaries. Added in v3; v1/v2 snapshots load with an empty map. + #[serde(default)] + pub consolidated: Vec<(String, Vec)>, } impl EngramSnapshot { @@ -61,12 +64,13 @@ mod tests { global_graph: None, visibility: vec![], session_agents: vec![], + consolidated: vec![], } } #[test] - fn snapshot_carries_version_two() { - assert_eq!(sample().version, 2); + fn snapshot_carries_version_three() { + assert_eq!(sample().version, 3); } #[test] @@ -74,10 +78,11 @@ mod tests { let snap = sample(); let bytes = snap.to_bytes().unwrap(); let back = EngramSnapshot::from_bytes(&bytes).unwrap(); - assert_eq!(back.version, 2); + assert_eq!(back.version, 3); assert_eq!(back.core_memory[0].facts, vec!["f".to_string()]); assert!(back.global_graph.is_none()); assert!(back.visibility.is_empty()); + assert!(back.consolidated.is_empty()); } #[test] @@ -89,7 +94,7 @@ mod tests { } #[test] - fn snapshot_version_is_two_and_carries_global_and_visibility() { + fn snapshot_carries_global_and_visibility() { let snap = EngramSnapshot { version: SNAPSHOT_VERSION, short_term: vec![], @@ -98,8 +103,9 @@ mod tests { global_graph: Some(crate::knowledge::global::GlobalGraphSnapshot::default()), visibility: vec![("s1".into(), crate::knowledge::global::Visibility::Shared)], session_agents: vec![("s1".into(), "agent-7".into())], + consolidated: vec![], }; - assert_eq!(snap.version, 2); + assert_eq!(snap.version, 3); let bytes = snap.to_bytes().unwrap(); let back = EngramSnapshot::from_bytes(&bytes).unwrap(); assert!(back.global_graph.is_some()); @@ -114,5 +120,44 @@ mod tests { assert!(back.global_graph.is_none()); assert!(back.visibility.is_empty()); assert!(back.session_agents.is_empty()); + assert!(back.consolidated.is_empty()); + } + + #[test] + fn snapshot_version_is_three_and_carries_consolidated() { + let snap = EngramSnapshot { + version: SNAPSHOT_VERSION, + short_term: vec![], + core_memory: vec![], + knowledge_graph: crate::knowledge::graph::GraphSnapshot::default(), + global_graph: None, + visibility: vec![], + session_agents: vec![], + consolidated: vec![( + "s1".to_string(), + vec![crate::consolidation::store::Summary { + id: "u1".into(), + text: "t".into(), + created_at_index: 1, + consumed_message_ids: vec!["m1".into()], + consumed_count: 1, + model: "mock".into(), + prompt_version: "summarize_v1".into(), + }], + )], + }; + assert_eq!(snap.version, 3); + let bytes = snap.to_bytes().unwrap(); + let back = EngramSnapshot::from_bytes(&bytes).unwrap(); + assert_eq!(back.consolidated.len(), 1); + assert_eq!(back.consolidated[0].1[0].id, "u1"); + } + + #[test] + fn v2_snapshot_without_consolidated_still_loads() { + // v2 JSON without the consolidated field must load cleanly. + let v2 = r#"{"version":2,"short_term":[],"core_memory":[],"knowledge_graph":{"sessions":[],"processed":[]},"global_graph":null,"visibility":[],"session_agents":[]}"#; + let back = EngramSnapshot::from_bytes(v2.as_bytes()).unwrap(); + assert!(back.consolidated.is_empty()); } } diff --git a/src/raft/state_machine.rs b/src/raft/state_machine.rs index 508bef5..ee9b2bd 100644 --- a/src/raft/state_machine.rs +++ b/src/raft/state_machine.rs @@ -9,6 +9,7 @@ use openraft::{ }; use redb::{Database, TableDefinition}; +use crate::consolidation::store::{ConsolidatedMemoryStore, Summary}; use crate::core::{CoreMemoryStore, ShortTermMemory}; use crate::knowledge::global::{GlobalGraph, Visibility}; use crate::knowledge::graph::KnowledgeGraph; @@ -37,6 +38,7 @@ struct SmInner { knowledge_graph: Arc>, knowledge_tx: mpsc::Sender, global_graph: Arc>, + consolidated: Arc, visibility: Arc>>, session_agents: Arc>>, metrics: Arc, @@ -56,6 +58,7 @@ impl EngStateMachineStore { knowledge_tx: mpsc::Sender, db: Arc, global_graph: Arc>, + consolidated: Arc, metrics: Arc, ) -> Self { { @@ -73,6 +76,7 @@ impl EngStateMachineStore { knowledge_graph, knowledge_tx, global_graph, + consolidated, visibility: Arc::new(RwLock::new(HashMap::new())), session_agents: Arc::new(RwLock::new(HashMap::new())), metrics, @@ -82,10 +86,6 @@ impl EngStateMachineStore { } } - pub(crate) fn inner_handle(&self) -> Arc> { - self.inner.clone() - } - /// Returns `(meta, payload_bytes)` of the persisted snapshot, if any. /// Called at startup before the Raft node starts (uncontended). pub(crate) fn load_snapshot_for_recovery( @@ -163,6 +163,11 @@ async fn build_payload(inner: &SmInner) -> Result<(EngramSnapshot, SnapshotMeta< .iter() .map(|(k, v)| (k.clone(), v.clone())) .collect(); + let consolidated = inner + .consolidated + .dump_all() + .await + .map_err(|e| sm_io_err(ErrorVerb::Read, e.to_string()))?; let payload = EngramSnapshot { version: crate::raft::snapshot::SNAPSHOT_VERSION, @@ -172,6 +177,7 @@ async fn build_payload(inner: &SmInner) -> Result<(EngramSnapshot, SnapshotMeta< global_graph, visibility, session_agents, + consolidated, }; let snapshot_id = format!( "{}-{}", @@ -246,7 +252,7 @@ impl RaftStateMachine for EngStateMachineStore { I::IntoIter: Send, { // Clone Arcs once so the lock is not held across async apply_cmd calls. - let (short_term, core_memory, embedding_tx, knowledge_graph, knowledge_tx, global_graph, visibility, session_agents, metrics) = { + let (short_term, core_memory, embedding_tx, knowledge_graph, knowledge_tx, global_graph, consolidated, visibility, session_agents, metrics) = { let inner = self.inner.lock().await; ( inner.short_term.clone(), @@ -255,6 +261,7 @@ impl RaftStateMachine for EngStateMachineStore { inner.knowledge_graph.clone(), inner.knowledge_tx.clone(), inner.global_graph.clone(), + inner.consolidated.clone(), inner.visibility.clone(), inner.session_agents.clone(), inner.metrics.clone(), @@ -272,7 +279,7 @@ impl RaftStateMachine for EngStateMachineStore { last_membership = Some(StoredMembership::new(Some(entry.log_id.clone()), mem.clone())); } if let EntryPayload::Normal(cmd) = entry.payload { - apply_cmd(cmd, &short_term, &core_memory, &embedding_tx, &knowledge_graph, &knowledge_tx, &global_graph, &visibility, &session_agents, &metrics, index).await; + apply_cmd(cmd, &short_term, &core_memory, &embedding_tx, &knowledge_graph, &knowledge_tx, &global_graph, &consolidated, &visibility, &session_agents, &metrics, index).await; } responses.push(CommandResponse::default()); } @@ -310,7 +317,7 @@ impl RaftStateMachine for EngStateMachineStore { .map_err(|e| sm_io_err(ErrorVerb::Read, e.to_string()))?; // Clone store/graph handles under a short lock so we don't hold it across awaits. - let (short_term, core_memory, knowledge_graph, global_graph, visibility, session_agents, db) = { + let (short_term, core_memory, knowledge_graph, global_graph, visibility, session_agents, consolidated, db) = { let inner = self.inner.lock().await; ( inner.short_term.clone(), @@ -319,6 +326,7 @@ impl RaftStateMachine for EngStateMachineStore { inner.global_graph.clone(), inner.visibility.clone(), inner.session_agents.clone(), + inner.consolidated.clone(), inner.db.clone(), ) }; @@ -346,6 +354,10 @@ impl RaftStateMachine for EngStateMachineStore { agents.insert(session_id, agent_id); } } + consolidated + .restore_all(payload.consolidated) + .await + .map_err(|e| sm_io_err(ErrorVerb::Write, e.to_string()))?; persist_snapshot(&db, meta, snapshot.get_ref())?; @@ -379,6 +391,7 @@ async fn apply_cmd( knowledge_graph: &Arc>, knowledge_tx: &mpsc::Sender, global_graph: &Arc>, + consolidated: &Arc, visibility: &Arc>>, session_agents: &Arc>>, metrics: &Arc, @@ -420,6 +433,9 @@ async fn apply_cmd( if let Err(e) = core_memory.delete_session(&session_id).await { tracing::error!(error = %e, session_id = %session_id, "failed to delete session from core memory"); } + if let Err(e) = consolidated.delete_session(&session_id).await { + tracing::error!(error = %e, session_id = %session_id, "failed to delete session from consolidated store"); + } // Signal embedding worker to delete from local LanceDB. let _ = embedding_tx.try_send(EmbeddingJob::DeleteSession { session_id: session_id.clone() }); knowledge_graph.write().await.delete_session(&session_id); @@ -471,6 +487,41 @@ async fn apply_cmd( session_agents.write().await.insert(session_id, agent); } } + MemoryCommand::ApplySummary { + session_id, + summary_id, + summary_text, + consumed_message_ids, + model, + prompt_version, + } => { + // One atomic replicated transition: store summary, trim raw messages, update metrics. + // The store no-ops a duplicate id, so replaying this entry on a lagging follower + // is safe — already-removed ids are gone and the store just skips the duplicate. + let summary = Summary { + id: summary_id, + text: summary_text, + created_at_index: index, + consumed_count: consumed_message_ids.len() as u64, + consumed_message_ids: consumed_message_ids.clone(), + model, + prompt_version, + }; + if let Err(e) = consolidated.add_summary(&session_id, summary).await { + tracing::error!(error = %e, session_id = %session_id, "failed to store summary"); + } + if let Err(e) = short_term.remove_messages(&session_id, &consumed_message_ids).await { + tracing::error!(error = %e, session_id = %session_id, "failed to trim consolidated messages"); + } + // ponytail: per-message vector delete deferred; stale vectors are a search-quality + // nit, not a correctness bug. Add EmbeddingJob::DeleteMessages when it matters. + tracing::debug!(session_id = %session_id, count = consumed_message_ids.len(), "vector cleanup deferred until session delete"); + metrics.increment_consolidations(); + metrics.increment_messages_consolidated(consumed_message_ids.len() as u64); + if let Ok(all) = consolidated.get_summaries(&session_id).await { + metrics.set_summaries(all.len()); + } + } MemoryCommand::NoOp => {} } } @@ -478,6 +529,7 @@ async fn apply_cmd( #[cfg(test)] mod tests { use super::*; + use crate::consolidation::store::{ConsolidatedMemoryStore, InMemoryConsolidatedStore}; use crate::core::{InMemoryCoreMemoryStore, InMemoryStore, InMemoryVectorStore}; use crate::knowledge::global::{GlobalGraph, Visibility}; use crate::knowledge::graph::KnowledgeGraph; @@ -495,6 +547,7 @@ mod tests { Arc>, Arc, Arc>, + Arc, tempfile::TempDir, ) { let short_term = Arc::new(InMemoryStore::default()); @@ -504,6 +557,7 @@ mod tests { let (know_tx, know_rx) = mpsc::channel(10); let kg = Arc::new(RwLock::new(KnowledgeGraph::new())); let gg = Arc::new(RwLock::new(GlobalGraph::new())); + let consolidated = Arc::new(InMemoryConsolidatedStore::default()); let dir = tempfile::tempdir().unwrap(); let db = Arc::new(redb::Database::create(dir.path().join("sm.redb")).unwrap()); let metrics = Arc::new(crate::metrics::AppMetrics::new().unwrap()); @@ -516,9 +570,10 @@ mod tests { know_tx, db, gg.clone(), + consolidated.clone() as Arc, metrics, ); - (sm, short_term, embed_rx, know_rx, kg, core_memory, gg, dir) + (sm, short_term, embed_rx, know_rx, kg, core_memory, gg, consolidated, dir) } fn make_entry(index: u64, cmd: MemoryCommand) -> openraft::Entry { @@ -530,7 +585,7 @@ mod tests { #[tokio::test] async fn add_message_writes_to_short_term() { - let (mut sm, short_term, _embed, _know, _kg, _cm, _gg, _dir) = make_sm(); + let (mut sm, short_term, _embed, _know, _kg, _cm, _gg, _cons, _dir) = make_sm(); sm.apply(vec![make_entry( 0, MemoryCommand::AddMessage { @@ -552,7 +607,7 @@ mod tests { #[tokio::test] async fn add_message_enqueues_embedding_job() { - let (mut sm, _st, mut embed_rx, _know, _kg, _cm, _gg, _dir) = make_sm(); + let (mut sm, _st, mut embed_rx, _know, _kg, _cm, _gg, _cons, _dir) = make_sm(); sm.apply(vec![make_entry( 0, MemoryCommand::AddMessage { @@ -573,7 +628,7 @@ mod tests { #[tokio::test] async fn delete_session_clears_redis_and_enqueues_lancedb_delete() { - let (mut sm, short_term, mut embed_rx, _know, _kg, _cm, _gg, _dir) = make_sm(); + let (mut sm, short_term, mut embed_rx, _know, _kg, _cm, _gg, _cons, _dir) = make_sm(); sm.apply(vec![ make_entry( 0, @@ -600,7 +655,7 @@ mod tests { #[tokio::test] async fn noop_command_is_ignored() { - let (mut sm, short_term, _embed, _know, _kg, _cm, _gg, _dir) = make_sm(); + let (mut sm, short_term, _embed, _know, _kg, _cm, _gg, _cons, _dir) = make_sm(); sm.apply(vec![make_entry(0, MemoryCommand::NoOp)]).await.unwrap(); let msgs = short_term.get_recent("any", 10).await.unwrap(); assert_eq!(msgs.len(), 0); @@ -608,7 +663,7 @@ mod tests { #[tokio::test] async fn add_message_enqueues_knowledge_job() { - let (mut sm, _st, _embed, mut know_rx, _kg, _cm, _gg, _dir) = make_sm(); + let (mut sm, _st, _embed, mut know_rx, _kg, _cm, _gg, _cons, _dir) = make_sm(); sm.apply(vec![make_entry(0, MemoryCommand::AddMessage { session_id: "s1".into(), message: MessagePayload { @@ -625,7 +680,7 @@ mod tests { #[tokio::test] async fn add_knowledge_updates_graph() { - let (mut sm, _st, _embed, _know, kg, _cm, _gg, _dir) = make_sm(); + let (mut sm, _st, _embed, _know, kg, _cm, _gg, _cons, _dir) = make_sm(); sm.apply(vec![make_entry(0, MemoryCommand::AddKnowledge { session_id: "s1".into(), message_id: "m1".into(), @@ -646,7 +701,7 @@ mod tests { #[tokio::test] async fn delete_session_clears_knowledge_graph() { - let (mut sm, _st, _embed, _know, kg, _cm, _gg, _dir) = make_sm(); + let (mut sm, _st, _embed, _know, kg, _cm, _gg, _cons, _dir) = make_sm(); sm.apply(vec![make_entry(0, MemoryCommand::AddKnowledge { session_id: "s1".into(), message_id: "m1".into(), entities: vec![Entity { name: "Alice".into(), entity_type: "Person".into(), attributes: HashMap::new() }], @@ -660,14 +715,14 @@ mod tests { #[tokio::test] async fn install_snapshot_sets_last_applied_to_meta_log_id() { - let (mut src, _st, _e, _k, _kg, _cm, _gg, _dir) = make_sm(); + let (mut src, _st, _e, _k, _kg, _cm, _gg, _cons, _dir) = make_sm(); src.apply(vec![make_entry(7, MemoryCommand::AddFact { session_id: "s1".into(), fact: "f".into(), })]).await.unwrap(); let mut builder = src.get_snapshot_builder().await; let snap = builder.build_snapshot().await.unwrap(); - let (mut dst, dst_st, _e2, _k2, dst_kg, _cm2, _gg2, _dir2) = make_sm(); + let (mut dst, dst_st, _e2, _k2, dst_kg, _cm2, _gg2, _cons2, _dir2) = make_sm(); let mut buf = dst.begin_receiving_snapshot().await.unwrap(); *buf = std::io::Cursor::new(snap.snapshot.get_ref().clone()); dst.install_snapshot(&snap.meta, buf).await.unwrap(); @@ -680,7 +735,7 @@ mod tests { #[tokio::test] async fn apply_build_install_reproduces_state() { - let (mut src, _st, _e, _k, _kg, src_cm, _gg, _dir) = make_sm(); + let (mut src, _st, _e, _k, _kg, src_cm, _gg, _cons, _dir) = make_sm(); src.apply(vec![ make_entry(0, MemoryCommand::AddKnowledge { session_id: "s1".into(), message_id: "m1".into(), @@ -697,7 +752,7 @@ mod tests { let mut builder = src.get_snapshot_builder().await; let snap = builder.build_snapshot().await.unwrap(); - let (mut dst, _st2, _e2, _k2, dst_kg, dst_cm, _gg2, _dir2) = make_sm(); + let (mut dst, _st2, _e2, _k2, dst_kg, dst_cm, _gg2, _cons2, _dir2) = make_sm(); let mut buf = dst.begin_receiving_snapshot().await.unwrap(); *buf = std::io::Cursor::new(snap.snapshot.get_ref().clone()); dst.install_snapshot(&snap.meta, buf).await.unwrap(); @@ -709,7 +764,7 @@ mod tests { #[tokio::test] async fn build_snapshot_meta_index_equals_last_applied() { - let (mut sm, _st, _e, _k, _kg, _cm, _gg, _dir) = make_sm(); + let (mut sm, _st, _e, _k, _kg, _cm, _gg, _cons, _dir) = make_sm(); for i in 0..=4u64 { sm.apply(vec![make_entry(i, MemoryCommand::AddFact { session_id: "s1".into(), fact: format!("f{i}"), @@ -722,7 +777,7 @@ mod tests { #[tokio::test] async fn build_then_get_current_snapshot_returns_same_index() { - let (mut sm, _st, _e, _k, _kg, _cm, _gg, _dir) = make_sm(); + let (mut sm, _st, _e, _k, _kg, _cm, _gg, _cons, _dir) = make_sm(); sm.apply(vec![make_entry(0, MemoryCommand::AddFact { session_id: "s1".into(), fact: "f".into(), })]).await.unwrap(); @@ -734,7 +789,7 @@ mod tests { #[tokio::test] async fn snapshot_payload_contains_applied_state() { - let (mut sm, _st, _e, _k, _kg, _cm, _gg, _dir) = make_sm(); + let (mut sm, _st, _e, _k, _kg, _cm, _gg, _cons, _dir) = make_sm(); sm.apply(vec![make_entry(0, MemoryCommand::AddKnowledge { session_id: "s1".into(), message_id: "m1".into(), entities: vec![Entity { name: "Alice".into(), entity_type: "Person".into(), attributes: HashMap::new() }], @@ -747,14 +802,14 @@ mod tests { let mut builder = sm.get_snapshot_builder().await; let snap = builder.build_snapshot().await.unwrap(); let payload = crate::raft::snapshot::EngramSnapshot::from_bytes(snap.snapshot.get_ref()).unwrap(); - assert_eq!(payload.version, 2); + assert_eq!(payload.version, 3); assert!(payload.knowledge_graph.sessions.iter().any(|s| s.session_id == "s1")); assert!(payload.core_memory.iter().any(|s| s.facts.contains(&"likes coffee".to_string()))); } #[tokio::test] async fn shared_session_knowledge_reaches_global_graph() { - let (mut sm, _st, _e, _k, _kg, _cm, gg, _dir) = make_sm(); + let (mut sm, _st, _e, _k, _kg, _cm, gg, _cons, _dir) = make_sm(); sm.apply(vec![make_entry(0, MemoryCommand::SetSessionVisibility { session_id: "s1".into(), visibility: Visibility::Shared, })]).await.unwrap(); @@ -768,7 +823,7 @@ mod tests { #[tokio::test] async fn private_session_knowledge_stays_out_of_global_graph() { - let (mut sm, _st, _e, _k, _kg, _cm, gg, _dir) = make_sm(); + let (mut sm, _st, _e, _k, _kg, _cm, gg, _cons, _dir) = make_sm(); sm.apply(vec![make_entry(0, MemoryCommand::AddKnowledge { session_id: "s1".into(), message_id: "m1".into(), entities: vec![Entity { name: "Secret".into(), entity_type: "Person".into(), attributes: HashMap::new() }], @@ -779,7 +834,7 @@ mod tests { #[tokio::test] async fn becoming_shared_backmerges_existing_session_knowledge() { - let (mut sm, _st, _e, _k, _kg, _cm, gg, _dir) = make_sm(); + let (mut sm, _st, _e, _k, _kg, _cm, gg, _cons, _dir) = make_sm(); sm.apply(vec![make_entry(0, MemoryCommand::AddKnowledge { session_id: "s1".into(), message_id: "m1".into(), entities: vec![Entity { name: "Alice".into(), entity_type: "Person".into(), attributes: HashMap::new() }], @@ -793,7 +848,7 @@ mod tests { #[tokio::test] async fn delete_session_prunes_global_contributions() { - let (mut sm, _st, _e, _k, _kg, _cm, gg, _dir) = make_sm(); + let (mut sm, _st, _e, _k, _kg, _cm, gg, _cons, _dir) = make_sm(); for (idx, sid) in [(0u64, "s1"), (1, "s2")] { sm.apply(vec![make_entry(idx, MemoryCommand::SetSessionVisibility { session_id: sid.into(), visibility: Visibility::Shared, @@ -817,7 +872,7 @@ mod tests { #[tokio::test] async fn registered_agent_id_flows_into_global_provenance() { - let (mut sm, _st, _e, _k, _kg, _cm, gg, _dir) = make_sm(); + let (mut sm, _st, _e, _k, _kg, _cm, gg, _cons, _dir) = make_sm(); sm.apply(vec![make_entry(0, MemoryCommand::RegisterSession { session_id: "s1".into(), agent_id: Some("agent-7".into()), })]).await.unwrap(); @@ -834,7 +889,7 @@ mod tests { #[tokio::test] async fn visibility_transitions_are_fully_reversible() { - let (mut sm, _st, _e, _k, _kg, _cm, gg, _dir) = make_sm(); + let (mut sm, _st, _e, _k, _kg, _cm, gg, _cons, _dir) = make_sm(); sm.apply(vec![make_entry(0, MemoryCommand::AddKnowledge { session_id: "s1".into(), message_id: "m1".into(), entities: vec![ @@ -858,4 +913,89 @@ mod tests { sm.apply(vec![vis(4, Visibility::Private)]).await.unwrap(); assert!(gg.read().await.all_entities().is_empty()); } + + #[tokio::test] + async fn apply_summary_stores_summary_and_trims_messages() { + let (mut sm, short_term, _embed, _know, _kg, _cm, _gg, cons, _dir) = make_sm(); + for (i, id) in [(0u64, "m1"), (1, "m2"), (2, "m3")] { + sm.apply(vec![make_entry(i, MemoryCommand::AddMessage { + session_id: "s1".into(), + message: MessagePayload { id: id.into(), role: "user".into(), content: format!("msg {id}"), timestamp: chrono::Utc::now() }, + })]).await.unwrap(); + } + sm.apply(vec![make_entry(3, MemoryCommand::ApplySummary { + session_id: "s1".into(), + summary_id: "u1".into(), + summary_text: "summary text".into(), + consumed_message_ids: vec!["m1".into(), "m2".into()], + model: "mock".into(), + prompt_version: "summarize_v1".into(), + })]).await.unwrap(); + + let summaries = cons.get_summaries("s1").await.unwrap(); + assert_eq!(summaries.len(), 1); + assert_eq!(summaries[0].id, "u1"); + assert_eq!(summaries[0].created_at_index, 3, "ordering index is the log index"); + assert_eq!(summaries[0].consumed_message_ids, vec!["m1".to_string(), "m2".to_string()]); + + let remaining = short_term.get_recent("s1", 10).await.unwrap(); + assert_eq!(remaining.len(), 1); + assert_eq!(remaining[0].content, "msg m3"); + } + + #[tokio::test] + async fn apply_summary_is_idempotent_by_id() { + let (mut sm, short_term, _embed, _know, _kg, _cm, _gg, cons, _dir) = make_sm(); + for (i, id) in [(0u64, "m1"), (1, "m2")] { + sm.apply(vec![make_entry(i, MemoryCommand::AddMessage { + session_id: "s1".into(), + message: MessagePayload { id: id.into(), role: "user".into(), content: id.into(), timestamp: chrono::Utc::now() }, + })]).await.unwrap(); + } + let cmd = MemoryCommand::ApplySummary { + session_id: "s1".into(), summary_id: "dup".into(), summary_text: "t".into(), + consumed_message_ids: vec!["m1".into()], model: "mock".into(), prompt_version: "summarize_v1".into(), + }; + sm.apply(vec![make_entry(2, cmd.clone())]).await.unwrap(); + sm.apply(vec![make_entry(3, cmd)]).await.unwrap(); + + assert_eq!(cons.get_summaries("s1").await.unwrap().len(), 1, "no duplicate summary"); + // m1 already trimmed; second apply must not trim m2 + let remaining = short_term.get_recent("s1", 10).await.unwrap(); + assert_eq!(remaining.len(), 1); + assert_eq!(remaining[0].content, "m2"); + } + + #[tokio::test] + async fn delete_session_clears_consolidated() { + let (mut sm, _st, _embed, _know, _kg, _cm, _gg, cons, _dir) = make_sm(); + sm.apply(vec![make_entry(0, MemoryCommand::ApplySummary { + session_id: "s1".into(), summary_id: "u1".into(), summary_text: "t".into(), + consumed_message_ids: vec![], model: "mock".into(), prompt_version: "summarize_v1".into(), + })]).await.unwrap(); + assert_eq!(cons.get_summaries("s1").await.unwrap().len(), 1); + sm.apply(vec![make_entry(1, MemoryCommand::DeleteSession { session_id: "s1".into() })]).await.unwrap(); + assert!(cons.get_summaries("s1").await.unwrap().is_empty()); + } + + #[tokio::test] + async fn snapshot_round_trip_preserves_summaries() { + let (mut sm, _st, _embed, _know, _kg, _cm, _gg, _cons, _dir) = make_sm(); + sm.apply(vec![make_entry(0, MemoryCommand::ApplySummary { + session_id: "s1".into(), summary_id: "u1".into(), summary_text: "kept".into(), + consumed_message_ids: vec![], model: "mock".into(), prompt_version: "summarize_v1".into(), + })]).await.unwrap(); + + let mut builder = sm.get_snapshot_builder().await; + let snap = builder.build_snapshot().await.unwrap(); + + let (mut sm2, _st2, _e2, _k2, _kg2, _cm2, _gg2, cons2, _dir2) = make_sm(); + let mut buf = sm2.begin_receiving_snapshot().await.unwrap(); + *buf = std::io::Cursor::new(snap.snapshot.get_ref().clone()); + sm2.install_snapshot(&snap.meta, buf).await.unwrap(); + + let restored = cons2.get_summaries("s1").await.unwrap(); + assert_eq!(restored.len(), 1); + assert_eq!(restored[0].text, "kept"); + } } diff --git a/src/raft/types.rs b/src/raft/types.rs index 2a7e4ab..077b759 100644 --- a/src/raft/types.rs +++ b/src/raft/types.rs @@ -51,6 +51,19 @@ pub enum MemoryCommand { }, /// Record an agent owner for a session (provenance for the global graph). RegisterSession { session_id: String, agent_id: Option }, + /// Apply a leader-produced summary: store it, trim the consumed raw messages. + /// One atomic replicated transition. Idempotent by `summary_id`. Only the leader + /// calls the LLM; followers apply this replicated artifact. `model` and + /// `prompt_version` are carried on the command (not read from node-local config) + /// so the stored summary is byte-identical on every node. + ApplySummary { + session_id: String, + summary_id: String, + summary_text: String, + consumed_message_ids: Vec, + model: String, + prompt_version: String, + }, /// No-op placeholder. Applied by the state machine without side effects. /// Reserved for future cluster operations (e.g., leadership probes). NoOp, @@ -146,4 +159,27 @@ mod tests { let back: MemoryCommand = serde_json::from_str(&json).unwrap(); assert!(matches!(back, MemoryCommand::SetSessionVisibility { visibility: Visibility::Shared, .. })); } + + #[test] + fn apply_summary_command_round_trips() { + let cmd = MemoryCommand::ApplySummary { + session_id: "s1".into(), + summary_id: "11111111-1111-1111-1111-111111111111".into(), + summary_text: "Alice works at OpenAI.".into(), + consumed_message_ids: vec!["m1".into(), "m2".into()], + model: "gpt-4o-mini".into(), + prompt_version: "summarize_v1".into(), + }; + let json = serde_json::to_string(&cmd).unwrap(); + let back: MemoryCommand = serde_json::from_str(&json).unwrap(); + match back { + MemoryCommand::ApplySummary { session_id, summary_id, consumed_message_ids, model, .. } => { + assert_eq!(session_id, "s1"); + assert_eq!(summary_id, "11111111-1111-1111-1111-111111111111"); + assert_eq!(consumed_message_ids.len(), 2); + assert_eq!(model, "gpt-4o-mini"); + } + _ => panic!("wrong variant"), + } + } } diff --git a/src/server.rs b/src/server.rs index 5bd5844..535021c 100644 --- a/src/server.rs +++ b/src/server.rs @@ -88,6 +88,31 @@ fn forward_to_redirect( MemoryServerError::Internal(format!("raft error: {e}")) } +/// Returns a `RedirectToLeader` error if this node is a follower. +/// +/// Unlike [`forward_to_redirect`], this consults the leader directly from Raft metrics +/// rather than a `client_write` rejection for handlers that don't submit a command +/// synchronously but still must route mutations to the leader (e.g. manual consolidate, +/// which enqueues an async job the leader runs). Returns `None` when this node is the +/// leader (or standalone), so the caller proceeds locally. +pub(crate) fn redirect_if_follower( + raft: &crate::raft::types::RaftHandle, + node_id: u64, + peer_http_addrs: &std::collections::HashMap, + path: &str, +) -> Option { + let leader = raft.metrics().borrow().current_leader; + if leader == Some(node_id) { + return None; + } + Some(match leader.and_then(|id| peer_http_addrs.get(&id)) { + Some(http_addr) => { + MemoryServerError::RedirectToLeader(format!("http://{http_addr}{path}")) + } + None => MemoryServerError::NoLeader, + }) +} + /// Submits a command to Raft and maps the result to a handler-ready `StatusCode`. /// /// Called only when `state.raft.is_some()`. Centralises the `client_write` call and @@ -134,6 +159,10 @@ pub struct AppState { pub knowledge_job_sender: tokio::sync::mpsc::Sender, /// Cluster-wide knowledge graph aggregating all Shared sessions. pub global_graph: Arc>, + /// Per-session consolidated summaries, shared with the Raft state machine. + pub consolidated: Arc, + /// Channel for handing consolidation jobs to the scheduler worker pool. + pub consolidation_tx: mpsc::Sender, } @@ -236,6 +265,14 @@ pub fn build_router(state: Arc) -> Router { .route("/sessions/{session_id}/knowledge/path", get(find_path)) .route("/sessions/{session_id}/knowledge/export", get(export_knowledge)) .route("/sessions/{session_id}/visibility", put(set_visibility)) + .route( + "/sessions/{session_id}/summaries", + get(crate::consolidation::handler::get_summaries), + ) + .route( + "/sessions/{session_id}/consolidate", + post(crate::consolidation::handler::post_consolidate), + ) .route("/knowledge/global", get(get_global)) .route("/knowledge/global/entities/{name}", get(get_global_entity)) .route("/knowledge/global/entities/{name}/sources", get(get_global_entity_sources)) @@ -369,7 +406,7 @@ async fn add_message( // Cluster mode: replicate through Raft. The state machine applies the command to // Redis and enqueues the embedding job on every node independently. if let Some(raft) = &state.raft { - return raft_write( + let status = raft_write( raft, MemoryCommand::AddMessage { session_id: session_id.clone(), @@ -378,7 +415,15 @@ async fn add_message( &state.peer_http_addrs, &format!("/sessions/{session_id}/messages"), ) - .await; + .await?; + // Nudge the consolidation scheduler. The worker re-checks leadership, the count + // threshold, and the in-flight guard, so a full queue or a follower is harmless. + let _ = state + .consolidation_tx + .try_send(crate::consolidation::scheduler::ConsolidationJob { + session_id: session_id.clone(), + }); + return Ok(status); } // Standalone mode: write directly to Redis and enqueue embedding. @@ -816,6 +861,13 @@ mod tests { global_graph: Arc::new(tokio::sync::RwLock::new( crate::knowledge::global::GlobalGraph::new(), )), + consolidated: Arc::new(crate::consolidation::store::InMemoryConsolidatedStore::default()), + consolidation_tx: { + let (tx, mut rx) = + tokio::sync::mpsc::channel::(16); + tokio::spawn(async move { while rx.recv().await.is_some() {} }); + tx + }, }) } @@ -1289,6 +1341,113 @@ mod tests { ); } + #[tokio::test] + async fn summaries_and_consolidate_routes_exist() { + let state = build_test_state(); + { + let s = crate::consolidation::store::Summary { + id: "u1".into(), + text: "kept".into(), + created_at_index: 1, + consumed_message_ids: vec!["m1".into()], + consumed_count: 1, + model: "mock".into(), + prompt_version: "summarize_v1".into(), + }; + state.consolidated.add_summary("s1", s).await.unwrap(); + } + let server = TestServer::new(build_router(state)).unwrap(); + + let got = server.get("/sessions/s1/summaries").await; + got.assert_status_ok(); + assert!(got.text().contains("kept")); + + // standalone (no raft): leader path always accepts. + let resp = server.post("/sessions/s1/consolidate").await; + assert!( + resp.status_code().is_success() || resp.status_code().as_u16() == 307, + "expected 2xx or 307 but got {}", + resp.status_code() + ); + } + + #[tokio::test] + async fn post_consolidate_drives_full_flow_to_summary_and_trim() { + // Wire a real scheduler to the same stores the router serves, mirroring app.rs + // standalone wiring, then drive the whole path over HTTP: POST /consolidate -> + // job -> mock summarize -> trim -> GET /summaries shows the result. + use crate::consolidation::scheduler::{ + consolidation_job_channel, spawn_consolidation_workers, + }; + use crate::consolidation::store::{ConsolidatedMemoryStore, InMemoryConsolidatedStore}; + use crate::core::ShortTermMemory; + use crate::knowledge::summarizer::{MockSummarizer, Summarizer}; + + let short_term = Arc::new(InMemoryStore::default()); + for i in 0..6 { + short_term + .add_message( + "s1", + crate::models::Message { + id: Some(format!("m{i}")), + role: "user".into(), + content: format!("content {i}"), + timestamp: None, + embedding_status: None, + }, + ) + .await + .unwrap(); + } + let consolidated = Arc::new(InMemoryConsolidatedStore::default()); + let metrics = Arc::new(AppMetrics::new().unwrap()); + let (consolidation_tx, rx) = consolidation_job_channel(16); + let summarizer: Arc = Arc::new(MockSummarizer); + // threshold 4, window 2: with 6 messages, summarize oldest 4, keep newest 2. + spawn_consolidation_workers( + summarizer, + None, + 0, + short_term.clone(), + consolidated.clone(), + metrics.clone(), + 4, + 2, + rx, + 1, + ); + + let base = build_test_state(); + let state = Arc::new(AppState { + short_term_memory: short_term.clone(), + consolidated: consolidated.clone(), + consolidation_tx, + ..(*base).clone() + }); + let server = TestServer::new(build_router(state)).unwrap(); + + server.post("/sessions/s1/consolidate").await.assert_status(StatusCode::ACCEPTED); + + // Worker runs async; poll until the summary lands. + let mut summaries: Vec = vec![]; + for _ in 0..40 { + summaries = consolidated.get_summaries("s1").await.unwrap(); + if !summaries.is_empty() { + break; + } + tokio::time::sleep(std::time::Duration::from_millis(25)).await; + } + assert_eq!(summaries.len(), 1, "consolidation produced a summary"); + assert_eq!(summaries[0].consumed_count, 4); + + let remaining = short_term.get_recent("s1", 10).await.unwrap(); + assert_eq!(remaining.len(), 2, "session trimmed back to the window"); + + let got = server.get("/sessions/s1/summaries").await; + got.assert_status_ok(); + assert!(got.text().contains(&summaries[0].id)); + } + #[tokio::test] async fn knowledge_routes_are_registered() { let server = TestServer::new(build_router(build_test_state())).unwrap(); diff --git a/src/stores/mod.rs b/src/stores/mod.rs index 1b6a074..2bf9c8a 100644 --- a/src/stores/mod.rs +++ b/src/stores/mod.rs @@ -1,7 +1,9 @@ mod lancedb; +mod redis_consolidated; mod redis_core_memory; mod redis_shortterm; pub use lancedb::LanceDBStore; +pub use redis_consolidated::RedisConsolidatedStore; pub use redis_core_memory::RedisCoreMemoryStore; pub use redis_shortterm::RedisShortTermMemory; diff --git a/src/stores/redis_consolidated.rs b/src/stores/redis_consolidated.rs new file mode 100644 index 0000000..ae58dac --- /dev/null +++ b/src/stores/redis_consolidated.rs @@ -0,0 +1,183 @@ +use std::error::Error as StdError; + +use async_trait::async_trait; +use futures::StreamExt; +use redis::{AsyncCommands, Client, aio::MultiplexedConnection}; + +use crate::consolidation::store::{ConsolidatedMemoryStore, Summary}; +use crate::core::MemoryError; + +#[derive(Debug, Clone)] +pub struct RedisConsolidatedStore { + connection: MultiplexedConnection, +} + +impl RedisConsolidatedStore { + pub fn new(connection: MultiplexedConnection) -> Self { + Self { connection } + } + + pub async fn connect(redis_url: &str) -> Result { + let client = Client::open(redis_url).map_err(memory_error)?; + let connection = client + .get_multiplexed_async_connection() + .await + .map_err(memory_error)?; + Ok(Self::new(connection)) + } + + fn session_key(session_id: &str) -> String { + format!("consolidated:{session_id}") + } +} + +#[async_trait] +impl ConsolidatedMemoryStore for RedisConsolidatedStore { + async fn add_summary(&self, session_id: &str, summary: Summary) -> Result<(), MemoryError> { + let key = Self::session_key(session_id); + let mut conn = self.connection.clone(); + + // Idempotent: skip if this summary id is already in the list. + let existing: Vec = conn.lrange(&key, 0, -1).await.map_err(memory_error)?; + for raw in &existing { + let s: Summary = serde_json::from_str(raw).map_err(memory_error)?; + if s.id == summary.id { + return Ok(()); + } + } + + let payload = serde_json::to_string(&summary).map_err(memory_error)?; + let _: usize = conn.rpush(&key, payload).await.map_err(memory_error)?; + Ok(()) + } + + async fn get_summaries(&self, session_id: &str) -> Result, MemoryError> { + let mut conn = self.connection.clone(); + let raw: Vec = conn + .lrange(Self::session_key(session_id), 0, -1) + .await + .map_err(memory_error)?; + raw.into_iter() + .map(|r| serde_json::from_str(&r).map_err(memory_error)) + .collect() + } + + async fn delete_session(&self, session_id: &str) -> Result<(), MemoryError> { + let mut conn = self.connection.clone(); + let _: usize = conn + .del(Self::session_key(session_id)) + .await + .map_err(memory_error)?; + Ok(()) + } + + async fn dump_all(&self) -> Result)>, MemoryError> { + let mut conn = self.connection.clone(); + let keys: Vec = { + let mut iter = conn + .scan_match::<_, String>("consolidated:*") + .await + .map_err(memory_error)?; + let mut collected = Vec::new(); + while let Some(key) = iter.next().await { + collected.push(key); + } + collected + }; + let mut out = Vec::new(); + for key in keys { + let session_id = session_id_from_key(&key); + let summaries = self.get_summaries(&session_id).await?; + out.push((session_id, summaries)); + } + Ok(out) + } + + async fn restore_all(&self, sessions: Vec<(String, Vec)>) -> Result<(), MemoryError> { + let mut conn = self.connection.clone(); + // Wipe existing data before restoring snapshot. + let existing: Vec = { + let mut iter = conn + .scan_match::<_, String>("consolidated:*") + .await + .map_err(memory_error)?; + let mut collected = Vec::new(); + while let Some(key) = iter.next().await { + collected.push(key); + } + collected + }; + for key in existing { + let _: usize = conn.del(key.as_str()).await.map_err(memory_error)?; + } + for (session_id, summaries) in sessions { + for summary in summaries { + self.add_summary(&session_id, summary).await?; + } + } + Ok(()) + } +} + +fn session_id_from_key(key: &str) -> String { + key.strip_prefix("consolidated:") + .unwrap_or(key) + .to_string() +} + +fn memory_error(error: impl StdError + Send + Sync + 'static) -> MemoryError { + MemoryError::Other(Box::new(error)) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::consolidation::store::{ConsolidatedMemoryStore, Summary}; + use testcontainers::{ + GenericImage, + core::{IntoContainerPort, WaitFor}, + runners::AsyncRunner, + }; + + const REDIS_PORT: u16 = 6379; + + async fn test_store() -> (RedisConsolidatedStore, testcontainers::ContainerAsync) { + let node = GenericImage::new("redis", "7.2.4") + .with_exposed_port(REDIS_PORT.tcp()) + .with_wait_for(WaitFor::message_on_stdout("Ready to accept connections")) + .start() + .await + .unwrap(); + let host = node.get_host().await.unwrap(); + let port = node.get_host_port_ipv4(REDIS_PORT.tcp()).await.unwrap(); + let url = format!("redis://{host}:{port}/"); + let store = RedisConsolidatedStore::connect(&url).await.unwrap(); + (store, node) + } + + fn summary(id: &str) -> Summary { + Summary { + id: id.into(), + text: "t".into(), + created_at_index: 1, + consumed_message_ids: vec!["m1".into()], + consumed_count: 1, + model: "mock".into(), + prompt_version: "summarize_v1".into(), + } + } + + #[tokio::test] + async fn redis_consolidated_round_trip_and_idempotent() { + let (store, _c) = test_store().await; + store.add_summary("s1", summary("a")).await.unwrap(); + store.add_summary("s1", summary("a")).await.unwrap(); // idempotent + store.add_summary("s1", summary("b")).await.unwrap(); + assert_eq!(store.get_summaries("s1").await.unwrap().len(), 2); + + let dump = store.dump_all().await.unwrap(); + let (fresh, _c2) = test_store().await; + fresh.restore_all(dump).await.unwrap(); + assert_eq!(fresh.get_summaries("s1").await.unwrap().len(), 2); + } +} diff --git a/src/stores/redis_shortterm.rs b/src/stores/redis_shortterm.rs index fcbeba3..7c11c74 100644 --- a/src/stores/redis_shortterm.rs +++ b/src/stores/redis_shortterm.rs @@ -109,7 +109,8 @@ impl ShortTermMemory for RedisShortTermMemory { } let mut connection = self.connection.clone(); - let start = -(count as isize); + // LRANGE start = -(count): clamp to 0 (fetch all) if count overflows isize. + let start: isize = if count > isize::MAX as usize { 0 } else { -(count as isize) }; let raw_messages: Vec = connection .lrange(Self::session_key(session_id), start, -1) .await @@ -246,6 +247,15 @@ impl ShortTermMemory for RedisShortTermMemory { } Ok(()) } + + async fn remove_messages(&self, session_id: &str, ids: &[String]) -> Result<(), MemoryError> { + if ids.is_empty() { + return Ok(()); + } + let mut messages = self.read_messages(session_id).await?; + messages.retain(|m| m.id.as_deref().map_or(true, |id| !ids.contains(&id.to_string()))); + self.write_messages(session_id, &messages).await + } } fn session_id_from_key(key: &str) -> String { @@ -296,6 +306,16 @@ mod tests { } } + fn message_with_id(id: &str, content: &str) -> Message { + Message { + id: Some(id.to_string()), + role: "user".to_string(), + content: content.to_string(), + timestamp: None, + embedding_status: None, + } + } + #[tokio::test] async fn dump_all_and_restore_all_round_trip() { let (store, _node) = test_store().await; @@ -313,4 +333,18 @@ mod tests { assert_eq!(store.get_recent("s1", 10).await.unwrap().len(), 1); assert_eq!(store.get_recent("s2", 10).await.unwrap().len(), 1); } + + #[tokio::test] + async fn remove_messages_by_id() { + let (store, _node) = test_store().await; + store.add_message("s1", message_with_id("m1", "first")).await.unwrap(); + store.add_message("s1", message_with_id("m2", "second")).await.unwrap(); + store.add_message("s1", message_with_id("m3", "third")).await.unwrap(); + + store.remove_messages("s1", &["m1".into(), "m3".into()]).await.unwrap(); + + let remaining = store.get_recent("s1", 10).await.unwrap(); + assert_eq!(remaining.len(), 1); + assert_eq!(remaining[0].id.as_deref(), Some("m2")); + } } diff --git a/tests/e2e_test.rs b/tests/e2e_test.rs index 4b8fc2f..6c098d5 100644 --- a/tests/e2e_test.rs +++ b/tests/e2e_test.rs @@ -133,6 +133,7 @@ async fn e2e_flow_uses_real_stores_and_background_worker() { knowledge_extractor: engram::config::KnowledgeExtractorType::OpenAI, raft_db_path: std::path::PathBuf::from("./data/raft/engram.redb"), snapshot_log_threshold: 1000, + ..Config::default() }; let embedding_provider: Arc = Arc::new( OpenAIEmbedder::new_with_base_url("test-key", mock_server.uri()).unwrap(), diff --git a/tests/integration_test.rs b/tests/integration_test.rs index 3137a88..cc74845 100644 --- a/tests/integration_test.rs +++ b/tests/integration_test.rs @@ -4,7 +4,7 @@ use std::time::Duration; use axum::http::StatusCode; use axum_test::TestServer; use engram::app::build_app_state_with_embedding_provider; -use engram::config::Config; +use engram::config::{Config, SummarizerType}; use engram::core::{EmbeddingProvider, ShortTermMemory}; use engram::embedding::OpenAIEmbedder; use engram::models::EmbeddingStatus; @@ -96,6 +96,7 @@ async fn setup_test_app() -> TestApp { knowledge_extractor: engram::config::KnowledgeExtractorType::OpenAI, raft_db_path: std::path::PathBuf::from("./data/raft/engram.redb"), snapshot_log_threshold: 1000, + ..Config::default() }; let embedding_provider: Arc = Arc::new( @@ -413,4 +414,101 @@ async fn deleting_a_session_removes_context_and_vector_results() { .await .unwrap(); assert!(vector_results.is_empty()); +} + +async fn setup_consolidation_test_app() -> TestApp { + let redis_container = GenericImage::new("redis", "7.2.4") + .with_exposed_port(REDIS_PORT.tcp()) + .with_wait_for(WaitFor::message_on_stdout("Ready to accept connections")) + .start() + .await + .unwrap(); + + let host = redis_container.get_host().await.unwrap(); + let port = redis_container.get_host_port_ipv4(REDIS_PORT.tcp()).await.unwrap(); + let redis_url = format!("redis://{host}:{port}/"); + + let mock_server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/v1/embeddings")) + .respond_with(ResponseTemplate::new(200).set_body_json(embedding_payload())) + .mount(&mock_server) + .await; + + let lance_db_dir = TempDir::new().unwrap(); + let config = Config { + redis_url, + openai_api_key: "test-key".to_string(), + openai_base_url: Some(mock_server.uri()), + lance_db_path: lance_db_dir.path().to_path_buf(), + embedding_dimension: 1536, + embedding_max_concurrency: 1, + mpsc_channel_size: 4, + short_term_count: 20, + knowledge_extractor: engram::config::KnowledgeExtractorType::Mock, + summarizer: SummarizerType::Mock, + consolidation_threshold: 3, + consolidation_target_window: 2, + ..Config::default() + }; + + let embedding_provider: Arc = + Arc::new(OpenAIEmbedder::new_with_base_url("test-key", mock_server.uri()).unwrap()); + let state = build_app_state_with_embedding_provider(&config, embedding_provider) + .await + .unwrap(); + let server = TestServer::new(build_router(state.clone())).unwrap(); + + TestApp { server, state, _lance_db_dir: lance_db_dir, _redis_container: redis_container, _mock_server: mock_server } +} + +#[tokio::test] +async fn consolidation_scheduler_fires_when_threshold_crossed_and_trims_short_term() { + // threshold=3, window=2: 4 messages → oldest 2 summarized, newest 2 kept. + // In standalone mode the add_message nudge is Raft-only, so we POST /consolidate explicitly. + // We wait for all embeddings to settle before consolidating: the embedding worker's + // update_message_status uses a non-atomic DEL+RPUSH that would race with remove_messages. + let app = setup_consolidation_test_app().await; + let session_id = create_session(&app.server).await; + + let mut message_ids = Vec::new(); + for i in 0..4u32 { + let id = Uuid::new_v4().to_string(); + add_user_message(&app.server, &session_id, &id, &format!("message number {i}")).await; + message_ids.push(id); + } + + // Wait for all embeddings to complete before triggering consolidation. + for id in &message_ids { + wait_for_terminal_status(app.state.short_term_memory.as_ref(), &session_id, id).await; + } + + app.server + .post(&format!("/sessions/{session_id}/consolidate")) + .await + .assert_status(StatusCode::ACCEPTED); + + // Poll GET /summaries until the async worker finishes. + let mut summaries = serde_json::Value::Array(vec![]); + for _ in 0..80 { + let resp = app.server.get(&format!("/sessions/{session_id}/summaries")).await; + resp.assert_status_ok(); + let body: serde_json::Value = resp.json(); + summaries = body["summaries"].clone(); + if summaries.as_array().map_or(false, |s| !s.is_empty()) { + break; + } + sleep(Duration::from_millis(25)).await; + } + + let summaries = summaries.as_array().expect("summaries must be an array"); + assert_eq!(summaries.len(), 1, "one summary produced"); + assert_eq!( + summaries[0]["consumed_count"].as_u64().unwrap(), + 2, + "oldest 2 consumed (4 total - 2 window)" + ); + + let remaining = app.state.short_term_memory.get_recent(&session_id, 10).await.unwrap(); + assert_eq!(remaining.len(), 2, "trimmed back to target window of 2"); } \ No newline at end of file diff --git a/tests/raft_write_test.rs b/tests/raft_write_test.rs index 231ae97..10d7b7f 100644 --- a/tests/raft_write_test.rs +++ b/tests/raft_write_test.rs @@ -24,6 +24,8 @@ async fn single_node_raft_write_commits_to_state_machine() { let knowledge_graph = Arc::new(tokio::sync::RwLock::new(engram::knowledge::graph::KnowledgeGraph::new())); let global_graph = Arc::new(tokio::sync::RwLock::new(engram::knowledge::GlobalGraph::new())); let (knowledge_tx, _knowledge_rx) = mpsc::channel(500); + let consolidated = Arc::new(engram::consolidation::store::InMemoryConsolidatedStore::default()) + as Arc; let metrics = Arc::new(engram::metrics::AppMetrics::new().unwrap()); let raft = build_raft_node( &config, @@ -34,6 +36,7 @@ async fn single_node_raft_write_commits_to_state_machine() { knowledge_graph, knowledge_tx, global_graph, + consolidated, metrics, ) .await diff --git a/tests/token_efficiency.rs b/tests/token_efficiency.rs index cee0b97..710ae68 100644 --- a/tests/token_efficiency.rs +++ b/tests/token_efficiency.rs @@ -83,6 +83,13 @@ fn build_test_state() -> Arc { global_graph: Arc::new(tokio::sync::RwLock::new( engram::knowledge::global::GlobalGraph::new(), )), + consolidated: Arc::new(engram::consolidation::store::InMemoryConsolidatedStore::default()), + consolidation_tx: { + let (tx, mut rx) = + tokio::sync::mpsc::channel::(16); + tokio::spawn(async move { while rx.recv().await.is_some() {} }); + tx + }, }) }