From 823f6fac5b1896c78773ae603513e41cd88edb5d Mon Sep 17 00:00:00 2001 From: Claude Date: Thu, 24 Sep 2026 15:40:22 +0000 Subject: [PATCH] Add chat-memory module: MessageWindowChatMemory, JDBC and Redis repositories, per-user conversation IDs Real PostgreSQL 16 and Redis Stack 7.4; 24 tests write output/01-24. Covers the 20-message default window with no property, the 36-character conversation_id, tool messages dropped by JdbcChatMemoryRepository, concurrent add() on one conversation, a 1.x table under the 2.0 repository, the Redis repository silently backing off for a custom ChatMemory, and the removal of PromptChatMemoryAdvisor. Co-Authored-By: Claude Sonnet 5 Claude-Session: https://claude.ai/code/session_01Ja4jkzrbQ4LQZBNrb5mkZE --- README.md | 1 + chat-memory/.gitignore | 1 + chat-memory/README.md | 67 ++++ chat-memory/output/01-window-semantics.txt | 16 + .../output/02-advisor-conversations.txt | 19 ++ .../output/03-missing-conversation-id.txt | 5 + .../output/04-autoconfigured-defaults.txt | 5 + chat-memory/output/05-no-window-property.txt | 7 + chat-memory/output/06-token-growth.txt | 18 ++ chat-memory/output/07-token-budget-memory.txt | 4 + chat-memory/output/08-jdbc-round-trip.txt | 16 + .../output/09-jdbc-conversation-id-limit.txt | 7 + .../output/10-jdbc-drops-tool-messages.txt | 13 + chat-memory/output/11-jdbc-lost-update.txt | 8 + .../output/12-jdbc-overlapping-saves.txt | 7 + .../output/13-jdbc-1x-table-upgrade.txt | 16 + .../output/14-jdbc-lost-update-locked.txt | 4 + chat-memory/output/15-redis-round-trip.txt | 16 + chat-memory/output/16-redis-lost-update.txt | 4 + .../output/17-redis-silent-fallback.txt | 11 + chat-memory/output/18-redis-ttl-and-cap.txt | 8 + chat-memory/output/19-redis-plain-fails.txt | 6 + chat-memory/output/20-ownership-enforced.txt | 7 + .../output/21-ownership-unenforced.txt | 7 + chat-memory/output/22-config-keys.txt | 11 + .../output/23-prompt-advisor-removed.txt | 21 ++ .../output/24-memory-with-tool-calling.txt | 15 + chat-memory/pom.xml | 82 +++++ chat-memory/scripts/capture-1x-compile.sh | 31 ++ chat-memory/scripts/run-all.sh | 13 + chat-memory/scripts/services-up.sh | 23 ++ .../src/broken/PromptAdvisorFrom1x.java | 12 + .../chatmemory/ChatMemoryApplication.java | 12 + .../chatmemory/config/MemoryConfig.java | 44 +++ .../chatmemory/config/RedisMemoryConfig.java | 55 ++++ .../memory/ConversationRegistry.java | 57 ++++ .../memory/ConversationService.java | 51 +++ .../memory/NotYourConversationException.java | 12 + .../memory/TokenBudgetChatMemory.java | 88 ++++++ .../ankurm/chatmemory/web/ChatController.java | 57 ++++ .../src/main/resources/application-jdbc.yml | 14 + .../src/main/resources/application-redis.yml | 14 + .../src/main/resources/application.yml | 25 ++ .../chatmemory/AdvisorConversationTest.java | 100 ++++++ .../com/ankurm/chatmemory/ConfigKeysTest.java | 59 ++++ .../ankurm/chatmemory/DefaultWindowTest.java | 62 ++++ .../ankurm/chatmemory/JdbcPostgresTest.java | 295 ++++++++++++++++++ .../ankurm/chatmemory/OwnershipLeakTest.java | 41 +++ .../com/ankurm/chatmemory/OwnershipTest.java | 42 +++ .../chatmemory/RedisCapsAndTtlTest.java | 68 ++++ .../chatmemory/RedisPlainFailureTest.java | 42 +++ .../chatmemory/RedisSilentFallbackTest.java | 82 +++++ .../com/ankurm/chatmemory/RedisStackTest.java | 108 +++++++ .../chatmemory/TokenBudgetChatMemoryTest.java | 48 +++ .../ankurm/chatmemory/TokenGrowthTest.java | 109 +++++++ .../chatmemory/ToolCallingMemoryTest.java | 74 +++++ .../chatmemory/WindowSemanticsTest.java | 62 ++++ .../com/ankurm/chatmemory/support/Http.java | 46 +++ .../support/InterleavingRepository.java | 87 ++++++ .../chatmemory/support/ScriptedChatModel.java | 75 +++++ .../chatmemory/support/ScriptedToolModel.java | 117 +++++++ .../ankurm/chatmemory/support/Services.java | 29 ++ .../com/ankurm/chatmemory/support/Show.java | 31 ++ .../chatmemory/support/TestModelConfig.java | 14 + .../ankurm/chatmemory/support/Transcript.java | 47 +++ 65 files changed, 2548 insertions(+) create mode 100644 chat-memory/.gitignore create mode 100644 chat-memory/README.md create mode 100644 chat-memory/output/01-window-semantics.txt create mode 100644 chat-memory/output/02-advisor-conversations.txt create mode 100644 chat-memory/output/03-missing-conversation-id.txt create mode 100644 chat-memory/output/04-autoconfigured-defaults.txt create mode 100644 chat-memory/output/05-no-window-property.txt create mode 100644 chat-memory/output/06-token-growth.txt create mode 100644 chat-memory/output/07-token-budget-memory.txt create mode 100644 chat-memory/output/08-jdbc-round-trip.txt create mode 100644 chat-memory/output/09-jdbc-conversation-id-limit.txt create mode 100644 chat-memory/output/10-jdbc-drops-tool-messages.txt create mode 100644 chat-memory/output/11-jdbc-lost-update.txt create mode 100644 chat-memory/output/12-jdbc-overlapping-saves.txt create mode 100644 chat-memory/output/13-jdbc-1x-table-upgrade.txt create mode 100644 chat-memory/output/14-jdbc-lost-update-locked.txt create mode 100644 chat-memory/output/15-redis-round-trip.txt create mode 100644 chat-memory/output/16-redis-lost-update.txt create mode 100644 chat-memory/output/17-redis-silent-fallback.txt create mode 100644 chat-memory/output/18-redis-ttl-and-cap.txt create mode 100644 chat-memory/output/19-redis-plain-fails.txt create mode 100644 chat-memory/output/20-ownership-enforced.txt create mode 100644 chat-memory/output/21-ownership-unenforced.txt create mode 100644 chat-memory/output/22-config-keys.txt create mode 100644 chat-memory/output/23-prompt-advisor-removed.txt create mode 100644 chat-memory/output/24-memory-with-tool-calling.txt create mode 100644 chat-memory/pom.xml create mode 100755 chat-memory/scripts/capture-1x-compile.sh create mode 100755 chat-memory/scripts/run-all.sh create mode 100755 chat-memory/scripts/services-up.sh create mode 100644 chat-memory/src/broken/PromptAdvisorFrom1x.java create mode 100644 chat-memory/src/main/java/com/ankurm/chatmemory/ChatMemoryApplication.java create mode 100644 chat-memory/src/main/java/com/ankurm/chatmemory/config/MemoryConfig.java create mode 100644 chat-memory/src/main/java/com/ankurm/chatmemory/config/RedisMemoryConfig.java create mode 100644 chat-memory/src/main/java/com/ankurm/chatmemory/memory/ConversationRegistry.java create mode 100644 chat-memory/src/main/java/com/ankurm/chatmemory/memory/ConversationService.java create mode 100644 chat-memory/src/main/java/com/ankurm/chatmemory/memory/NotYourConversationException.java create mode 100644 chat-memory/src/main/java/com/ankurm/chatmemory/memory/TokenBudgetChatMemory.java create mode 100644 chat-memory/src/main/java/com/ankurm/chatmemory/web/ChatController.java create mode 100644 chat-memory/src/main/resources/application-jdbc.yml create mode 100644 chat-memory/src/main/resources/application-redis.yml create mode 100644 chat-memory/src/main/resources/application.yml create mode 100644 chat-memory/src/test/java/com/ankurm/chatmemory/AdvisorConversationTest.java create mode 100644 chat-memory/src/test/java/com/ankurm/chatmemory/ConfigKeysTest.java create mode 100644 chat-memory/src/test/java/com/ankurm/chatmemory/DefaultWindowTest.java create mode 100644 chat-memory/src/test/java/com/ankurm/chatmemory/JdbcPostgresTest.java create mode 100644 chat-memory/src/test/java/com/ankurm/chatmemory/OwnershipLeakTest.java create mode 100644 chat-memory/src/test/java/com/ankurm/chatmemory/OwnershipTest.java create mode 100644 chat-memory/src/test/java/com/ankurm/chatmemory/RedisCapsAndTtlTest.java create mode 100644 chat-memory/src/test/java/com/ankurm/chatmemory/RedisPlainFailureTest.java create mode 100644 chat-memory/src/test/java/com/ankurm/chatmemory/RedisSilentFallbackTest.java create mode 100644 chat-memory/src/test/java/com/ankurm/chatmemory/RedisStackTest.java create mode 100644 chat-memory/src/test/java/com/ankurm/chatmemory/TokenBudgetChatMemoryTest.java create mode 100644 chat-memory/src/test/java/com/ankurm/chatmemory/TokenGrowthTest.java create mode 100644 chat-memory/src/test/java/com/ankurm/chatmemory/ToolCallingMemoryTest.java create mode 100644 chat-memory/src/test/java/com/ankurm/chatmemory/WindowSemanticsTest.java create mode 100644 chat-memory/src/test/java/com/ankurm/chatmemory/support/Http.java create mode 100644 chat-memory/src/test/java/com/ankurm/chatmemory/support/InterleavingRepository.java create mode 100644 chat-memory/src/test/java/com/ankurm/chatmemory/support/ScriptedChatModel.java create mode 100644 chat-memory/src/test/java/com/ankurm/chatmemory/support/ScriptedToolModel.java create mode 100644 chat-memory/src/test/java/com/ankurm/chatmemory/support/Services.java create mode 100644 chat-memory/src/test/java/com/ankurm/chatmemory/support/Show.java create mode 100644 chat-memory/src/test/java/com/ankurm/chatmemory/support/TestModelConfig.java create mode 100644 chat-memory/src/test/java/com/ankurm/chatmemory/support/Transcript.java diff --git a/README.md b/README.md index 37dec4d..5032a73 100644 --- a/README.md +++ b/README.md @@ -12,5 +12,6 @@ Runnable companion code for the Spring AI articles on [ankurm.com](https://ankur | [`tool-calling/`](tool-calling) | `@Tool` methods, `ToolCallingAdvisor` (the advisor-layer replacement for Spring AI 1.x's per-model tool loop), `returnDirect`, `ToolContext`, and `ToolSearchToolCallingAdvisor` for progressive disclosure across a 230-tool synthetic library -- every test driven by a hand-written `ScriptedChatModel`, no live model anywhere. Spring Boot 4.1.1, Spring AI 2.0.1, Java 25. | [Tool Calling in Spring AI 2.0](https://ankurm.com/spring-ai-2-0-tool-calling/) | | [`structured-output/`](structured-output) | `ChatClient.entity()` mapping LLM responses to Java records, lists and maps; `StructuredOutputValidationAdvisor` retrying non-conforming JSON with a real enum-constrained schema, including a captured run that exhausts every retry without throwing. Spring Boot 4.1.1, Spring AI 2.0.1, Java 25. | [Structured Output in Spring AI 2.0](https://ankurm.com/spring-ai-2-0-structured-output/) | | [`ollama-local/`](ollama-local) | Chat and embeddings against a real local `qwen2.5:0.5b`/`all-minilm`, no API key, driven by a Testcontainers-managed Ollama container started from a baked image; a confirmed model unload via `keep_alive: 0` and `/api/ps`, not a scripted model anywhere. Spring Boot 4.1.1, Spring AI 2.0.1, Testcontainers 2.0.5, Java 25. | [Run LLMs Locally with Spring AI and Ollama](https://ankurm.com/spring-ai-2-0-ollama-local/) | +| [`chat-memory/`](chat-memory) | `MessageChatMemoryAdvisor`, `MessageWindowChatMemory`, the JDBC and Redis `ChatMemoryRepository`, per-user conversation IDs and a token-budget memory of our own, with the traps reproduced against a real PostgreSQL 16 and Redis Stack: a 36-character `conversation_id`, tool messages dropped on save, concurrent writers, a 1.x table under the 2.0 repository, and a Redis repository that silently steps aside for a custom `ChatMemory`. Spring Boot 4.1.1, Spring AI 2.0.1, Java 25. | [Chat Memory in Spring AI 2.0: JDBC, Redis and Windowed Conversations](https://ankurm.com/spring-ai-2-0-chat-memory-jdbc-redis-windowed-conversations/) | Upgrading from Spring AI 1.x: [migration guide](https://ankurm.com/spring-ai-1-to-2-migration-guide/). diff --git a/chat-memory/.gitignore b/chat-memory/.gitignore new file mode 100644 index 0000000..2f7896d --- /dev/null +++ b/chat-memory/.gitignore @@ -0,0 +1 @@ +target/ diff --git a/chat-memory/README.md b/chat-memory/README.md new file mode 100644 index 0000000..3dd5b3a --- /dev/null +++ b/chat-memory/README.md @@ -0,0 +1,67 @@ +# chat-memory + +Companion code for [Chat Memory in Spring AI 2.0: JDBC, Redis and Windowed Conversations](https://ankurm.com/spring-ai-2-0-chat-memory-jdbc-redis-windowed-conversations/), part of the [Spring AI series](../README.md) on ankurm.com. + +`MessageChatMemoryAdvisor` + `MessageWindowChatMemory` + a `ChatMemoryRepository` (in-memory, PostgreSQL through JDBC, or Redis Stack), with per-user conversation IDs, a token-budget memory of our own, and every trap the article describes reproduced by a test that writes its own transcript. + +There is no live model anywhere: [`ScriptedChatModel`](src/test/java/com/ankurm/chatmemory/support/ScriptedChatModel.java) records the prompts it is sent, which is what a memory test needs to prove -- *which messages reached the model*. The databases are real. + +## Versions + +| Component | Version | +|---|---| +| Spring Boot | 4.1.1 | +| Spring AI | 2.0.1 (GA 2026-06-12 for 2.0.0; 2.0.1 on Maven Central 2026-08-20) | +| PostgreSQL | 16.13 | +| Redis Stack | 7.4.7 (RediSearch 2.10.20, RedisJSON) | +| Jedis (pulled in by the Redis starter) | 7.4.1 | +| Java | 25 (LTS) | + +## Quickstart + +```bash +scripts/services-up.sh # PostgreSQL 16 (apt), Redis Stack (docker), plain Redis (apt) -- optional +scripts/run-all.sh # runs all 24 tests and regenerates output/01 .. 24 +``` + +Tests that need a service which is not listening are **skipped with a message**, not failed, so the in-memory tests still run with nothing installed. Run the app itself with `mvn spring-boot:run` (in-memory), `-Dspring-boot.run.profiles=jdbc` or `=redis`; it needs `OPENAI_API_KEY` for real answers. + +## What's here + +| File | What it shows | +|---|---| +| [`config/MemoryConfig.java`](src/main/java/com/ankurm/chatmemory/config/MemoryConfig.java) | The `ChatMemory` bean (window size is not configurable by property) and the `ChatClient` with `MessageChatMemoryAdvisor` | +| [`config/RedisMemoryConfig.java`](src/main/java/com/ankurm/chatmemory/config/RedisMemoryConfig.java) | Why the Redis repository is built by hand: the autoconfigured one silently steps aside for a custom `ChatMemory` | +| [`memory/ConversationRegistry.java`](src/main/java/com/ankurm/chatmemory/memory/ConversationRegistry.java) | Who owns which conversation; `app.ownership.enforce=false` reproduces the leak | +| [`memory/ConversationService.java`](src/main/java/com/ankurm/chatmemory/memory/ConversationService.java) | Passing `ChatMemory.CONVERSATION_ID` per request | +| [`memory/TokenBudgetChatMemory.java`](src/main/java/com/ankurm/chatmemory/memory/TokenBudgetChatMemory.java) | A `ChatMemory` that prunes by tokens instead of by message count | +| [`web/ChatController.java`](src/main/java/com/ankurm/chatmemory/web/ChatController.java) | Four endpoints; the user comes from an `X-User` header only so the tests need no login | +| [`application*.yml`](src/main/resources) | Profiles `jdbc` and `redis`; the default excludes both repository autoconfigurations | +| [`src/broken/PromptAdvisorFrom1x.java`](src/broken/PromptAdvisorFrom1x.java) | Not compiled by the build; `scripts/capture-1x-compile.sh` compiles it against 2.0.1 and 1.1.8 | + +## Output files + +Every file is written by the test named in the right column (or by the script), and every console block in the article is quoted from one of them. + +| File | Written by | +|---|---| +| `01-window-semantics.txt` | `WindowSemanticsTest` | +| `02-advisor-conversations.txt`, `03-missing-conversation-id.txt` | `AdvisorConversationTest` | +| `04-autoconfigured-defaults.txt`, `05-no-window-property.txt` | `DefaultWindowTest` | +| `06-token-growth.txt` | `TokenGrowthTest` | +| `07-token-budget-memory.txt` | `TokenBudgetChatMemoryTest` | +| `08` .. `14` (JDBC round trip, 36-character limit, dropped tool messages, lost update, overlapping saves, 1.x table upgrade, locked fix) | `JdbcPostgresTest` | +| `15-redis-round-trip.txt`, `16-redis-lost-update.txt` | `RedisStackTest` | +| `17-redis-silent-fallback.txt` | `RedisSilentFallbackTest` | +| `18-redis-ttl-and-cap.txt` | `RedisCapsAndTtlTest` | +| `19-redis-plain-fails.txt` | `RedisPlainFailureTest` | +| `20-ownership-enforced.txt`, `21-ownership-unenforced.txt` | `OwnershipTest`, `OwnershipLeakTest` | +| `22-config-keys.txt` | `ConfigKeysTest` | +| `23-prompt-advisor-removed.txt` | `scripts/capture-1x-compile.sh` | +| `24-memory-with-tool-calling.txt` | `ToolCallingMemoryTest` | + +Two consecutive `mvn test` runs produce byte-identical files. `12-jdbc-overlapping-saves.txt` deliberately prints no exact count: two `saveAll` calls that overlap produce 3 or 6 messages depending on timing, and the file says so instead of pinning one run. + +## Requirements + +JDK 25, Maven. PostgreSQL 16 and Redis Stack are optional (see above); Docker is only used to run Redis Stack. diff --git a/chat-memory/output/01-window-semantics.txt b/chat-memory/output/01-window-semantics.txt new file mode 100644 index 0000000..0f6a15a --- /dev/null +++ b/chat-memory/output/01-window-semantics.txt @@ -0,0 +1,16 @@ +# MessageWindowChatMemory, maxMessages = 4 and 5 + +maxMessages = 4. One system message, then one user+assistant turn at a time: + after turn 1: S:SYS-1 | U:u1 | A:a1 + after turn 2: S:SYS-1 | U:u2 | A:a2 + after turn 3: S:SYS-1 | U:u3 | A:a3 + after turn 4: S:SYS-1 | U:u4 | A:a4 + +A second, different system message arrives: + after SYS-2: U:u4 | A:a4 | S:SYS-2 + +maxMessages = 5 (odd), no system message: + after turn 1: U:u1 | A:a1 + after turn 2: U:u1 | A:a1 | U:u2 | A:a2 + after turn 3: U:u2 | A:a2 | U:u3 | A:a3 + after turn 4: U:u3 | A:a3 | U:u4 | A:a4 diff --git a/chat-memory/output/02-advisor-conversations.txt b/chat-memory/output/02-advisor-conversations.txt new file mode 100644 index 0000000..138e594 --- /dev/null +++ b/chat-memory/output/02-advisor-conversations.txt @@ -0,0 +1,19 @@ +# MessageChatMemoryAdvisor: what the model is sent + +conv-A, call 1: "My name is Priya" + reply: Nice to meet you, Priya. + model was sent: S:You are a terse assistant. | U:My name is Priya + +conv-A, call 2: "What is my name?" + reply: Your name is Priya. + model was sent: S:You are a terse assistant. | U:My name is Priya | A:Nice to meet you, Priya. | U:What is my name? + +conv-B, call 1: "What is my name?" (different conversation ID) + reply: I do not know your name yet. + model was sent: S:You are a terse assistant. | U:What is my name? + +Stored for conv-A afterwards: + USER My name is Priya + ASSISTANT Nice to meet you, Priya. + USER What is my name? + ASSISTANT Your name is Priya. diff --git a/chat-memory/output/03-missing-conversation-id.txt b/chat-memory/output/03-missing-conversation-id.txt new file mode 100644 index 0000000..753f92d --- /dev/null +++ b/chat-memory/output/03-missing-conversation-id.txt @@ -0,0 +1,5 @@ +# No conversation ID on the request + +chat.prompt().user("hello").call().content() // no advisors(...) param + threw java.lang.IllegalArgumentException + message: conversationId cannot be null diff --git a/chat-memory/output/04-autoconfigured-defaults.txt b/chat-memory/output/04-autoconfigured-defaults.txt new file mode 100644 index 0000000..a80a7f1 --- /dev/null +++ b/chat-memory/output/04-autoconfigured-defaults.txt @@ -0,0 +1,5 @@ +# ChatMemoryAutoConfiguration, nothing configured + +ChatMemory bean: MessageWindowChatMemory +ChatMemoryRepository bean: InMemoryChatMemoryRepository +15 turns (30 messages) added -> 20 stored, oldest kept: u6 diff --git a/chat-memory/output/05-no-window-property.txt b/chat-memory/output/05-no-window-property.txt new file mode 100644 index 0000000..5f3362f --- /dev/null +++ b/chat-memory/output/05-no-window-property.txt @@ -0,0 +1,7 @@ +# Guessing a window-size property + +spring.ai.chat.memory.max-messages = 6 -> 30 messages added, 20 stored +spring.ai.chat.memory.window-size = 6 -> 30 messages added, 20 stored +spring.ai.chat.memory.repository.max-messages = 6 -> 30 messages added, 20 stored + +Spring Boot ignores unknown keys silently, so none of these did anything. diff --git a/chat-memory/output/06-token-growth.txt b/chat-memory/output/06-token-growth.txt new file mode 100644 index 0000000..3a5bffb --- /dev/null +++ b/chat-memory/output/06-token-growth.txt @@ -0,0 +1,18 @@ +# Input tokens per call, 30 turns, four pruning strategies + +Encoding: o200k_base (JTokkit estimate). Each turn: ~25-token question, ~40-token answer, ~6-token system prompt. + +turn unbounded window=10 window=20 budget=400tok +1 26 26 26 26 +2 79 79 79 79 +5 238 238 238 238 +10 503 291 503 397 +11 556 291 556 397 +15 768 291 556 397 +20 1033 291 556 397 +21 1086 291 556 397 +25 1298 291 556 397 +30 1563 291 556 397 + +sum 23835 7935 13765 10426 +sum as % of unbounded: window=10 33%, window=20 58%, budget=400tok 44% diff --git a/chat-memory/output/07-token-budget-memory.txt b/chat-memory/output/07-token-budget-memory.txt new file mode 100644 index 0000000..17d30e9 --- /dev/null +++ b/chat-memory/output/07-token-budget-memory.txt @@ -0,0 +1,4 @@ +# TokenBudgetChatMemory, budget = 120 tokens + +three short turns: S:Be terse. | U:short question 1 | A:short answer 1 | U:short question 2 | A:short answer 2 | U:short question 3 | A:short answer 3 +then one pasted trace: 7 messages kept: S:Be terse. | U:short question 2 | A:short answer 2 | U:short question 3 | A:short answer 3 | U: | A:Check the datasource URL. diff --git a/chat-memory/output/08-jdbc-round-trip.txt b/chat-memory/output/08-jdbc-round-trip.txt new file mode 100644 index 0000000..d8b29cc --- /dev/null +++ b/chat-memory/output/08-jdbc-round-trip.txt @@ -0,0 +1,16 @@ +# JdbcChatMemoryRepository on PostgreSQL 16 + +Table created by initialize-schema=always: + conversation_id character varying (36) + content text + type character varying (10) + timestamp timestamp without time zone + sequence_id bigint + +After one call, rows for the conversation: + seq=0 USER My name is Priya + seq=1 ASSISTANT Nice to meet you, Priya. + +New repository + new memory over the same database, then "What is my name?": + reply: Your name is Priya. + findConversationIds() returns every conversation in the table: [6f0a1c9e-2b7d-4c55-9a53-0d1f3e7a2b10] diff --git a/chat-memory/output/09-jdbc-conversation-id-limit.txt b/chat-memory/output/09-jdbc-conversation-id-limit.txt new file mode 100644 index 0000000..3b4536b --- /dev/null +++ b/chat-memory/output/09-jdbc-conversation-id-limit.txt @@ -0,0 +1,7 @@ +# conversation_id is VARCHAR(36) + +conversation ID: "user:asmhatre@example.com:thread:2026-09-24" (43 characters) +model calls made for this request: 0 (the user message is saved before the model is called) +thrown: org.springframework.dao.DataIntegrityViolationException +root cause: ERROR: value too long for type character varying(36) +rows stored for it afterwards: 0 diff --git a/chat-memory/output/10-jdbc-drops-tool-messages.txt b/chat-memory/output/10-jdbc-drops-tool-messages.txt new file mode 100644 index 0000000..f58a6b8 --- /dev/null +++ b/chat-memory/output/10-jdbc-drops-tool-messages.txt @@ -0,0 +1,13 @@ +# What JdbcChatMemoryRepository.saveAll keeps + +saveAll() given 4 messages: + USER Weather in Mumbai? + ASSISTANT (tool call weather) + TOOL (tool call weather) + ASSISTANT 31C and humid. + +findByConversationId() returns 2: + USER Weather in Mumbai? + ASSISTANT 31C and humid. + +Logged by the repository: WARN JdbcChatMemoryRepository does not support tool call messages. Some messages were filtered out for conversation: tool-conv diff --git a/chat-memory/output/11-jdbc-lost-update.txt b/chat-memory/output/11-jdbc-lost-update.txt new file mode 100644 index 0000000..8e13511 --- /dev/null +++ b/chat-memory/output/11-jdbc-lost-update.txt @@ -0,0 +1,8 @@ +# Two concurrent add() calls on one conversation + +seeded with: u0, a0 +two writers each add one message; both read the conversation before either writes it back +first writer to save: adds a 3-message list, saved +second writer to save: adds a different 3-message list, replaces the first +stored afterwards: [u0, a0, ] (3 messages, 4 were sent) +the first writer's message is GONE, the second writer's is kept diff --git a/chat-memory/output/12-jdbc-overlapping-saves.txt b/chat-memory/output/12-jdbc-overlapping-saves.txt new file mode 100644 index 0000000..2079cd1 --- /dev/null +++ b/chat-memory/output/12-jdbc-overlapping-saves.txt @@ -0,0 +1,7 @@ +# Two saveAll() calls released at the same instant + +seeded with u0, a0; two writers add from-A and from-B; both saveAll() calls start together +The outcome differs from run to run, so the exact count is not printed here. Across repeated runs on +PostgreSQL 16 (READ COMMITTED) it is one of two things, and never the 4 messages that were sent: + 3 messages -> one writer's message is gone + 6 messages -> both transactions inserted, history is duplicated, sequence_id values collide diff --git a/chat-memory/output/13-jdbc-1x-table-upgrade.txt b/chat-memory/output/13-jdbc-1x-table-upgrade.txt new file mode 100644 index 0000000..95cd2c3 --- /dev/null +++ b/chat-memory/output/13-jdbc-1x-table-upgrade.txt @@ -0,0 +1,16 @@ +# A Spring AI 1.x table under the 2.0 repository + +A 1.1.x-shaped table (no sequence_id column) holds 2 rows for "old-conv". + +2.0.1 repository, findByConversationId("old-conv"): + ERROR: column "sequence_id" does not exist + +Migration: add the column, backfill it from the timestamp order, then make it NOT NULL: + ALTER TABLE SPRING_AI_CHAT_MEMORY ADD COLUMN sequence_id BIGINT + UPDATE SPRING_AI_CHAT_MEMORY m SET sequence_id = s.rn FROM (SELECT ctid AS c, row_number() OVER (PARTITION ... + ALTER TABLE SPRING_AI_CHAT_MEMORY ALTER COLUMN sequence_id SET NOT NULL + CREATE INDEX IF NOT EXISTS SPRING_AI_CHAT_MEMORY_CONVERSATION_ID_SEQUENCE_ID_IDX ON SPRING_AI_CHAT_MEMORY(c... + +findByConversationId("old-conv") after the migration: + USER My name is Priya + ASSISTANT Nice to meet you, Priya. diff --git a/chat-memory/output/14-jdbc-lost-update-locked.txt b/chat-memory/output/14-jdbc-lost-update-locked.txt new file mode 100644 index 0000000..ab1367a --- /dev/null +++ b/chat-memory/output/14-jdbc-lost-update-locked.txt @@ -0,0 +1,4 @@ +# Same two writers, per-conversation lock + +stored afterwards: 4 messages +contains from-A: true, from-B: true diff --git a/chat-memory/output/15-redis-round-trip.txt b/chat-memory/output/15-redis-round-trip.txt new file mode 100644 index 0000000..786de81 --- /dev/null +++ b/chat-memory/output/15-redis-round-trip.txt @@ -0,0 +1,16 @@ +# RedisChatMemoryRepository on Redis Stack + +ChatMemoryRepository bean: RedisChatMemoryRepository +call 1 reply: Nice to meet you, Priya. + +Keys under the default prefix "chat-memory:" (3: one JSON document per message, plus a counter): + chat-memory:6f0a1c9e-2b7d-4c55-9a53-0d1f3e7a2b10: + chat-memory:6f0a1c9e-2b7d-4c55-9a53-0d1f3e7a2b10: + chat-memory:counter:6f0a1c9e-2b7d-4c55-9a53-0d1f3e7a2b10 + +Fields of one document: [content, conversation_id, metadata, timestamp, type] +TTL on each key (spring.ai.chat.memory.repository.redis.time-to-live=24h): between 86390 and 86400 s -> true + +Same conversation, second call "What is my name?": + reply: Your name is Priya. + stored: U:My name is Priya | A:Nice to meet you, Priya. | U:What is my name? | A:Your name is Priya. diff --git a/chat-memory/output/16-redis-lost-update.txt b/chat-memory/output/16-redis-lost-update.txt new file mode 100644 index 0000000..d416d48 --- /dev/null +++ b/chat-memory/output/16-redis-lost-update.txt @@ -0,0 +1,4 @@ +# Two concurrent add() calls, Redis repository + +seeded with u0, a0; two writers add one message each, both read before either writes back +stored afterwards: [u0, a0, ] (3 messages, 4 were sent) diff --git a/chat-memory/output/17-redis-silent-fallback.txt b/chat-memory/output/17-redis-silent-fallback.txt new file mode 100644 index 0000000..beaa5ba --- /dev/null +++ b/chat-memory/output/17-redis-silent-fallback.txt @@ -0,0 +1,11 @@ +# Redis starter + your own ChatMemory bean + +application has: spring-ai-starter-model-chat-memory-repository-redis, profile "redis", Redis Stack reachable, +and its own @Bean ChatMemory (needed to set a window other than 20) + +ChatMemoryRepository bean in the context: InMemoryChatMemoryRepository +Redis keys under "chat-memory:" before / after one chat call: 0 / 0 + +Condition report for the autoconfigured Redis repository bean: + @ConditionalOnMissingBean (types: org.springframework.ai.chat.memory.repository.redis.RedisChatMemoryRepository,org.springframework.ai.chat.memory.ChatMemory,org.springframework.ai.chat.memory.ChatMemoryRepository; SearchStrategy: all) found beans of type 'org.springframework.ai.chat.memory.ChatMemory' chatMemory +WARN/ERROR log lines mentioning Redis or ChatMemory, captured over the whole test class: 0 diff --git a/chat-memory/output/18-redis-ttl-and-cap.txt b/chat-memory/output/18-redis-ttl-and-cap.txt new file mode 100644 index 0000000..381c31a --- /dev/null +++ b/chat-memory/output/18-redis-ttl-and-cap.txt @@ -0,0 +1,8 @@ +# Redis time-to-live = 2s, max-messages-per-conversation = 4 + +saved 2 messages; readable now: U:u1 | A:a1 +3 seconds later: (gone: Redis expired the keys) + +saved 6 messages with max-messages-per-conversation=4, read back 4: + U:u1 | A:a1 | U:u2 | A:a2 +keys in Redis for it: 7 diff --git a/chat-memory/output/19-redis-plain-fails.txt b/chat-memory/output/19-redis-plain-fails.txt new file mode 100644 index 0000000..bcaa634 --- /dev/null +++ b/chat-memory/output/19-redis-plain-fails.txt @@ -0,0 +1,6 @@ +# RedisChatMemoryRepository against plain Redis 7.0 (no modules) + +server: redis_version:7.0.15 +RedisChatMemoryRepository.builder()...initializeSchema(true).build() threw: + java.lang.IllegalStateException + root cause: JedisDataException: ERR unknown command 'FT._LIST', with args beginning with: diff --git a/chat-memory/output/20-ownership-enforced.txt b/chat-memory/output/20-ownership-enforced.txt new file mode 100644 index 0000000..8f9936b --- /dev/null +++ b/chat-memory/output/20-ownership-enforced.txt @@ -0,0 +1,7 @@ +# app.ownership.enforce=true (default) + +alice starts a conversation and says "My name is Priya" + +bob: GET /conversations//messages -> 403 +bob: POST /conversations//messages -> 403 +alice: GET /conversations//messages -> 200, 2 messages diff --git a/chat-memory/output/21-ownership-unenforced.txt b/chat-memory/output/21-ownership-unenforced.txt new file mode 100644 index 0000000..12e1ed8 --- /dev/null +++ b/chat-memory/output/21-ownership-unenforced.txt @@ -0,0 +1,7 @@ +# app.ownership.enforce=false + +alice starts a conversation and says "My name is Priya" + +bob: GET /conversations//messages -> 200 + body: [{"role":"USER","text":"My name is Priya"},{"role":"ASSISTANT","text":"Nice to meet you, Priya."}] +bob: POST "What is my name?" -> 200 {"reply":"Your name is Priya."} diff --git a/chat-memory/output/22-config-keys.txt b/chat-memory/output/22-config-keys.txt new file mode 100644 index 0000000..1db1d2a --- /dev/null +++ b/chat-memory/output/22-config-keys.txt @@ -0,0 +1,11 @@ +# Do the chat-memory property keys bind? + +spring.ai.chat.memory.repository.jdbc.initialize-schema=always -> initializeSchema = ALWAYS + +redis.host / port / time-to-live / max-messages-per-conversation = cache.internal / 6390 / PT24H / 200 + +Three guesses, each of which looks right and binds nothing: + spring.ai.chat.memory.redis.port=6390 (no "repository" segment) + spring.ai.chat.memory.repository.redis.ttl=24h (the key is time-to-live) + spring.data.redis.port=6390 (Spring Data Redis's key; this repository builds its own Jedis client) + -> port = 6379, time-to-live = null diff --git a/chat-memory/output/23-prompt-advisor-removed.txt b/chat-memory/output/23-prompt-advisor-removed.txt new file mode 100644 index 0000000..8de7d74 --- /dev/null +++ b/chat-memory/output/23-prompt-advisor-removed.txt @@ -0,0 +1,21 @@ +# PromptChatMemoryAdvisor: Spring AI 1.1.8 vs 2.0.1 + +$ javac PromptAdvisorFrom1x.java # classpath: spring-ai-client-chat 2.0.1 +PromptAdvisorFrom1x.java:3: error: cannot find symbol +import org.springframework.ai.chat.client.advisor.PromptChatMemoryAdvisor; + ^ + symbol: class PromptChatMemoryAdvisor + location: package org.springframework.ai.chat.client.advisor +PromptAdvisorFrom1x.java:10: error: cannot find symbol + return PromptChatMemoryAdvisor.builder(chatMemory).build(); + ^ + symbol: variable PromptChatMemoryAdvisor + location: class PromptAdvisorFrom1x +2 errors + +$ javac PromptAdvisorFrom1x.java # classpath: spring-ai-client-chat 1.1.8 first +PromptAdvisorFrom1x.java:10: warning: [removal] PromptChatMemoryAdvisor in org.springframework.ai.chat.client.advisor has been deprecated and marked for removal + return PromptChatMemoryAdvisor.builder(chatMemory).build(); + ^ +1 warning +compiled: PromptAdvisorFrom1x.class diff --git a/chat-memory/output/24-memory-with-tool-calling.txt b/chat-memory/output/24-memory-with-tool-calling.txt new file mode 100644 index 0000000..99e10b5 --- /dev/null +++ b/chat-memory/output/24-memory-with-tool-calling.txt @@ -0,0 +1,15 @@ +# MessageChatMemoryAdvisor + ToolCallingAdvisor + +Advisor.DEFAULT_CHAT_MEMORY_PRECEDENCE_ORDER = HIGHEST_PRECEDENCE + 200 +ToolCallingAdvisor.DEFAULT_ORDER = HIGHEST_PRECEDENCE + 300 + +answer: It is 31C in Mumbai. +model calls: 2 + +What the model was sent on each call: + call 1: U:Weather in Mumbai? + call 2: U:Weather in Mumbai? | A:[tool call currentWeather] | T:[tool result "31C in Mumbai"] + +Stored in memory afterwards: + USER Weather in Mumbai? + ASSISTANT It is 31C in Mumbai. diff --git a/chat-memory/pom.xml b/chat-memory/pom.xml new file mode 100644 index 0000000..090f22b --- /dev/null +++ b/chat-memory/pom.xml @@ -0,0 +1,82 @@ + + + 4.0.0 + + + org.springframework.boot + spring-boot-starter-parent + 4.1.1 + + + + com.ankurm + chat-memory + 1.0.0 + chat-memory + Chat memory in Spring AI 2.0: MessageWindowChatMemory, JDBC and Redis ChatMemoryRepository, per-user conversation IDs, pruning. + + + 25 + 2.0.1 + + + + + + org.springframework.ai + spring-ai-bom + ${spring-ai.version} + pom + import + + + + + + + org.springframework.boot + spring-boot-starter-webmvc + + + org.springframework.ai + spring-ai-starter-model-openai + + + org.springframework.ai + spring-ai-starter-model-chat-memory-repository-jdbc + + + org.springframework.ai + spring-ai-starter-model-chat-memory-repository-redis + + + org.postgresql + postgresql + runtime + + + + org.springframework.boot + spring-boot-starter-test + test + + + + + + + org.springframework.boot + spring-boot-maven-plugin + + + org.apache.maven.plugins + maven-surefire-plugin + + -Duser.timezone=UTC -Dstdout.encoding=UTF-8 -Dfile.encoding=UTF-8 + + + + + diff --git a/chat-memory/scripts/capture-1x-compile.sh b/chat-memory/scripts/capture-1x-compile.sh new file mode 100755 index 0000000..d4e9b0d --- /dev/null +++ b/chat-memory/scripts/capture-1x-compile.sh @@ -0,0 +1,31 @@ +#!/usr/bin/env bash +# Compiles src/broken/PromptAdvisorFrom1x.java against the 2.0.1 classpath, then against the +# 1.1.8 spring-ai-client-chat + spring-ai-model jars, and writes both compiler results to +# output/23-prompt-advisor-removed.txt. Needs the module built once online (mvn -q test-compile). +set -euo pipefail +cd "$(dirname "$0")/.." +OUT=output/23-prompt-advisor-removed.txt +mvn -q -B dependency:build-classpath -Dmdep.outputFile=target/classpath.txt >/dev/null 2>&1 +CP2="$(cat target/classpath.txt)" +WORK="$(mktemp -d)" +for a in spring-ai-client-chat spring-ai-model; do + # Maven Central occasionally answers 429/5xx to a burst of requests: retry with a pause. + for attempt in 1 2 3 4 5; do + curl -sfL -o "$WORK/$a-1.1.8.jar" "https://repo1.maven.org/maven2/org/springframework/ai/$a/1.1.8/$a-1.1.8.jar" && break + sleep $((attempt * 3)) + done +done +CP1="$WORK/spring-ai-client-chat-1.1.8.jar:$WORK/spring-ai-model-1.1.8.jar:$CP2" +# put the 1.1.8 jars FIRST so their classes win over the 2.0.1 ones on the classpath +{ + echo "# PromptChatMemoryAdvisor: Spring AI 1.1.8 vs 2.0.1" + echo + echo "\$ javac PromptAdvisorFrom1x.java # classpath: spring-ai-client-chat 2.0.1" + javac -proc:none -d "$WORK/out2" -cp "$CP2" src/broken/PromptAdvisorFrom1x.java 2>&1 | grep -v JAVA_TOOL | sed "s#src/broken/##" || true + echo + echo "\$ javac PromptAdvisorFrom1x.java # classpath: spring-ai-client-chat 1.1.8 first" + mkdir -p "$WORK/out1" + if javac -proc:none -d "$WORK/out1" -cp "$CP1" src/broken/PromptAdvisorFrom1x.java 2>&1 | grep -v JAVA_TOOL | sed "s#src/broken/##"; then :; fi + ls "$WORK/out1" | sed 's/^/compiled: /' +} > "$OUT" +cat "$OUT" diff --git a/chat-memory/scripts/run-all.sh b/chat-memory/scripts/run-all.sh new file mode 100755 index 0000000..101b0dd --- /dev/null +++ b/chat-memory/scripts/run-all.sh @@ -0,0 +1,13 @@ +#!/usr/bin/env bash +# Regenerates every file under output/. The test suite writes 01-22 and 24 itself through the +# Transcript helper; 23 comes from scripts/capture-1x-compile.sh. +# +# Backing services: scripts/services-up.sh starts PostgreSQL 16, Redis Stack (docker) and a plain +# Redis. Tests that need one that is not listening are SKIPPED with a message, not failed, so the +# in-memory tests (01-07, 20-22, 24) still run without any of them. +set -euo pipefail +cd "$(dirname "$0")/.." +rm -rf target +mvn -q -B test +scripts/capture-1x-compile.sh >/dev/null +ls output diff --git a/chat-memory/scripts/services-up.sh b/chat-memory/scripts/services-up.sh new file mode 100755 index 0000000..31cee5a --- /dev/null +++ b/chat-memory/scripts/services-up.sh @@ -0,0 +1,23 @@ +#!/usr/bin/env bash +# Starts the three backing services the tests need, with no docker-compose: +# * PostgreSQL 16 on 127.0.0.1:5440 (apt: postgresql-16; runs as the postgres user) +# * Redis Stack on 127.0.0.1:6390 (docker: redis/redis-stack-server -- Redis + RediSearch + RedisJSON) +# * plain Redis on 127.0.0.1:6391 (apt: redis-server, NO modules; only used to capture a failure) +# If you already have these, skip this script and set CM_PG_URL / CM_REDIS_STACK_PORT / CM_REDIS_PLAIN_PORT. +set -euo pipefail + +PGDATA="${PGDATA:-/tmp/pgchat/data}"; PGPORT="${PGPORT:-5440}"; SOCKDIR="$(dirname "$PGDATA")" +BIN="$(ls -d /usr/lib/postgresql/*/bin 2>/dev/null | sort -V | tail -1)" +as_pg() { if [ "$(id -u)" = 0 ]; then su postgres -c "$*"; else bash -c "$*"; fi; } +mkdir -p "$SOCKDIR"; [ "$(id -u)" = 0 ] && chown postgres "$SOCKDIR" +[ -d "$PGDATA/base" ] || as_pg "'$BIN/initdb' -D '$PGDATA' -A trust >'$SOCKDIR/initdb.log' 2>&1" +as_pg "'$BIN/pg_ctl' -D '$PGDATA' status" >/dev/null 2>&1 || \ + as_pg "'$BIN/pg_ctl' -D '$PGDATA' -o \"-p $PGPORT -c listen_addresses=127.0.0.1 -c unix_socket_directories=$SOCKDIR\" -l '$SOCKDIR/pg.log' -w start" +PSQL="$BIN/psql -h $SOCKDIR -p $PGPORT -U postgres -v ON_ERROR_STOP=1 -q" +as_pg "$PSQL -d postgres -tc \"select 1 from pg_roles where rolname='chat'\"" | grep -q 1 || as_pg "$PSQL -d postgres -c \"create role chat login superuser password 'chat'\"" +as_pg "$PSQL -d postgres -tc \"select 1 from pg_database where datname='chatdb'\"" | grep -q 1 || as_pg "$PSQL -d postgres -c 'create database chatdb owner chat'" + +docker rm -f cm-redis-stack >/dev/null 2>&1 || true +docker run -d --name cm-redis-stack -p 6390:6379 redis/redis-stack-server:latest >/dev/null +redis-server --port 6391 --save "" --appendonly no --daemonize yes >/dev/null +echo "postgres 127.0.0.1:$PGPORT (chatdb / chat / chat), redis-stack :6390, plain redis :6391" diff --git a/chat-memory/src/broken/PromptAdvisorFrom1x.java b/chat-memory/src/broken/PromptAdvisorFrom1x.java new file mode 100644 index 0000000..a362678 --- /dev/null +++ b/chat-memory/src/broken/PromptAdvisorFrom1x.java @@ -0,0 +1,12 @@ +// Deliberately NOT part of the build (src/broken is not a source root). Compiled by +// scripts/capture-1x-compile.sh to capture the compiler's own message. +import org.springframework.ai.chat.client.advisor.PromptChatMemoryAdvisor; +import org.springframework.ai.chat.memory.ChatMemory; + +class PromptAdvisorFrom1x { + + Object build(ChatMemory chatMemory) { + // Spring AI 1.x: memory rendered into the system prompt as text + return PromptChatMemoryAdvisor.builder(chatMemory).build(); + } +} diff --git a/chat-memory/src/main/java/com/ankurm/chatmemory/ChatMemoryApplication.java b/chat-memory/src/main/java/com/ankurm/chatmemory/ChatMemoryApplication.java new file mode 100644 index 0000000..443475a --- /dev/null +++ b/chat-memory/src/main/java/com/ankurm/chatmemory/ChatMemoryApplication.java @@ -0,0 +1,12 @@ +package com.ankurm.chatmemory; + +import org.springframework.boot.SpringApplication; +import org.springframework.boot.autoconfigure.SpringBootApplication; + +@SpringBootApplication +public class ChatMemoryApplication { + + public static void main(String[] args) { + SpringApplication.run(ChatMemoryApplication.class, args); + } +} diff --git a/chat-memory/src/main/java/com/ankurm/chatmemory/config/MemoryConfig.java b/chat-memory/src/main/java/com/ankurm/chatmemory/config/MemoryConfig.java new file mode 100644 index 0000000..e243cdb --- /dev/null +++ b/chat-memory/src/main/java/com/ankurm/chatmemory/config/MemoryConfig.java @@ -0,0 +1,44 @@ +package com.ankurm.chatmemory.config; + +import org.springframework.ai.chat.client.ChatClient; +import org.springframework.ai.chat.client.advisor.MessageChatMemoryAdvisor; +import org.springframework.ai.chat.memory.ChatMemory; +import org.springframework.ai.chat.memory.ChatMemoryRepository; +import org.springframework.ai.chat.memory.MessageWindowChatMemory; +import org.springframework.ai.chat.model.ChatModel; +import org.springframework.beans.factory.annotation.Value; +import org.springframework.context.annotation.Bean; +import org.springframework.context.annotation.Configuration; + +/** + * Two beans, and both are choices the autoconfiguration would otherwise make for you: + * + *
    + *
  • {@link ChatMemory}: Spring AI 2.0.1 autoconfigures a {@code MessageWindowChatMemory} + * with a hard-coded window of 20 messages and no property to change it, so a window of + * any other size means declaring this bean yourself. It sits on whichever + * {@link ChatMemoryRepository} bean exists (in-memory, JDBC or Redis).
  • + *
  • {@link ChatClient}: one advisor, {@link MessageChatMemoryAdvisor}, that loads the + * conversation before every call and appends the new turn after it.
  • + *
+ */ +@Configuration +public class MemoryConfig { + + @Bean + ChatMemory chatMemory(ChatMemoryRepository repository, + @Value("${app.memory.max-messages:20}") int maxMessages) { + return MessageWindowChatMemory.builder() + .chatMemoryRepository(repository) + .maxMessages(maxMessages) + .build(); + } + + @Bean + ChatClient chatClient(ChatModel chatModel, ChatMemory chatMemory) { + return ChatClient.builder(chatModel) + .defaultSystem("You are a terse assistant.") + .defaultAdvisors(MessageChatMemoryAdvisor.builder(chatMemory).build()) + .build(); + } +} diff --git a/chat-memory/src/main/java/com/ankurm/chatmemory/config/RedisMemoryConfig.java b/chat-memory/src/main/java/com/ankurm/chatmemory/config/RedisMemoryConfig.java new file mode 100644 index 0000000..b46d4fc --- /dev/null +++ b/chat-memory/src/main/java/com/ankurm/chatmemory/config/RedisMemoryConfig.java @@ -0,0 +1,55 @@ +package com.ankurm.chatmemory.config; + +import java.time.Duration; + +import org.springframework.ai.chat.memory.repository.redis.RedisChatMemoryRepository; +import org.springframework.ai.model.chat.memory.repository.redis.autoconfigure.RedisChatMemoryRepositoryProperties; +import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty; +import org.springframework.context.annotation.Bean; +import org.springframework.context.annotation.Configuration; +import org.springframework.context.annotation.Profile; +import redis.clients.jedis.RedisClient; + +/** + * Builds the {@link RedisChatMemoryRepository} by hand from the autoconfigured properties and + * client. Why this class exists at all is the second-most surprising thing in the article: in + * Spring AI 2.0.1 the autoconfigured Redis repository bean is + * {@code @ConditionalOnMissingBean(RedisChatMemoryRepository, ChatMemory, ChatMemoryRepository)}. + * The moment the application declares its own {@code ChatMemory} bean (which it must, to change + * the window from 20), the autoconfigured Redis repository silently steps aside and the + * {@code InMemoryChatMemoryRepository} default takes over. Nothing is logged and everything + * appears to work until the first restart. + * + *

Set {@code app.memory.redis.explicit-repository=false} to remove this class from the + * context and reproduce that. + */ +@Configuration +@Profile("redis") +@ConditionalOnProperty(name = "app.memory.redis.explicit-repository", havingValue = "true", matchIfMissing = true) +public class RedisMemoryConfig { + + @Bean + RedisChatMemoryRepository redisChatMemoryRepository(RedisClient jedisClient, RedisChatMemoryRepositoryProperties props) { + RedisChatMemoryRepository.Builder builder = RedisChatMemoryRepository.builder() + .jedisClient(jedisClient) + .indexName(props.getIndexName()) + .keyPrefix(props.getKeyPrefix()); + Duration ttl = props.getTimeToLive(); + if (ttl != null) { + builder.timeToLive(ttl); + } + if (props.getInitializeSchema() != null) { + builder.initializeSchema(props.getInitializeSchema()); + } + if (props.getMaxMessagesPerConversation() != null) { + builder.maxMessagesPerConversation(props.getMaxMessagesPerConversation()); + } + if (props.getMaxConversationIds() != null) { + builder.maxConversationIds(props.getMaxConversationIds()); + } + if (props.getMetadataFields() != null) { + builder.metadataFields(props.getMetadataFields()); + } + return builder.build(); + } +} diff --git a/chat-memory/src/main/java/com/ankurm/chatmemory/memory/ConversationRegistry.java b/chat-memory/src/main/java/com/ankurm/chatmemory/memory/ConversationRegistry.java new file mode 100644 index 0000000..32b2771 --- /dev/null +++ b/chat-memory/src/main/java/com/ankurm/chatmemory/memory/ConversationRegistry.java @@ -0,0 +1,57 @@ +package com.ankurm.chatmemory.memory; + +import java.util.List; +import java.util.Map; +import java.util.UUID; +import java.util.concurrent.ConcurrentHashMap; + +import org.springframework.beans.factory.annotation.Value; +import org.springframework.stereotype.Component; + +/** + * Who owns which conversation. {@code ChatMemory} has no idea: it stores whatever it is handed + * under whatever ID it is handed, so "user A must not read user B's history" is entirely the + * application's job. This map is the smallest honest version of that job. In production it is a + * table ({@code conversation_id, user_id}) next to the memory table. + * + *

Set {@code app.ownership.enforce=false} to switch the check off and reproduce the leak the + * article shows. + */ +@Component +public class ConversationRegistry { + + private final Map ownerByConversation = new ConcurrentHashMap<>(); + private final boolean enforce; + + public ConversationRegistry(@Value("${app.ownership.enforce:true}") boolean enforce) { + this.enforce = enforce; + } + + /** A UUID is exactly 36 characters, the width of the JDBC repository's conversation_id column. */ + public String create(String user) { + String id = UUID.randomUUID().toString(); + ownerByConversation.put(id, user); + return id; + } + + public void requireOwner(String user, String conversationId) { + if (!enforce) { + return; + } + if (!user.equals(ownerByConversation.get(conversationId))) { + throw new NotYourConversationException(conversationId); + } + } + + public List conversationsOf(String user) { + return ownerByConversation.entrySet().stream() + .filter(e -> e.getValue().equals(user)) + .map(Map.Entry::getKey) + .sorted() + .toList(); + } + + public void forget(String conversationId) { + ownerByConversation.remove(conversationId); + } +} diff --git a/chat-memory/src/main/java/com/ankurm/chatmemory/memory/ConversationService.java b/chat-memory/src/main/java/com/ankurm/chatmemory/memory/ConversationService.java new file mode 100644 index 0000000..51c81ab --- /dev/null +++ b/chat-memory/src/main/java/com/ankurm/chatmemory/memory/ConversationService.java @@ -0,0 +1,51 @@ +package com.ankurm.chatmemory.memory; + +import java.util.List; + +import org.springframework.ai.chat.client.ChatClient; +import org.springframework.ai.chat.memory.ChatMemory; +import org.springframework.ai.chat.messages.Message; +import org.springframework.stereotype.Service; + +/** + * The only class that talks to both the {@link ChatClient} and the {@link ChatMemory}. The + * conversation ID reaches the advisor as a request parameter ({@link ChatMemory#CONVERSATION_ID}); + * the advisor reads it back out of the request context. + */ +@Service +public class ConversationService { + + private final ChatClient chatClient; + private final ChatMemory chatMemory; + private final ConversationRegistry registry; + + public ConversationService(ChatClient chatClient, ChatMemory chatMemory, ConversationRegistry registry) { + this.chatClient = chatClient; + this.chatMemory = chatMemory; + this.registry = registry; + } + + public String start(String user) { + return registry.create(user); + } + + public String send(String user, String conversationId, String text) { + registry.requireOwner(user, conversationId); + return chatClient.prompt() + .user(text) + .advisors(a -> a.param(ChatMemory.CONVERSATION_ID, conversationId)) + .call() + .content(); + } + + public List history(String user, String conversationId) { + registry.requireOwner(user, conversationId); + return chatMemory.get(conversationId); + } + + public void delete(String user, String conversationId) { + registry.requireOwner(user, conversationId); + chatMemory.clear(conversationId); + registry.forget(conversationId); + } +} diff --git a/chat-memory/src/main/java/com/ankurm/chatmemory/memory/NotYourConversationException.java b/chat-memory/src/main/java/com/ankurm/chatmemory/memory/NotYourConversationException.java new file mode 100644 index 0000000..ff737bc --- /dev/null +++ b/chat-memory/src/main/java/com/ankurm/chatmemory/memory/NotYourConversationException.java @@ -0,0 +1,12 @@ +package com.ankurm.chatmemory.memory; + +import org.springframework.http.HttpStatus; +import org.springframework.web.bind.annotation.ResponseStatus; + +@ResponseStatus(HttpStatus.FORBIDDEN) +public class NotYourConversationException extends RuntimeException { + + public NotYourConversationException(String conversationId) { + super("conversation " + conversationId + " does not belong to this user"); + } +} diff --git a/chat-memory/src/main/java/com/ankurm/chatmemory/memory/TokenBudgetChatMemory.java b/chat-memory/src/main/java/com/ankurm/chatmemory/memory/TokenBudgetChatMemory.java new file mode 100644 index 0000000..2012bed --- /dev/null +++ b/chat-memory/src/main/java/com/ankurm/chatmemory/memory/TokenBudgetChatMemory.java @@ -0,0 +1,88 @@ +package com.ankurm.chatmemory.memory; + +import java.util.ArrayList; +import java.util.List; + +import org.springframework.ai.chat.memory.ChatMemory; +import org.springframework.ai.chat.memory.ChatMemoryRepository; +import org.springframework.ai.chat.messages.Message; +import org.springframework.ai.chat.messages.MessageType; +import org.springframework.ai.tokenizer.TokenCountEstimator; + +/** + * A {@link ChatMemory} that prunes by tokens instead of by message count. {@code + * MessageWindowChatMemory} keeps "the last N messages", which is a poor proxy for cost: ten + * short messages and ten pasted stack traces are both "10". This one drops the oldest + * non-system messages until the whole conversation fits a token budget, and never leaves a + * reply at the front with its question gone. + * + *

It is a {@code ChatMemory}, so it drops into {@code MessageChatMemoryAdvisor.builder(...)} + * exactly where the window did. The token count is an estimate (the estimator you pass in), not + * the provider's own billing count. + */ +public final class TokenBudgetChatMemory implements ChatMemory { + + private final ChatMemoryRepository repository; + private final TokenCountEstimator estimator; + private final int maxTokens; + + public TokenBudgetChatMemory(ChatMemoryRepository repository, TokenCountEstimator estimator, int maxTokens) { + this.repository = repository; + this.estimator = estimator; + this.maxTokens = maxTokens; + } + + @Override + public void add(String conversationId, List messages) { + List all = new ArrayList<>(repository.findByConversationId(conversationId)); + all.addAll(messages); + repository.saveAll(conversationId, prune(all)); + } + + @Override + public List get(String conversationId) { + return repository.findByConversationId(conversationId); + } + + @Override + public void clear(String conversationId) { + repository.deleteByConversationId(conversationId); + } + + List prune(List all) { + List kept = new ArrayList<>(all); + while (tokens(kept) > maxTokens) { + int oldest = firstIndexOfNonSystem(kept); + if (oldest < 0) { + break; // only system messages left: nothing more can be dropped + } + kept.remove(oldest); + // A reply whose question was just dropped is noise; drop leading non-user messages too. + while (!kept.isEmpty()) { + int first = firstIndexOfNonSystem(kept); + if (first < 0 || kept.get(first).getMessageType() == MessageType.USER) { + break; + } + kept.remove(first); + } + } + return kept; + } + + private int tokens(List messages) { + int total = 0; + for (Message m : messages) { + total += estimator.estimate(m.getText()); + } + return total; + } + + private static int firstIndexOfNonSystem(List messages) { + for (int i = 0; i < messages.size(); i++) { + if (messages.get(i).getMessageType() != MessageType.SYSTEM) { + return i; + } + } + return -1; + } +} diff --git a/chat-memory/src/main/java/com/ankurm/chatmemory/web/ChatController.java b/chat-memory/src/main/java/com/ankurm/chatmemory/web/ChatController.java new file mode 100644 index 0000000..3ef42f1 --- /dev/null +++ b/chat-memory/src/main/java/com/ankurm/chatmemory/web/ChatController.java @@ -0,0 +1,57 @@ +package com.ankurm.chatmemory.web; + +import java.util.List; +import java.util.Map; + +import com.ankurm.chatmemory.memory.ConversationService; +import org.springframework.web.bind.annotation.DeleteMapping; +import org.springframework.web.bind.annotation.GetMapping; +import org.springframework.web.bind.annotation.PathVariable; +import org.springframework.web.bind.annotation.PostMapping; +import org.springframework.web.bind.annotation.RequestBody; +import org.springframework.web.bind.annotation.RequestHeader; +import org.springframework.web.bind.annotation.RestController; + +/** + * The user comes from an {@code X-User} header purely so the tests need no login. In a real + * service it is {@code Principal#getName()} from Spring Security -- never a header the client + * chooses. + */ +@RestController +public class ChatController { + + public record SendRequest(String text) { + } + + public record MessageView(String role, String text) { + } + + private final ConversationService conversations; + + public ChatController(ConversationService conversations) { + this.conversations = conversations; + } + + @PostMapping("/conversations") + Map start(@RequestHeader("X-User") String user) { + return Map.of("conversationId", conversations.start(user)); + } + + @PostMapping("/conversations/{id}/messages") + Map send(@RequestHeader("X-User") String user, @PathVariable String id, + @RequestBody SendRequest request) { + return Map.of("reply", conversations.send(user, id, request.text())); + } + + @GetMapping("/conversations/{id}/messages") + List history(@RequestHeader("X-User") String user, @PathVariable String id) { + return conversations.history(user, id).stream() + .map(m -> new MessageView(m.getMessageType().name(), m.getText())) + .toList(); + } + + @DeleteMapping("/conversations/{id}") + void delete(@RequestHeader("X-User") String user, @PathVariable String id) { + conversations.delete(user, id); + } +} diff --git a/chat-memory/src/main/resources/application-jdbc.yml b/chat-memory/src/main/resources/application-jdbc.yml new file mode 100644 index 0000000..d0ec366 --- /dev/null +++ b/chat-memory/src/main/resources/application-jdbc.yml @@ -0,0 +1,14 @@ +spring: + datasource: + url: ${CM_PG_URL:jdbc:postgresql://127.0.0.1:5440/chatdb} + username: chat + password: chat + ai: + chat: + memory: + repository: + jdbc: + initialize-schema: always + autoconfigure: + exclude: + - org.springframework.ai.model.chat.memory.repository.redis.autoconfigure.RedisChatMemoryRepositoryAutoConfiguration diff --git a/chat-memory/src/main/resources/application-redis.yml b/chat-memory/src/main/resources/application-redis.yml new file mode 100644 index 0000000..80ac865 --- /dev/null +++ b/chat-memory/src/main/resources/application-redis.yml @@ -0,0 +1,14 @@ +spring: + ai: + chat: + memory: + repository: + redis: + host: 127.0.0.1 + port: ${CM_REDIS_STACK_PORT:6390} + time-to-live: 24h + max-messages-per-conversation: 200 + autoconfigure: + exclude: + - org.springframework.boot.jdbc.autoconfigure.DataSourceAutoConfiguration + - org.springframework.ai.model.chat.memory.repository.jdbc.autoconfigure.JdbcChatMemoryRepositoryAutoConfiguration diff --git a/chat-memory/src/main/resources/application.yml b/chat-memory/src/main/resources/application.yml new file mode 100644 index 0000000..83b1183 --- /dev/null +++ b/chat-memory/src/main/resources/application.yml @@ -0,0 +1,25 @@ +spring: + application: + name: chat-memory + ai: + model: + chat: openai + openai: + api-key: ${OPENAI_API_KEY:sk-not-set} + chat: + model: gpt-5-mini + # Default profile = the in-memory repository. Spring AI's JDBC and Redis repository + # autoconfigurations are on the classpath (both starters are dependencies) and would otherwise + # demand a DataSource / a Redis server at startup, so they are excluded here and re-enabled per + # profile in application-jdbc.yml and application-redis.yml. + autoconfigure: + exclude: + - org.springframework.boot.jdbc.autoconfigure.DataSourceAutoConfiguration + - org.springframework.ai.model.chat.memory.repository.jdbc.autoconfigure.JdbcChatMemoryRepositoryAutoConfiguration + - org.springframework.ai.model.chat.memory.repository.redis.autoconfigure.RedisChatMemoryRepositoryAutoConfiguration + +app: + memory: + max-messages: 20 + ownership: + enforce: true diff --git a/chat-memory/src/test/java/com/ankurm/chatmemory/AdvisorConversationTest.java b/chat-memory/src/test/java/com/ankurm/chatmemory/AdvisorConversationTest.java new file mode 100644 index 0000000..43aa337 --- /dev/null +++ b/chat-memory/src/test/java/com/ankurm/chatmemory/AdvisorConversationTest.java @@ -0,0 +1,100 @@ +package com.ankurm.chatmemory; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +import java.util.List; + +import com.ankurm.chatmemory.support.ScriptedChatModel; +import com.ankurm.chatmemory.support.Show; +import com.ankurm.chatmemory.support.TestModelConfig; +import com.ankurm.chatmemory.support.Transcript; +import org.junit.jupiter.api.Test; +import org.springframework.ai.chat.client.ChatClient; +import org.springframework.ai.chat.memory.ChatMemory; +import org.springframework.ai.chat.prompt.Prompt; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.context.annotation.Import; + +/** + * The advisor end to end on the in-memory repository: which messages reach the model on each + * call, what ends up in storage, and what happens when the conversation ID is left out. + */ +@SpringBootTest(properties = "spring.ai.model.chat=none") +@Import(TestModelConfig.class) +class AdvisorConversationTest { + + @Autowired + ChatClient chat; + @Autowired + ChatMemory memory; + @Autowired + ScriptedChatModel model; + + private String ask(String conversationId, String text) { + return chat.prompt().user(text) + .advisors(a -> a.param(ChatMemory.CONVERSATION_ID, conversationId)) + .call().content(); + } + + private static String sent(Prompt prompt) { + return Show.inline(prompt.getInstructions()); + } + + @Test + void memoryIsPerConversationAndCarriesTheHistoryIntoEveryPrompt() { + try (Transcript t = new Transcript("02-advisor-conversations.txt", "MessageChatMemoryAdvisor: what the model is sent")) { + String a = "conv-A"; + String b = "conv-B"; + + t.line("conv-A, call 1: \"My name is Priya\""); + t.line(" reply: %s", ask(a, "My name is Priya")); + t.line(" model was sent: %s", sent(model.lastPrompt())); + + t.blank().line("conv-A, call 2: \"What is my name?\""); + String replyA = ask(a, "What is my name?"); + t.line(" reply: %s", replyA); + t.line(" model was sent: %s", sent(model.lastPrompt())); + + t.blank().line("conv-B, call 1: \"What is my name?\" (different conversation ID)"); + String replyB = ask(b, "What is my name?"); + t.line(" reply: %s", replyB); + t.line(" model was sent: %s", sent(model.lastPrompt())); + + t.blank().line("Stored for conv-A afterwards:"); + memory.get(a).forEach(m -> t.line(" %s", Show.one(m))); + + assertThat(replyA).isEqualTo("Your name is Priya."); + assertThat(replyB).isEqualTo("I do not know your name yet."); + } + } + + @Test + void leavingOutTheConversationIdIsAnErrorNotASharedDefault() { + try (Transcript t = new Transcript("03-missing-conversation-id.txt", "No conversation ID on the request")) { + t.line("chat.prompt().user(\"hello\").call().content() // no advisors(...) param"); + Throwable thrown = null; + try { + chat.prompt().user("hello").call().content(); + } + catch (RuntimeException e) { + thrown = e; + } + assertThat(thrown).isNotNull(); + t.line(" threw %s", thrown.getClass().getName()); + t.line(" message: %s", thrown.getMessage()); + assertThatThrownBy(() -> chat.prompt().user("hello").call().content()) + .isInstanceOf(RuntimeException.class); + } + } + + @Test + void oneCallOnAFreshConversationSendsOnlyTheSystemPromptAndTheUserTurn() { + List before = List.copyOf(model.prompts()); + ask("fresh", "hi"); + Prompt p = model.lastPrompt(); + assertThat(model.prompts()).hasSize(before.size() + 1); + assertThat(p.getInstructions()).hasSize(2); + } +} diff --git a/chat-memory/src/test/java/com/ankurm/chatmemory/ConfigKeysTest.java b/chat-memory/src/test/java/com/ankurm/chatmemory/ConfigKeysTest.java new file mode 100644 index 0000000..c01a526 --- /dev/null +++ b/chat-memory/src/test/java/com/ankurm/chatmemory/ConfigKeysTest.java @@ -0,0 +1,59 @@ +package com.ankurm.chatmemory; + +import static org.assertj.core.api.Assertions.assertThat; + +import java.util.Map; + +import com.ankurm.chatmemory.support.Transcript; +import org.junit.jupiter.api.Test; +import org.springframework.ai.model.chat.memory.repository.jdbc.autoconfigure.JdbcChatMemoryRepositoryProperties; +import org.springframework.ai.model.chat.memory.repository.redis.autoconfigure.RedisChatMemoryRepositoryProperties; +import org.springframework.boot.context.properties.bind.Bindable; +import org.springframework.boot.context.properties.bind.Binder; +import org.springframework.boot.context.properties.source.MapConfigurationPropertySource; + +/** + * Spring Boot silently ignores property keys nothing binds to. This test binds the keys the + * application's YAML uses to the real properties classes, then binds two plausible-looking wrong + * spellings to show that they change nothing. + */ +class ConfigKeysTest { + + private static T bind(String prefix, T target, Map props) { + new Binder(new MapConfigurationPropertySource(props)).bind(prefix, Bindable.ofInstance(target)); + return target; + } + + @Test + void theKeysInTheYamlBindAndTheGuessesDoNot() { + try (Transcript t = new Transcript("22-config-keys.txt", "Do the chat-memory property keys bind?")) { + JdbcChatMemoryRepositoryProperties jdbc = bind("spring.ai.chat.memory.repository.jdbc", + new JdbcChatMemoryRepositoryProperties(), Map.of("spring.ai.chat.memory.repository.jdbc.initialize-schema", "always")); + t.line("spring.ai.chat.memory.repository.jdbc.initialize-schema=always -> initializeSchema = %s", jdbc.getInitializeSchema()); + assertThat(jdbc.getInitializeSchema().name()).isEqualTo("ALWAYS"); + + RedisChatMemoryRepositoryProperties redis = bind("spring.ai.chat.memory.repository.redis", + new RedisChatMemoryRepositoryProperties(), Map.of( + "spring.ai.chat.memory.repository.redis.host", "cache.internal", + "spring.ai.chat.memory.repository.redis.port", "6390", + "spring.ai.chat.memory.repository.redis.time-to-live", "24h", + "spring.ai.chat.memory.repository.redis.max-messages-per-conversation", "200")); + t.blank().line("redis.host / port / time-to-live / max-messages-per-conversation = %s / %d / %s / %d", + redis.getHost(), redis.getPort(), redis.getTimeToLive(), redis.getMaxMessagesPerConversation()); + assertThat(redis.getPort()).isEqualTo(6390); + + RedisChatMemoryRepositoryProperties guess = bind("spring.ai.chat.memory.repository.redis", + new RedisChatMemoryRepositoryProperties(), Map.of( + "spring.ai.chat.memory.redis.port", "6390", + "spring.ai.chat.memory.repository.redis.ttl", "24h", + "spring.data.redis.port", "6390")); + t.blank().line("Three guesses, each of which looks right and binds nothing:"); + t.line(" spring.ai.chat.memory.redis.port=6390 (no \"repository\" segment)"); + t.line(" spring.ai.chat.memory.repository.redis.ttl=24h (the key is time-to-live)"); + t.line(" spring.data.redis.port=6390 (Spring Data Redis's key; this repository builds its own Jedis client)"); + t.line(" -> port = %d, time-to-live = %s", guess.getPort(), guess.getTimeToLive()); + assertThat(guess.getPort()).isEqualTo(6379); + assertThat(guess.getTimeToLive()).isNull(); + } + } +} diff --git a/chat-memory/src/test/java/com/ankurm/chatmemory/DefaultWindowTest.java b/chat-memory/src/test/java/com/ankurm/chatmemory/DefaultWindowTest.java new file mode 100644 index 0000000..ec16315 --- /dev/null +++ b/chat-memory/src/test/java/com/ankurm/chatmemory/DefaultWindowTest.java @@ -0,0 +1,62 @@ +package com.ankurm.chatmemory; + +import static org.assertj.core.api.Assertions.assertThat; + +import java.util.List; + +import com.ankurm.chatmemory.support.Transcript; +import org.junit.jupiter.api.Test; +import org.springframework.ai.chat.memory.ChatMemory; +import org.springframework.ai.chat.memory.ChatMemoryRepository; +import org.springframework.ai.chat.messages.AssistantMessage; +import org.springframework.ai.chat.messages.Message; +import org.springframework.ai.chat.messages.UserMessage; +import org.springframework.ai.model.chat.memory.autoconfigure.ChatMemoryAutoConfiguration; +import org.springframework.boot.autoconfigure.AutoConfigurations; +import org.springframework.boot.test.context.runner.ApplicationContextRunner; + +/** What you get with no configuration at all -- and what happens to the properties you might expect to exist. */ +class DefaultWindowTest { + + private final ApplicationContextRunner runner = new ApplicationContextRunner() + .withConfiguration(AutoConfigurations.of(ChatMemoryAutoConfiguration.class)); + + private static int storedAfter30Messages(ChatMemory memory) { + for (int n = 1; n <= 15; n++) { + memory.add("c", List.of(new UserMessage("u" + n), new AssistantMessage("a" + n))); + } + List kept = memory.get("c"); + return kept.size(); + } + + @Test + void autoconfigurationGivesAnInMemoryRepositoryAndAWindowOfTwenty() { + try (Transcript t = new Transcript("04-autoconfigured-defaults.txt", "ChatMemoryAutoConfiguration, nothing configured")) { + runner.run(ctx -> { + ChatMemory memory = ctx.getBean(ChatMemory.class); + ChatMemoryRepository repo = ctx.getBean(ChatMemoryRepository.class); + t.line("ChatMemory bean: %s", memory.getClass().getSimpleName()); + t.line("ChatMemoryRepository bean: %s", repo.getClass().getSimpleName()); + int kept = storedAfter30Messages(memory); + t.line("15 turns (30 messages) added -> %d stored, oldest kept: %s", kept, memory.get("c").getFirst().getText()); + assertThat(memory).isInstanceOf(org.springframework.ai.chat.memory.MessageWindowChatMemory.class); + assertThat(kept).isEqualTo(20); + }); + } + } + + @Test + void thereIsNoSpringAiPropertyForTheWindowSize() { + try (Transcript t = new Transcript("05-no-window-property.txt", "Guessing a window-size property")) { + for (String key : List.of("spring.ai.chat.memory.max-messages", "spring.ai.chat.memory.window-size", + "spring.ai.chat.memory.repository.max-messages")) { + runner.withPropertyValues(key + "=6").run(ctx -> { + int kept = storedAfter30Messages(ctx.getBean(ChatMemory.class)); + t.line("%-52s = 6 -> 30 messages added, %d stored", key, kept); + assertThat(kept).isEqualTo(20); + }); + } + t.blank().line("Spring Boot ignores unknown keys silently, so none of these did anything."); + } + } +} diff --git a/chat-memory/src/test/java/com/ankurm/chatmemory/JdbcPostgresTest.java b/chat-memory/src/test/java/com/ankurm/chatmemory/JdbcPostgresTest.java new file mode 100644 index 0000000..c1d406e --- /dev/null +++ b/chat-memory/src/test/java/com/ankurm/chatmemory/JdbcPostgresTest.java @@ -0,0 +1,295 @@ +package com.ankurm.chatmemory; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +import java.util.List; +import java.util.Map; +import java.util.concurrent.BrokenBarrierException; +import java.util.concurrent.CyclicBarrier; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.TimeoutException; +import java.util.stream.Collectors; + +import com.ankurm.chatmemory.support.InterleavingRepository; +import com.ankurm.chatmemory.support.ScriptedChatModel; +import com.ankurm.chatmemory.support.Services; +import com.ankurm.chatmemory.support.Show; +import com.ankurm.chatmemory.support.TestModelConfig; +import com.ankurm.chatmemory.support.Transcript; +import org.junit.jupiter.api.BeforeAll; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.springframework.ai.chat.client.ChatClient; +import org.springframework.ai.chat.client.advisor.MessageChatMemoryAdvisor; +import org.springframework.ai.chat.memory.ChatMemory; +import org.springframework.ai.chat.memory.ChatMemoryRepository; +import org.springframework.ai.chat.memory.MessageWindowChatMemory; +import org.springframework.ai.chat.memory.repository.jdbc.JdbcChatMemoryRepository; +import org.springframework.ai.chat.messages.AssistantMessage; +import org.springframework.ai.chat.messages.Message; +import org.springframework.ai.chat.messages.ToolResponseMessage; +import org.springframework.ai.chat.messages.UserMessage; +import org.springframework.ai.chat.prompt.Prompt; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.boot.test.system.CapturedOutput; +import org.springframework.boot.test.system.OutputCaptureExtension; +import org.springframework.context.annotation.Import; +import org.springframework.jdbc.core.JdbcTemplate; + +/** The autoconfigured JdbcChatMemoryRepository against a real PostgreSQL 16 (scripts/services-up.sh). */ +@SpringBootTest(properties = { "spring.ai.model.chat=none", "spring.profiles.active=jdbc" }) +@Import(TestModelConfig.class) +@ExtendWith(OutputCaptureExtension.class) +class JdbcPostgresTest { + + @BeforeAll + static void needPostgres() { + Services.require("PostgreSQL", Services.PG_PORT); + } + + @Autowired + ChatClient chat; + @Autowired + ChatMemory memory; + @Autowired + ChatMemoryRepository repository; + @Autowired + JdbcTemplate jdbc; + @Autowired + ScriptedChatModel model; + + @BeforeEach + void clean() { + jdbc.update("DELETE FROM SPRING_AI_CHAT_MEMORY"); + } + + private String ask(String id, String text) { + return chat.prompt().user(text).advisors(a -> a.param(ChatMemory.CONVERSATION_ID, id)).call().content(); + } + + @Test + void historySurvivesARestartBecauseItLivesInThePostgresTable() { + try (Transcript t = new Transcript("08-jdbc-round-trip.txt", "JdbcChatMemoryRepository on PostgreSQL 16")) { + List> columns = jdbc.queryForList(""" + SELECT column_name, data_type, character_maximum_length + FROM information_schema.columns WHERE table_name = 'spring_ai_chat_memory' ORDER BY ordinal_position"""); + t.line("Table created by initialize-schema=always:"); + columns.forEach(c -> t.line(" %-16s %-28s %s", c.get("column_name"), c.get("data_type"), + c.get("character_maximum_length") == null ? "" : "(" + c.get("character_maximum_length") + ")")); + + String id = "6f0a1c9e-2b7d-4c55-9a53-0d1f3e7a2b10"; + ask(id, "My name is Priya"); + t.blank().line("After one call, rows for the conversation:"); + jdbc.queryForList("SELECT sequence_id, type, content FROM SPRING_AI_CHAT_MEMORY WHERE conversation_id = ? ORDER BY sequence_id", id) + .forEach(r -> t.line(" seq=%s %-9s %s", r.get("sequence_id"), r.get("type"), r.get("content"))); + + // "Restart": a brand-new repository and memory over the same database. Nothing is shared in the JVM. + JdbcChatMemoryRepository fresh = JdbcChatMemoryRepository.builder().jdbcTemplate(jdbc) + .dialect(new org.springframework.ai.chat.memory.repository.jdbc.PostgresChatMemoryRepositoryDialect()).build(); + ChatMemory freshMemory = MessageWindowChatMemory.builder().chatMemoryRepository(fresh).build(); + ChatClient freshClient = ChatClient.builder(model) + .defaultAdvisors(MessageChatMemoryAdvisor.builder(freshMemory).build()).build(); + String reply = freshClient.prompt().user("What is my name?") + .advisors(a -> a.param(ChatMemory.CONVERSATION_ID, id)).call().content(); + t.blank().line("New repository + new memory over the same database, then \"What is my name?\":"); + t.line(" reply: %s", reply); + assertThat(reply).isEqualTo("Your name is Priya."); + t.line(" findConversationIds() returns every conversation in the table: %s", fresh.findConversationIds()); + } + } + + @Test + void aConversationIdLongerThan36CharactersFailsBeforeTheModelIsCalled() { + try (Transcript t = new Transcript("09-jdbc-conversation-id-limit.txt", "conversation_id is VARCHAR(36)")) { + String tooLong = "user:asmhatre@example.com:thread:2026-09-24"; // 43 characters + t.line("conversation ID: \"%s\" (%d characters)", tooLong, tooLong.length()); + int callsBefore = model.prompts().size(); + Throwable thrown = null; + try { + ask(tooLong, "hello"); + } + catch (RuntimeException e) { + thrown = e; + } + assertThat(thrown).isNotNull(); + Throwable root = thrown; + while (root.getCause() != null) { + root = root.getCause(); + } + t.line("model calls made for this request: %d (the user message is saved before the model is called)", model.prompts().size() - callsBefore); + t.line("thrown: %s", thrown.getClass().getName()); + t.line("root cause: %s", root.getMessage().lines().findFirst().orElse("")); + t.line("rows stored for it afterwards: %d", + jdbc.queryForObject("SELECT count(*) FROM SPRING_AI_CHAT_MEMORY", Integer.class)); + assertThat(model.prompts().size() - callsBefore).isZero(); + } + } + + @Test + void toolCallMessagesAreDroppedOnSave(CapturedOutput output) { + try (Transcript t = new Transcript("10-jdbc-drops-tool-messages.txt", "What JdbcChatMemoryRepository.saveAll keeps")) { + String id = "tool-conv"; + AssistantMessage withToolCall = AssistantMessage.builder().content("") + .toolCalls(List.of(new AssistantMessage.ToolCall("call-1", "function", "weather", "{\"city\":\"Mumbai\"}"))).build(); + ToolResponseMessage toolResult = ToolResponseMessage.builder() + .responses(List.of(new ToolResponseMessage.ToolResponse("call-1", "weather", "31C"))).build(); + List turn = List.of(new UserMessage("Weather in Mumbai?"), withToolCall, toolResult, + new AssistantMessage("31C and humid.")); + t.line("saveAll() given 4 messages:"); + turn.forEach(m -> t.line(" %-9s %s", m.getMessageType(), m.getText().isEmpty() ? "(tool call weather)" : m.getText())); + repository.saveAll(id, turn); + List back = repository.findByConversationId(id); + t.blank().line("findByConversationId() returns %d:", back.size()); + back.forEach(m -> t.line(" %s", Show.one(m))); + String warning = output.getAll().lines().filter(l -> l.contains("JdbcChatMemoryRepository")).findFirst().orElse("(none)"); + t.blank().line("Logged by the repository: WARN %s", warning.replaceAll("^.*JdbcChatMemoryRepository\\s+:\\s+", "")); + assertThat(back).hasSize(2); + } + } + + private List twoWritersOnOneConversation(ChatMemory memoryUnderTest) throws Exception { + repository.saveAll("shared", List.of(new UserMessage("u0"), new AssistantMessage("a0"))); // seed without the barrier + Thread a = Thread.ofPlatform().start(() -> memoryUnderTest.add("shared", new UserMessage("from-A"))); + Thread b = Thread.ofPlatform().start(() -> memoryUnderTest.add("shared", new UserMessage("from-B"))); + a.join(); + b.join(); + return memoryUnderTest.get("shared").stream().map(Message::getText).collect(Collectors.toList()); + } + + @Test + void twoWritersOnOneConversationLoseAMessageBecauseSaveAllReplacesTheWholeConversation() throws Exception { + try (Transcript t = new Transcript("11-jdbc-lost-update.txt", "Two concurrent add() calls on one conversation")) { + InterleavingRepository racy = new InterleavingRepository(repository, false); + List stored = twoWritersOnOneConversation( + MessageWindowChatMemory.builder().chatMemoryRepository(racy).build()); + String first = racy.saveOrder.getFirst(); + String last = racy.saveOrder.getLast(); + t.line("seeded with: u0, a0"); + t.line("two writers each add one message; both read the conversation before either writes it back"); + t.line("first writer to save: adds a 3-message list, saved"); + t.line("second writer to save: adds a different 3-message list, replaces the first"); + t.line("stored afterwards: %s (%d messages, 4 were sent)", + stored.toString().replace(last, ""), stored.size()); + t.line("the first writer's message is %s, the second writer's is %s", + stored.contains(first) ? "kept" : "GONE", stored.contains(last) ? "kept" : "gone"); + assertThat(stored).containsExactly("u0", "a0", last); + } + } + + @Test + void saveAllsThatOverlapInTimeAreWorseThanALostUpdate() throws Exception { + try (Transcript t = new Transcript("12-jdbc-overlapping-saves.txt", "Two saveAll() calls released at the same instant")) { + InterleavingRepository racy = new InterleavingRepository(repository, true); + List stored = twoWritersOnOneConversation( + MessageWindowChatMemory.builder().chatMemoryRepository(racy).build()); + t.line("seeded with u0, a0; two writers add from-A and from-B; both saveAll() calls start together"); + t.line("The outcome differs from run to run, so the exact count is not printed here. Across repeated runs on"); + t.line("PostgreSQL 16 (READ COMMITTED) it is one of two things, and never the 4 messages that were sent:"); + t.line(" 3 messages -> one writer's message is gone"); + t.line(" 6 messages -> both transactions inserted, history is duplicated, sequence_id values collide"); + assertThat(stored.size()).isIn(3, 6); + } + } + + /** The smallest possible fix that works inside one JVM: one lock per conversation. */ + private static final class LockingChatMemory implements ChatMemory { + private final ChatMemory delegate; + private final java.util.concurrent.ConcurrentHashMap locks = new java.util.concurrent.ConcurrentHashMap<>(); + + LockingChatMemory(ChatMemory delegate) { + this.delegate = delegate; + } + + @Override + public void add(String id, List messages) { + synchronized (locks.computeIfAbsent(id, k -> new Object())) { + delegate.add(id, messages); + } + } + + @Override + public List get(String id) { + return delegate.get(id); + } + + @Override + public void clear(String id) { + delegate.clear(id); + } + } + + @Test + void aPerConversationLockRemovesTheLossOnASingleInstance() throws Exception { + try (Transcript t = new Transcript("14-jdbc-lost-update-locked.txt", "Same two writers, per-conversation lock")) { + MessageWindowChatMemory racy = MessageWindowChatMemory.builder() + .chatMemoryRepository(new InterleavingRepository(repository, false)).build(); + List stored = twoWritersOnOneConversation(new LockingChatMemory(racy)); + t.line("stored afterwards: %d messages", stored.size()); + t.line("contains from-A: %s, from-B: %s", stored.contains("from-A"), stored.contains("from-B")); + assertThat(stored).hasSize(4).contains("from-A", "from-B"); + } + } + + @Test + void aOneXTableStopsWorkingAfterTheUpgrade() { + Services.require("PostgreSQL", Services.PG_PORT); + try (Transcript t = new Transcript("13-jdbc-1x-table-upgrade.txt", "A Spring AI 1.x table under the 2.0 repository")) { + jdbc.execute("DROP SCHEMA IF EXISTS legacy1x CASCADE"); + jdbc.execute("CREATE SCHEMA legacy1x"); + jdbc.execute(""" + CREATE TABLE legacy1x.SPRING_AI_CHAT_MEMORY ( + conversation_id VARCHAR(36) NOT NULL, + content TEXT NOT NULL, + type VARCHAR(10) NOT NULL CHECK (type IN ('USER', 'ASSISTANT', 'SYSTEM', 'TOOL')), + "timestamp" TIMESTAMP NOT NULL)"""); + jdbc.update("INSERT INTO legacy1x.SPRING_AI_CHAT_MEMORY VALUES ('old-conv', 'My name is Priya', 'USER', now() - interval '2 minutes')"); + jdbc.update("INSERT INTO legacy1x.SPRING_AI_CHAT_MEMORY VALUES ('old-conv', 'Nice to meet you, Priya.', 'ASSISTANT', now() - interval '1 minute')"); + t.line("A 1.1.x-shaped table (no sequence_id column) holds 2 rows for \"old-conv\"."); + + var ds = new org.springframework.jdbc.datasource.DriverManagerDataSource( + "jdbc:postgresql://127.0.0.1:" + Services.PG_PORT + "/chatdb?currentSchema=legacy1x", "chat", "chat"); + JdbcTemplate legacy = new JdbcTemplate(ds); + JdbcChatMemoryRepository repo = JdbcChatMemoryRepository.builder().jdbcTemplate(legacy) + .dialect(new org.springframework.ai.chat.memory.repository.jdbc.PostgresChatMemoryRepositoryDialect()).build(); + + t.blank().line("2.0.1 repository, findByConversationId(\"old-conv\"):"); + Throwable read = null; + try { + repo.findByConversationId("old-conv"); + } + catch (RuntimeException e) { + read = e; + } + assertThat(read).isNotNull(); + t.line(" %s", firstLine(read)); + + t.blank().line("Migration: add the column, backfill it from the timestamp order, then make it NOT NULL:"); + String[] steps = { + "ALTER TABLE SPRING_AI_CHAT_MEMORY ADD COLUMN sequence_id BIGINT", + "UPDATE SPRING_AI_CHAT_MEMORY m SET sequence_id = s.rn FROM (SELECT ctid AS c, row_number() OVER (PARTITION BY conversation_id ORDER BY \"timestamp\") - 1 AS rn FROM SPRING_AI_CHAT_MEMORY) s WHERE m.ctid = s.c", + "ALTER TABLE SPRING_AI_CHAT_MEMORY ALTER COLUMN sequence_id SET NOT NULL", + "CREATE INDEX IF NOT EXISTS SPRING_AI_CHAT_MEMORY_CONVERSATION_ID_SEQUENCE_ID_IDX ON SPRING_AI_CHAT_MEMORY(conversation_id, sequence_id)" + }; + for (String step : steps) { + legacy.execute(step); + t.line(" %s", step.length() > 110 ? step.substring(0, 107) + "..." : step); + } + List migrated = repo.findByConversationId("old-conv"); + t.blank().line("findByConversationId(\"old-conv\") after the migration:"); + migrated.forEach(m -> t.line(" %s", Show.one(m))); + assertThat(migrated).extracting(Message::getText).containsExactly("My name is Priya", "Nice to meet you, Priya."); + jdbc.execute("DROP SCHEMA legacy1x CASCADE"); + } + } + + private static String firstLine(Throwable t) { + Throwable root = t; + while (root.getCause() != null) { + root = root.getCause(); + } + return root.getMessage().lines().findFirst().orElse(""); + } +} diff --git a/chat-memory/src/test/java/com/ankurm/chatmemory/OwnershipLeakTest.java b/chat-memory/src/test/java/com/ankurm/chatmemory/OwnershipLeakTest.java new file mode 100644 index 0000000..671aa97 --- /dev/null +++ b/chat-memory/src/test/java/com/ankurm/chatmemory/OwnershipLeakTest.java @@ -0,0 +1,41 @@ +package com.ankurm.chatmemory; + +import static org.assertj.core.api.Assertions.assertThat; + +import com.ankurm.chatmemory.support.Http; +import com.ankurm.chatmemory.support.TestModelConfig; +import com.ankurm.chatmemory.support.Transcript; +import org.junit.jupiter.api.Test; +import org.springframework.beans.factory.annotation.Value; +import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.context.annotation.Import; + +/** The same endpoints with the ownership check off: a conversation ID is a bearer token for someone's history. */ +@SpringBootTest(webEnvironment = SpringBootTest.WebEnvironment.RANDOM_PORT, + properties = { "spring.ai.model.chat=none", "app.ownership.enforce=false" }) +@Import(TestModelConfig.class) +class OwnershipLeakTest { + + @Value("${local.server.port}") + int port; + + @Test + void anyoneWhoHoldsTheConversationIdCanReadAndContinueIt() { + try (Transcript t = new Transcript("21-ownership-unenforced.txt", "app.ownership.enforce=false")) { + Http http = new Http(port); + String id = Http.field(http.post("alice", "/conversations", "").body(), "conversationId"); + http.post("alice", "/conversations/" + id + "/messages", "{\"text\":\"My name is Priya\"}"); + t.line("alice starts a conversation and says \"My name is Priya\""); + + Http.Reply read = http.get("bob", "/conversations/" + id + "/messages"); + t.blank().line("bob: GET /conversations//messages -> %d", read.status()); + t.line(" body: %s", read.body()); + Http.Reply write = http.post("bob", "/conversations/" + id + "/messages", "{\"text\":\"What is my name?\"}"); + t.line("bob: POST \"What is my name?\" -> %d %s", write.status(), write.body()); + + assertThat(read.status()).isEqualTo(200); + assertThat(read.body()).contains("My name is Priya"); + assertThat(write.body()).contains("Your name is Priya."); + } + } +} diff --git a/chat-memory/src/test/java/com/ankurm/chatmemory/OwnershipTest.java b/chat-memory/src/test/java/com/ankurm/chatmemory/OwnershipTest.java new file mode 100644 index 0000000..03b7741 --- /dev/null +++ b/chat-memory/src/test/java/com/ankurm/chatmemory/OwnershipTest.java @@ -0,0 +1,42 @@ +package com.ankurm.chatmemory; + +import static org.assertj.core.api.Assertions.assertThat; + +import com.ankurm.chatmemory.support.Http; +import com.ankurm.chatmemory.support.TestModelConfig; +import com.ankurm.chatmemory.support.Transcript; +import org.junit.jupiter.api.Test; +import org.springframework.beans.factory.annotation.Value; +import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.context.annotation.Import; + +/** Per-user conversations over HTTP, with the ownership check on (the default). */ +@SpringBootTest(webEnvironment = SpringBootTest.WebEnvironment.RANDOM_PORT, properties = "spring.ai.model.chat=none") +@Import(TestModelConfig.class) +class OwnershipTest { + + @Value("${local.server.port}") + int port; + + @Test + void anotherUserGetsA403ForSomeoneElsesConversationId() { + try (Transcript t = new Transcript("20-ownership-enforced.txt", "app.ownership.enforce=true (default)")) { + Http http = new Http(port); + String id = Http.field(http.post("alice", "/conversations", "").body(), "conversationId"); + http.post("alice", "/conversations/" + id + "/messages", "{\"text\":\"My name is Priya\"}"); + t.line("alice starts a conversation and says \"My name is Priya\""); + + Http.Reply read = http.get("bob", "/conversations/" + id + "/messages"); + t.blank().line("bob: GET /conversations//messages -> %d", read.status()); + Http.Reply write = http.post("bob", "/conversations/" + id + "/messages", "{\"text\":\"What is my name?\"}"); + t.line("bob: POST /conversations//messages -> %d", write.status()); + Http.Reply own = http.get("alice", "/conversations/" + id + "/messages"); + t.line("alice: GET /conversations//messages -> %d, %d messages", own.status(), + own.body().split("\"role\"").length - 1); + + assertThat(read.status()).isEqualTo(403); + assertThat(write.status()).isEqualTo(403); + assertThat(own.status()).isEqualTo(200); + } + } +} diff --git a/chat-memory/src/test/java/com/ankurm/chatmemory/RedisCapsAndTtlTest.java b/chat-memory/src/test/java/com/ankurm/chatmemory/RedisCapsAndTtlTest.java new file mode 100644 index 0000000..3060ebc --- /dev/null +++ b/chat-memory/src/test/java/com/ankurm/chatmemory/RedisCapsAndTtlTest.java @@ -0,0 +1,68 @@ +package com.ankurm.chatmemory; + +import static org.assertj.core.api.Assertions.assertThat; + +import java.util.List; +import java.util.Set; + +import com.ankurm.chatmemory.support.Services; +import com.ankurm.chatmemory.support.Show; +import com.ankurm.chatmemory.support.TestModelConfig; +import com.ankurm.chatmemory.support.Transcript; +import org.junit.jupiter.api.BeforeAll; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.springframework.ai.chat.memory.ChatMemoryRepository; +import org.springframework.ai.chat.messages.AssistantMessage; +import org.springframework.ai.chat.messages.Message; +import org.springframework.ai.chat.messages.UserMessage; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.context.annotation.Import; +import redis.clients.jedis.RedisClient; + +/** time-to-live and max-messages-per-conversation, with a 2-second TTL so the expiry is observable. */ +@SpringBootTest(properties = { "spring.ai.model.chat=none", "spring.profiles.active=redis", + "spring.ai.chat.memory.repository.redis.time-to-live=2s", + "spring.ai.chat.memory.repository.redis.max-messages-per-conversation=4" }) +@Import(TestModelConfig.class) +class RedisCapsAndTtlTest { + + @BeforeAll + static void needRedisStack() { + Services.require("Redis Stack", Services.REDIS_STACK_PORT); + } + + @Autowired + ChatMemoryRepository repository; + @Autowired + RedisClient redis; + + @BeforeEach + void clean() { + Set keys = redis.keys("chat-memory:*"); + if (!keys.isEmpty()) { + redis.del(keys.toArray(new String[0])); + } + } + + @Test + void conversationsExpireAndTheRepositoryHasItsOwnMessageCap() throws Exception { + try (Transcript t = new Transcript("18-redis-ttl-and-cap.txt", "Redis time-to-live = 2s, max-messages-per-conversation = 4")) { + repository.saveAll("ttl-conv", List.of(new UserMessage("u1"), new AssistantMessage("a1"))); + t.line("saved 2 messages; readable now: %s", Show.inline(repository.findByConversationId("ttl-conv"))); + Thread.sleep(3000); + List after = repository.findByConversationId("ttl-conv"); + t.line("3 seconds later: %s", after.isEmpty() ? "(gone: Redis expired the keys)" : Show.inline(after)); + assertThat(after).isEmpty(); + + repository.saveAll("cap-conv", List.of(new UserMessage("u1"), new AssistantMessage("a1"), new UserMessage("u2"), + new AssistantMessage("a2"), new UserMessage("u3"), new AssistantMessage("a3"))); + List capped = repository.findByConversationId("cap-conv"); + t.blank().line("saved 6 messages with max-messages-per-conversation=4, read back %d:", capped.size()); + t.line(" %s", Show.inline(capped)); + Set keys = redis.keys("chat-memory:*"); + t.line("keys in Redis for it: %d", keys.size()); + } + } +} diff --git a/chat-memory/src/test/java/com/ankurm/chatmemory/RedisPlainFailureTest.java b/chat-memory/src/test/java/com/ankurm/chatmemory/RedisPlainFailureTest.java new file mode 100644 index 0000000..ad71415 --- /dev/null +++ b/chat-memory/src/test/java/com/ankurm/chatmemory/RedisPlainFailureTest.java @@ -0,0 +1,42 @@ +package com.ankurm.chatmemory; + +import static org.assertj.core.api.Assertions.assertThat; + +import com.ankurm.chatmemory.support.Services; +import com.ankurm.chatmemory.support.Transcript; +import org.junit.jupiter.api.BeforeAll; +import org.junit.jupiter.api.Test; +import org.springframework.ai.chat.memory.repository.redis.RedisChatMemoryRepository; +import redis.clients.jedis.RedisClient; + +/** The Redis repository is not built on plain Redis: it needs RediSearch and RedisJSON. Point it at a plain server and read the error. */ +class RedisPlainFailureTest { + + @BeforeAll + static void needPlainRedis() { + Services.require("plain Redis", Services.REDIS_PLAIN_PORT); + } + + @Test + void plainRedisWithoutModulesFailsAtStartup() { + try (Transcript t = new Transcript("19-redis-plain-fails.txt", "RedisChatMemoryRepository against plain Redis 7.0 (no modules)")) { + RedisClient plain = RedisClient.create("redis://127.0.0.1:" + Services.REDIS_PLAIN_PORT); + t.line("server: %s", plain.info("server").lines().filter(l -> l.startsWith("redis_version")).findFirst().orElse("?")); + Throwable thrown = null; + try { + RedisChatMemoryRepository.builder().jedisClient(plain).initializeSchema(true).build(); + } + catch (RuntimeException e) { + thrown = e; + } + assertThat(thrown).isNotNull(); + Throwable root = thrown; + while (root.getCause() != null) { + root = root.getCause(); + } + t.line("RedisChatMemoryRepository.builder()...initializeSchema(true).build() threw:"); + t.line(" %s", thrown.getClass().getName()); + t.line(" root cause: %s: %s", root.getClass().getSimpleName(), root.getMessage()); + } + } +} diff --git a/chat-memory/src/test/java/com/ankurm/chatmemory/RedisSilentFallbackTest.java b/chat-memory/src/test/java/com/ankurm/chatmemory/RedisSilentFallbackTest.java new file mode 100644 index 0000000..1a4a1e3 --- /dev/null +++ b/chat-memory/src/test/java/com/ankurm/chatmemory/RedisSilentFallbackTest.java @@ -0,0 +1,82 @@ +package com.ankurm.chatmemory; + +import static org.assertj.core.api.Assertions.assertThat; + +import java.util.Set; + +import com.ankurm.chatmemory.support.Services; +import com.ankurm.chatmemory.support.TestModelConfig; +import com.ankurm.chatmemory.support.Transcript; +import org.junit.jupiter.api.BeforeAll; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.springframework.boot.test.system.CapturedOutput; +import org.springframework.boot.test.system.OutputCaptureExtension; +import org.springframework.ai.chat.client.ChatClient; +import org.springframework.ai.chat.memory.ChatMemory; +import org.springframework.ai.chat.memory.ChatMemoryRepository; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.boot.autoconfigure.condition.ConditionEvaluationReport; +import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.context.ConfigurableApplicationContext; +import org.springframework.context.annotation.Import; +import redis.clients.jedis.RedisClient; + +/** + * The Redis starter, the Redis profile, a running Redis Stack -- and the history goes to the JVM heap. + * Same application as RedisStackTest with one property flipped: app.memory.redis.explicit-repository=false. + */ +@SpringBootTest(properties = { "spring.ai.model.chat=none", "spring.profiles.active=redis", + "app.memory.redis.explicit-repository=false" }) +@Import(TestModelConfig.class) +@ExtendWith(OutputCaptureExtension.class) +class RedisSilentFallbackTest { + + @BeforeAll + static void needRedisStack() { + Services.require("Redis Stack", Services.REDIS_STACK_PORT); + } + + @Autowired + ChatClient chat; + @Autowired + ChatMemoryRepository repository; + @Autowired + ConfigurableApplicationContext context; + + @Test + void aCustomChatMemoryBeanMakesTheRedisRepositoryStepAsideWithoutAWord(CapturedOutput output) { + try (Transcript t = new Transcript("17-redis-silent-fallback.txt", "Redis starter + your own ChatMemory bean")) { + RedisClient probe = RedisClient.create("redis://127.0.0.1:" + Services.REDIS_STACK_PORT); + Set leftovers = probe.keys("chat-memory:*"); + if (!leftovers.isEmpty()) { + probe.del(leftovers.toArray(new String[0])); + } + Set before = probe.keys("chat-memory:*"); + t.line("application has: spring-ai-starter-model-chat-memory-repository-redis, profile \"redis\", Redis Stack reachable,"); + t.line("and its own @Bean ChatMemory (needed to set a window other than 20)"); + t.blank().line("ChatMemoryRepository bean in the context: %s", repository.getClass().getSimpleName()); + + chat.prompt().user("hello").advisors(a -> a.param(ChatMemory.CONVERSATION_ID, "fallback-conv")).call().content(); + Set after = probe.keys("chat-memory:*"); + t.line("Redis keys under \"chat-memory:\" before / after one chat call: %d / %d", before.size(), after.size()); + + ConditionEvaluationReport report = ConditionEvaluationReport.get(context.getBeanFactory()); + String outcome = report.getConditionAndOutcomesBySource().entrySet().stream() + .filter(e -> e.getKey().endsWith("RedisChatMemoryRepositoryAutoConfiguration#redisChatMemoryRepository")) + .flatMap(e -> e.getValue().stream()) + .map(o -> o.getOutcome().getMessage()).findFirst().orElse("(no condition report)"); + t.blank().line("Condition report for the autoconfigured Redis repository bean:"); + t.line(" %s", outcome.replace("SearchStrategy: all", "SearchStrategy: all")); + long warnings = output.getAll().lines() + .filter(l -> (l.contains(" WARN ") || l.contains(" ERROR ")) && (l.contains("Redis") || l.contains("ChatMemory"))) + .count(); + t.line("WARN/ERROR log lines mentioning Redis or ChatMemory, captured over the whole test class: %d", warnings); + assertThat(warnings).isZero(); + + assertThat(repository.getClass().getSimpleName()).isEqualTo("InMemoryChatMemoryRepository"); + assertThat(after).hasSameSizeAs(before); + assertThat(outcome).contains("org.springframework.ai.chat.memory.ChatMemory"); + } + } +} diff --git a/chat-memory/src/test/java/com/ankurm/chatmemory/RedisStackTest.java b/chat-memory/src/test/java/com/ankurm/chatmemory/RedisStackTest.java new file mode 100644 index 0000000..2abcac1 --- /dev/null +++ b/chat-memory/src/test/java/com/ankurm/chatmemory/RedisStackTest.java @@ -0,0 +1,108 @@ +package com.ankurm.chatmemory; + +import static org.assertj.core.api.Assertions.assertThat; + +import java.util.ArrayList; +import java.util.List; +import java.util.Set; + +import com.ankurm.chatmemory.support.InterleavingRepository; +import com.ankurm.chatmemory.support.ScriptedChatModel; +import com.ankurm.chatmemory.support.Services; +import com.ankurm.chatmemory.support.Show; +import com.ankurm.chatmemory.support.TestModelConfig; +import com.ankurm.chatmemory.support.Transcript; +import org.junit.jupiter.api.BeforeAll; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.springframework.ai.chat.client.ChatClient; +import org.springframework.ai.chat.memory.ChatMemory; +import org.springframework.ai.chat.memory.ChatMemoryRepository; +import org.springframework.ai.chat.memory.MessageWindowChatMemory; +import org.springframework.ai.chat.messages.AssistantMessage; +import org.springframework.ai.chat.messages.Message; +import org.springframework.ai.chat.messages.UserMessage; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.context.annotation.Import; +import redis.clients.jedis.RedisClient; + +/** RedisChatMemoryRepository against a real Redis Stack 7.4 (Redis + RediSearch + RedisJSON). */ +@SpringBootTest(properties = { "spring.ai.model.chat=none", "spring.profiles.active=redis" }) +@Import(TestModelConfig.class) +class RedisStackTest { + + @BeforeAll + static void needRedisStack() { + Services.require("Redis Stack", Services.REDIS_STACK_PORT); + } + + @Autowired + ChatClient chat; + @Autowired + ChatMemory memory; + @Autowired + ChatMemoryRepository repository; + @Autowired + ScriptedChatModel model; + @Autowired + RedisClient redis; + + @BeforeEach + void clean() { + Set keys = redis.keys("chat-memory:*"); + if (!keys.isEmpty()) { + redis.del(keys.toArray(new String[0])); + } + } + + @Test + @SuppressWarnings("unchecked") + void historyLivesInRedisJsonDocumentsWithAnExpiry() { + try (Transcript t = new Transcript("15-redis-round-trip.txt", "RedisChatMemoryRepository on Redis Stack")) { + t.line("ChatMemoryRepository bean: %s", repository.getClass().getSimpleName()); + String id = "6f0a1c9e-2b7d-4c55-9a53-0d1f3e7a2b10"; + String reply = chat.prompt().user("My name is Priya").advisors(a -> a.param(ChatMemory.CONVERSATION_ID, id)).call().content(); + t.line("call 1 reply: %s", reply); + + List keys = new ArrayList<>(redis.keys("chat-memory:*")); + keys.sort(String::compareTo); + t.blank().line("Keys under the default prefix \"chat-memory:\" (%d: one JSON document per message, plus a counter):", keys.size()); + for (String key : keys) { + t.line(" %s", key.replaceAll("\\d{10,}", "")); + } + Object doc = redis.jsonGet(keys.get(0)); + List fields = new ArrayList<>(((java.util.Map) doc).keySet()); + fields.sort(String::compareTo); + t.blank().line("Fields of one document: %s", fields); + long ttl = redis.ttl(keys.getFirst()); + t.line("TTL on each key (spring.ai.chat.memory.repository.redis.time-to-live=24h): between 86390 and 86400 s -> %s", + ttl > 86390 && ttl <= 86400); + assertThat(ttl).isBetween(86390L, 86400L); + + t.blank().line("Same conversation, second call \"What is my name?\":"); + t.line(" reply: %s", chat.prompt().user("What is my name?").advisors(a -> a.param(ChatMemory.CONVERSATION_ID, id)).call().content()); + t.line(" stored: %s", Show.inline(memory.get(id))); + assertThat(memory.get(id)).hasSize(4); + assertThat(repository.findConversationIds()).containsExactly(id); + } + } + + @Test + void twoWritersOnOneConversationLoseAMessageInRedisToo() throws Exception { + try (Transcript t = new Transcript("16-redis-lost-update.txt", "Two concurrent add() calls, Redis repository")) { + repository.saveAll("shared", List.of(new UserMessage("u0"), new AssistantMessage("a0"))); + InterleavingRepository racy = new InterleavingRepository(repository, false); + MessageWindowChatMemory racyMemory = MessageWindowChatMemory.builder().chatMemoryRepository(racy).build(); + Thread a = Thread.ofPlatform().start(() -> racyMemory.add("shared", new UserMessage("from-A"))); + Thread b = Thread.ofPlatform().start(() -> racyMemory.add("shared", new UserMessage("from-B"))); + a.join(); + b.join(); + List stored = racyMemory.get("shared").stream().map(Message::getText).toList(); + String last = racy.saveOrder.getLast(); + t.line("seeded with u0, a0; two writers add one message each, both read before either writes back"); + t.line("stored afterwards: %s (%d messages, 4 were sent)", stored.toString().replace(last, ""), stored.size()); + assertThat(stored).containsExactly("u0", "a0", last); + } + } +} diff --git a/chat-memory/src/test/java/com/ankurm/chatmemory/TokenBudgetChatMemoryTest.java b/chat-memory/src/test/java/com/ankurm/chatmemory/TokenBudgetChatMemoryTest.java new file mode 100644 index 0000000..b0dbda4 --- /dev/null +++ b/chat-memory/src/test/java/com/ankurm/chatmemory/TokenBudgetChatMemoryTest.java @@ -0,0 +1,48 @@ +package com.ankurm.chatmemory; + +import static org.assertj.core.api.Assertions.assertThat; + +import java.util.List; + +import com.ankurm.chatmemory.memory.TokenBudgetChatMemory; +import com.ankurm.chatmemory.support.Show; +import com.ankurm.chatmemory.support.Transcript; +import com.knuddels.jtokkit.api.EncodingType; +import org.junit.jupiter.api.Test; +import org.springframework.ai.chat.memory.InMemoryChatMemoryRepository; +import org.springframework.ai.chat.messages.AssistantMessage; +import org.springframework.ai.chat.messages.Message; +import org.springframework.ai.chat.messages.MessageType; +import org.springframework.ai.chat.messages.SystemMessage; +import org.springframework.ai.chat.messages.UserMessage; +import org.springframework.ai.tokenizer.JTokkitTokenCountEstimator; +import org.springframework.ai.tokenizer.TokenCountEstimator; + +/** The custom token-budget memory: one long message costs what it costs, however few messages there are. */ +class TokenBudgetChatMemoryTest { + + private static final TokenCountEstimator ESTIMATOR = new JTokkitTokenCountEstimator(EncodingType.O200K_BASE); + + @Test + void aLongMessageEvictsMoreHistoryThanAWindowOfTheSameCountWould() { + try (Transcript t = new Transcript("07-token-budget-memory.txt", "TokenBudgetChatMemory, budget = 120 tokens")) { + TokenBudgetChatMemory memory = new TokenBudgetChatMemory(new InMemoryChatMemoryRepository(), ESTIMATOR, 120); + memory.add("c", new SystemMessage("Be terse.")); + for (int n = 1; n <= 3; n++) { + memory.add("c", List.of(new UserMessage("short question " + n), new AssistantMessage("short answer " + n))); + } + t.line("three short turns: %s", Show.inline(memory.get("c"))); + + String pasted = "java.lang.IllegalStateException: Failed to load ApplicationContext ".repeat(8); + memory.add("c", List.of(new UserMessage(pasted), new AssistantMessage("Check the datasource URL."))); + List after = memory.get("c"); + t.line("then one pasted trace: %d messages kept: %s", after.size(), + Show.inline(after).replace(pasted, "")); + + assertThat(after.getFirst().getMessageType()).isEqualTo(MessageType.SYSTEM); + assertThat(after.get(1).getMessageType()).isEqualTo(MessageType.USER); + assertThat(after.stream().mapToInt(m -> ESTIMATOR.estimate(m.getText())).sum()).isLessThanOrEqualTo(120); + assertThat(after.size()).isLessThan(9); + } + } +} diff --git a/chat-memory/src/test/java/com/ankurm/chatmemory/TokenGrowthTest.java b/chat-memory/src/test/java/com/ankurm/chatmemory/TokenGrowthTest.java new file mode 100644 index 0000000..7b5ec3a --- /dev/null +++ b/chat-memory/src/test/java/com/ankurm/chatmemory/TokenGrowthTest.java @@ -0,0 +1,109 @@ +package com.ankurm.chatmemory; + +import static org.assertj.core.api.Assertions.assertThat; + +import java.util.ArrayList; +import java.util.List; +import java.util.function.Supplier; + +import com.ankurm.chatmemory.memory.TokenBudgetChatMemory; +import com.ankurm.chatmemory.support.Transcript; +import com.knuddels.jtokkit.api.EncodingType; +import org.junit.jupiter.api.Test; +import org.springframework.ai.chat.client.ChatClient; +import org.springframework.ai.chat.client.advisor.MessageChatMemoryAdvisor; +import org.springframework.ai.chat.memory.ChatMemory; +import org.springframework.ai.chat.memory.InMemoryChatMemoryRepository; +import org.springframework.ai.chat.memory.MessageWindowChatMemory; +import org.springframework.ai.chat.messages.AssistantMessage; +import org.springframework.ai.chat.messages.Message; +import org.springframework.ai.chat.model.ChatModel; +import org.springframework.ai.chat.model.ChatResponse; +import org.springframework.ai.chat.model.Generation; +import org.springframework.ai.chat.prompt.Prompt; +import org.springframework.ai.tokenizer.JTokkitTokenCountEstimator; +import org.springframework.ai.tokenizer.TokenCountEstimator; +import reactor.core.publisher.Flux; + +/** + * The cost side of memory: how many input tokens each call carries as a conversation grows, + * under four pruning strategies. Token counts are JTokkit estimates with the o200k_base encoding + * (the family OpenAI's current models use), not a provider's invoice. + */ +class TokenGrowthTest { + + private static final int TURNS = 30; + private static final TokenCountEstimator ESTIMATOR = new JTokkitTokenCountEstimator(EncodingType.O200K_BASE); + + private static final String ANSWER = "The staging deploy failed because the readiness probe timed out before the new pods " + + "finished warming their caches, so the rollout was halted and the previous ReplicaSet kept serving traffic."; + + /** Answers with a fixed paragraph and records how many tokens each prompt carried. */ + private static final class CountingModel implements ChatModel { + final List promptTokens = new ArrayList<>(); + + @Override + public ChatResponse call(Prompt prompt) { + int total = 0; + for (Message m : prompt.getInstructions()) { + total += ESTIMATOR.estimate(m.getText()); + } + promptTokens.add(total); + return new ChatResponse(List.of(new Generation(new AssistantMessage(ANSWER)))); + } + + @Override + public Flux stream(Prompt prompt) { + return Flux.just(call(prompt)); + } + } + + private static List run(Supplier memoryFactory) { + CountingModel model = new CountingModel(); + ChatMemory memory = memoryFactory.get(); + ChatClient client = ChatClient.builder(model) + .defaultSystem("You are a terse assistant.") + .defaultAdvisors(MessageChatMemoryAdvisor.builder(memory).build()) + .build(); + for (int turn = 1; turn <= TURNS; turn++) { + client.prompt() + .user("Question " + turn + " about the pipeline: why did stage " + turn + " fail on the staging cluster last night?") + .advisors(a -> a.param(ChatMemory.CONVERSATION_ID, "c")) + .call().content(); + } + return model.promptTokens; + } + + private static MessageWindowChatMemory window(int max) { + return MessageWindowChatMemory.builder().chatMemoryRepository(new InMemoryChatMemoryRepository()).maxMessages(max).build(); + } + + @Test + void inputTokensPerCallGrowLinearlyUntilThePruningStrategyCapsThem() { + List unbounded = run(() -> window(10_000)); + List w10 = run(() -> window(10)); + List w20 = run(() -> window(20)); + List budget = run(() -> new TokenBudgetChatMemory(new InMemoryChatMemoryRepository(), ESTIMATOR, 400)); + + try (Transcript t = new Transcript("06-token-growth.txt", "Input tokens per call, 30 turns, four pruning strategies")) { + t.line("Encoding: o200k_base (JTokkit estimate). Each turn: ~25-token question, ~40-token answer, ~6-token system prompt."); + t.blank().line("%-6s %10s %10s %10s %14s", "turn", "unbounded", "window=10", "window=20", "budget=400tok"); + for (int turn : new int[] {1, 2, 5, 10, 11, 15, 20, 21, 25, 30}) { + t.line("%-6d %10d %10d %10d %14d", turn, unbounded.get(turn - 1), w10.get(turn - 1), w20.get(turn - 1), budget.get(turn - 1)); + } + long su = unbounded.stream().mapToLong(Integer::longValue).sum(); + long s10 = w10.stream().mapToLong(Integer::longValue).sum(); + long s20 = w20.stream().mapToLong(Integer::longValue).sum(); + long sb = budget.stream().mapToLong(Integer::longValue).sum(); + t.blank().line("%-6s %10d %10d %10d %14d", "sum", su, s10, s20, sb); + t.line("sum as %% of unbounded: window=10 %.0f%%, window=20 %.0f%%, budget=400tok %.0f%%", + 100.0 * s10 / su, 100.0 * s20 / su, 100.0 * sb / su); + + assertThat(unbounded.get(29)).isGreaterThan(unbounded.get(0) * 20); + assertThat(w10.get(29)).isEqualTo(w10.get(15)); + assertThat(w20.get(29)).isEqualTo(w20.get(25)); + assertThat(budget.get(29)).isLessThanOrEqualTo(400 + 90); + assertThat(s10).isLessThan(su); + } + } +} diff --git a/chat-memory/src/test/java/com/ankurm/chatmemory/ToolCallingMemoryTest.java b/chat-memory/src/test/java/com/ankurm/chatmemory/ToolCallingMemoryTest.java new file mode 100644 index 0000000..5ee6c83 --- /dev/null +++ b/chat-memory/src/test/java/com/ankurm/chatmemory/ToolCallingMemoryTest.java @@ -0,0 +1,74 @@ +package com.ankurm.chatmemory; + +import static org.assertj.core.api.Assertions.assertThat; + +import java.util.List; + +import com.ankurm.chatmemory.support.ScriptedToolModel; +import com.ankurm.chatmemory.support.Show; +import com.ankurm.chatmemory.support.Transcript; +import org.junit.jupiter.api.Test; +import org.springframework.ai.chat.client.ChatClient; +import org.springframework.ai.chat.client.advisor.MessageChatMemoryAdvisor; +import org.springframework.ai.chat.client.advisor.ToolCallingAdvisor; +import org.springframework.ai.chat.memory.ChatMemory; +import org.springframework.ai.chat.memory.InMemoryChatMemoryRepository; +import org.springframework.ai.chat.memory.MessageWindowChatMemory; +import org.springframework.ai.chat.messages.AssistantMessage; +import org.springframework.ai.chat.messages.Message; +import org.springframework.ai.support.ToolCallbacks; +import org.springframework.ai.tool.annotation.Tool; +import org.springframework.ai.tool.annotation.ToolParam; + +/** + * Memory and tool calling together. The 2.0 advisor precedence puts the memory advisor + * (HIGHEST_PRECEDENCE + 200) outside the ToolCallingAdvisor (+ 300), so the memory sees one + * user message going in and one final answer coming out -- not the tool round trips between them. + */ +class ToolCallingMemoryTest { + + static class Weather { + @Tool(description = "Current temperature for a city") + String currentWeather(@ToolParam(description = "city name") String city) { + return "31C in " + city; + } + } + + private static ScriptedToolModel model() { + return ScriptedToolModel.builder() + .thenCallTools(new AssistantMessage.ToolCall("call-1", "function", "currentWeather", "{\"city\":\"Mumbai\"}")) + .thenRespond("It is 31C in Mumbai.") + .build(); + } + + @Test + void memoryOutsideTheToolLoopStoresOnlyTheQuestionAndTheFinalAnswer() { + try (Transcript t = new Transcript("24-memory-with-tool-calling.txt", "MessageChatMemoryAdvisor + ToolCallingAdvisor")) { + ScriptedToolModel model = model(); + ChatMemory memory = MessageWindowChatMemory.builder().chatMemoryRepository(new InMemoryChatMemoryRepository()).build(); + ChatClient client = ChatClient.builder(model) + .defaultToolCallbacks(ToolCallbacks.from(new Weather())) + .defaultAdvisors(ToolCallingAdvisor.builder().build(), MessageChatMemoryAdvisor.builder(memory).build()) + .build(); + + t.line("Advisor.DEFAULT_CHAT_MEMORY_PRECEDENCE_ORDER = HIGHEST_PRECEDENCE + %d", + org.springframework.ai.chat.client.advisor.api.Advisor.DEFAULT_CHAT_MEMORY_PRECEDENCE_ORDER - org.springframework.core.Ordered.HIGHEST_PRECEDENCE); + t.line("ToolCallingAdvisor.DEFAULT_ORDER = HIGHEST_PRECEDENCE + %d", + ToolCallingAdvisor.DEFAULT_ORDER - org.springframework.core.Ordered.HIGHEST_PRECEDENCE); + t.blank(); + String answer = client.prompt().user("Weather in Mumbai?") + .advisors(a -> a.param(ChatMemory.CONVERSATION_ID, "c")).call().content(); + t.line("answer: %s", answer); + t.line("model calls: %d", model.callCount()); + t.blank().line("What the model was sent on each call:"); + for (int i = 0; i < model.capturedPrompts().size(); i++) { + t.line(" call %d: %s", i + 1, Show.inline(model.capturedPrompts().get(i).getInstructions())); + } + List stored = memory.get("c"); + t.blank().line("Stored in memory afterwards:"); + stored.forEach(m -> t.line(" %s", Show.one(m))); + assertThat(stored).hasSize(2); + assertThat(stored.get(1).getText()).isEqualTo("It is 31C in Mumbai."); + } + } +} diff --git a/chat-memory/src/test/java/com/ankurm/chatmemory/WindowSemanticsTest.java b/chat-memory/src/test/java/com/ankurm/chatmemory/WindowSemanticsTest.java new file mode 100644 index 0000000..94e6d22 --- /dev/null +++ b/chat-memory/src/test/java/com/ankurm/chatmemory/WindowSemanticsTest.java @@ -0,0 +1,62 @@ +package com.ankurm.chatmemory; + +import static org.assertj.core.api.Assertions.assertThat; + +import java.util.List; + +import com.ankurm.chatmemory.support.Show; +import com.ankurm.chatmemory.support.Transcript; +import org.junit.jupiter.api.Test; +import org.springframework.ai.chat.memory.InMemoryChatMemoryRepository; +import org.springframework.ai.chat.memory.MessageWindowChatMemory; +import org.springframework.ai.chat.messages.AssistantMessage; +import org.springframework.ai.chat.messages.Message; +import org.springframework.ai.chat.messages.SystemMessage; +import org.springframework.ai.chat.messages.UserMessage; + +/** + * What {@code MessageWindowChatMemory} actually keeps, measured. No ChatClient, no advisor, no + * model: the memory object on its own, fed one turn at a time. + */ +class WindowSemanticsTest { + + private static MessageWindowChatMemory window(int max) { + return MessageWindowChatMemory.builder() + .chatMemoryRepository(new InMemoryChatMemoryRepository()) + .maxMessages(max) + .build(); + } + + private static void turn(MessageWindowChatMemory memory, int n) { + memory.add("c", List.of(new UserMessage("u" + n), new AssistantMessage("a" + n))); + } + + @Test + void windowKeepsTheSystemMessageAndTrimsToAUserBoundary() { + try (Transcript t = new Transcript("01-window-semantics.txt", "MessageWindowChatMemory, maxMessages = 4 and 5")) { + MessageWindowChatMemory four = window(4); + four.add("c", new SystemMessage("SYS-1")); + t.line("maxMessages = 4. One system message, then one user+assistant turn at a time:"); + for (int n = 1; n <= 4; n++) { + turn(four, n); + t.line(" after turn %d: %s", n, Show.inline(four.get("c"))); + } + assertThat(Show.inline(four.get("c"))).isEqualTo("S:SYS-1 | U:u4 | A:a4"); + + four.add("c", new SystemMessage("SYS-2")); + t.blank().line("A second, different system message arrives:"); + t.line(" after SYS-2: %s", Show.inline(four.get("c"))); + assertThat(four.get("c")).noneMatch(m -> m.getText().equals("SYS-1")); + + MessageWindowChatMemory five = window(5); + t.blank().line("maxMessages = 5 (odd), no system message:"); + for (int n = 1; n <= 4; n++) { + turn(five, n); + t.line(" after turn %d: %s", n, Show.inline(five.get("c"))); + } + List kept = five.get("c"); + assertThat(kept).hasSize(4); + assertThat(kept.getFirst().getText()).isEqualTo("u3"); + } + } +} diff --git a/chat-memory/src/test/java/com/ankurm/chatmemory/support/Http.java b/chat-memory/src/test/java/com/ankurm/chatmemory/support/Http.java new file mode 100644 index 0000000..e9c2418 --- /dev/null +++ b/chat-memory/src/test/java/com/ankurm/chatmemory/support/Http.java @@ -0,0 +1,46 @@ +package com.ankurm.chatmemory.support; + +import java.net.URI; +import java.net.http.HttpClient; +import java.net.http.HttpRequest; +import java.net.http.HttpResponse; +import java.util.regex.Matcher; +import java.util.regex.Pattern; + +/** A synchronous JDK HttpClient -- one send() per request, so a transcript line is one request. */ +public final class Http { + + public record Reply(int status, String body) { + } + + private final HttpClient client = HttpClient.newHttpClient(); + private final String base; + + public Http(int port) { + this.base = "http://127.0.0.1:" + port; + } + + public Reply post(String user, String path, String json) { + return send(HttpRequest.newBuilder(URI.create(base + path)).header("X-User", user) + .header("Content-Type", "application/json").POST(HttpRequest.BodyPublishers.ofString(json)).build()); + } + + public Reply get(String user, String path) { + return send(HttpRequest.newBuilder(URI.create(base + path)).header("X-User", user).GET().build()); + } + + private Reply send(HttpRequest request) { + try { + HttpResponse r = client.send(request, HttpResponse.BodyHandlers.ofString()); + return new Reply(r.statusCode(), r.body()); + } + catch (Exception e) { + throw new IllegalStateException(e); + } + } + + public static String field(String json, String name) { + Matcher m = Pattern.compile("\"" + name + "\":\"([^\"]*)\"").matcher(json); + return m.find() ? m.group(1) : null; + } +} diff --git a/chat-memory/src/test/java/com/ankurm/chatmemory/support/InterleavingRepository.java b/chat-memory/src/test/java/com/ankurm/chatmemory/support/InterleavingRepository.java new file mode 100644 index 0000000..9b7b829 --- /dev/null +++ b/chat-memory/src/test/java/com/ankurm/chatmemory/support/InterleavingRepository.java @@ -0,0 +1,87 @@ +package com.ankurm.chatmemory.support; + +import java.util.List; +import java.util.concurrent.BrokenBarrierException; +import java.util.concurrent.CopyOnWriteArrayList; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.CyclicBarrier; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.TimeoutException; +import java.util.concurrent.atomic.AtomicBoolean; + +import org.springframework.ai.chat.memory.ChatMemoryRepository; +import org.springframework.ai.chat.messages.Message; + +/** + * Holds every reader at a barrier until both have read, reproducing the interleaving a busy + * database produces by chance. {@code overlapWrites} additionally releases both writers into + * {@code saveAll} at the same instant; otherwise the second writer waits for the first to finish. + */ +public final class InterleavingRepository implements ChatMemoryRepository { + private final ChatMemoryRepository delegate; + private final boolean overlapWrites; + private final CyclicBarrier bothHaveRead = new CyclicBarrier(2); + private final CyclicBarrier bothAboutToWrite = new CyclicBarrier(2); + private final CountDownLatch firstWriterDone = new CountDownLatch(1); + private final AtomicBoolean firstWriter = new AtomicBoolean(true); + public final List saveOrder = new CopyOnWriteArrayList<>(); + + public InterleavingRepository(ChatMemoryRepository delegate, boolean overlapWrites) { + this.delegate = delegate; + this.overlapWrites = overlapWrites; + } + + @Override + public List findConversationIds() { + return delegate.findConversationIds(); + } + + @Override + public List findByConversationId(String id) { + List read = delegate.findByConversationId(id); + await(bothHaveRead); + return read; + } + + @Override + public void saveAll(String id, List messages) { + String writer = messages.getLast().getText(); + if (overlapWrites) { + await(bothAboutToWrite); + delegate.saveAll(id, messages); + return; + } + if (firstWriter.compareAndSet(true, false)) { + saveOrder.add(writer); + delegate.saveAll(id, messages); + firstWriterDone.countDown(); + } + else { + try { + firstWriterDone.await(1500, TimeUnit.MILLISECONDS); + } + catch (InterruptedException e) { + throw new IllegalStateException(e); + } + saveOrder.add(writer); + delegate.saveAll(id, messages); + } + } + + @Override + public void deleteByConversationId(String id) { + delegate.deleteByConversationId(id); + } + + private static void await(CyclicBarrier barrier) { + try { + barrier.await(1500, TimeUnit.MILLISECONDS); + } + catch (TimeoutException | BrokenBarrierException e) { + // the other writer is blocked behind a lock (or the barrier already gave up): carry on + } + catch (InterruptedException e) { + throw new IllegalStateException(e); + } + } +} diff --git a/chat-memory/src/test/java/com/ankurm/chatmemory/support/ScriptedChatModel.java b/chat-memory/src/test/java/com/ankurm/chatmemory/support/ScriptedChatModel.java new file mode 100644 index 0000000..8f05f60 --- /dev/null +++ b/chat-memory/src/test/java/com/ankurm/chatmemory/support/ScriptedChatModel.java @@ -0,0 +1,75 @@ +package com.ankurm.chatmemory.support; + +import java.util.Collections; +import java.util.List; +import java.util.concurrent.CopyOnWriteArrayList; +import java.util.regex.Matcher; +import java.util.regex.Pattern; + +import org.springframework.ai.chat.messages.AssistantMessage; +import org.springframework.ai.chat.messages.Message; +import org.springframework.ai.chat.messages.MessageType; +import org.springframework.ai.chat.model.ChatModel; +import org.springframework.ai.chat.model.ChatResponse; +import org.springframework.ai.chat.model.Generation; +import org.springframework.ai.chat.prompt.Prompt; +import reactor.core.publisher.Flux; + +/** + * A scripted {@link ChatModel}: no network, no API key, no randomness. It records every + * {@link Prompt} it is sent, which is the whole point -- what a memory test needs to prove is + * which messages reached the model, and that is exactly what a recorded prompt shows. + * + *

Its only "intelligence" is one rule that makes memory observable: if the last user message + * asks for the user's name, it looks for "my name is X" in the earlier messages of the + * same prompt and repeats X, or says it does not know. Nothing here says anything about what a + * real model would reply; it proves what the model would have been given. + */ +public class ScriptedChatModel implements ChatModel { + + private static final Pattern NAME = Pattern.compile("(?i)my name is (\\w+)"); + + private final List prompts = new CopyOnWriteArrayList<>(); + + public List prompts() { + return Collections.unmodifiableList(prompts); + } + + public Prompt lastPrompt() { + return prompts.getLast(); + } + + @Override + public ChatResponse call(Prompt prompt) { + prompts.add(prompt); + return new ChatResponse(List.of(new Generation(new AssistantMessage(reply(prompt))))); + } + + /** {@code ChatModel#stream} does not delegate to {@code call}; it throws unless overridden. */ + @Override + public Flux stream(Prompt prompt) { + return Flux.just(call(prompt)); + } + + private static String reply(Prompt prompt) { + List messages = prompt.getInstructions(); + Message last = messages.getLast(); + String text = last.getText(); + Matcher self = NAME.matcher(text); + if (self.find()) { + return "Nice to meet you, " + self.group(1) + "."; + } + if (text.toLowerCase().contains("what is my name")) { + for (Message m : messages.subList(0, messages.size() - 1)) { + if (m.getMessageType() == MessageType.USER) { + Matcher earlier = NAME.matcher(m.getText()); + if (earlier.find()) { + return "Your name is " + earlier.group(1) + "."; + } + } + } + return "I do not know your name yet."; + } + return "Noted."; + } +} diff --git a/chat-memory/src/test/java/com/ankurm/chatmemory/support/ScriptedToolModel.java b/chat-memory/src/test/java/com/ankurm/chatmemory/support/ScriptedToolModel.java new file mode 100644 index 0000000..6f19cd5 --- /dev/null +++ b/chat-memory/src/test/java/com/ankurm/chatmemory/support/ScriptedToolModel.java @@ -0,0 +1,117 @@ +package com.ankurm.chatmemory.support; + +import org.springframework.ai.chat.messages.AssistantMessage; +import org.springframework.ai.chat.model.ChatModel; +import org.springframework.ai.chat.model.ChatResponse; +import org.springframework.ai.chat.model.Generation; +import org.springframework.ai.chat.prompt.ChatOptions; +import org.springframework.ai.chat.prompt.Prompt; +import org.springframework.ai.model.tool.ToolCallingChatOptions; + +import java.util.ArrayDeque; +import java.util.ArrayList; +import java.util.Deque; +import java.util.List; +import java.util.Queue; +import java.util.concurrent.CopyOnWriteArrayList; + +/** + * A hand-written {@link ChatModel} that returns a pre-programmed queue of {@link ChatResponse}s, + * one per {@link #call(Prompt)} invocation, instead of calling a real LLM API. {@code ChatModel} + * has exactly one abstract method -- confirmed with {@code javap} against + * {@code spring-ai-model-2.0.1.jar}, everything else on the interface has a default implementation + * -- so this is the entire surface a deterministic test double needs to implement. + * + *

{@link org.springframework.ai.chat.client.advisor.ToolCallingAdvisor ToolCallingAdvisor} (and + * its subclass, {@code ToolSearchToolCallingAdvisor}) drive the tool-calling loop by calling the + * underlying {@link ChatModel} once per round: once to get the model's first response (which may + * contain tool calls), then once more per round of tool results fed back in, until a response + * comes back with no tool calls. Queuing responses here lets a test assert the exact shape of that + * loop -- how many rounds it took, what tool calls appeared, what the final answer was -- without + * an API key, network access, or the nondeterminism of an actual model. + * + *

Every {@link Prompt} the advisor sends is recorded in {@link #capturedPrompts()} so a test can + * also assert on what the advisor sent back on the next round -- in particular, that a + * {@code ToolResponseMessage} was appended after tool execution. + */ +public class ScriptedToolModel implements ChatModel { + + private final Queue script; + private final List capturedPrompts = new CopyOnWriteArrayList<>(); + + private ScriptedToolModel(Deque script) { + this.script = script; + } + + public static Builder builder() { + return new Builder(); + } + + @Override + public ChatResponse call(Prompt prompt) { + this.capturedPrompts.add(prompt); + ChatResponse next = this.script.poll(); + if (next == null) { + throw new IllegalStateException( + "ScriptedChatModel ran out of queued responses after " + this.capturedPrompts.size() + + " calls. Prompts so far: " + this.capturedPrompts); + } + return next; + } + + public List capturedPrompts() { + return List.copyOf(this.capturedPrompts); + } + + /** + * {@code DefaultChatClientUtils} builds every outgoing {@link Prompt}'s options from + * {@code chatModel.getOptions().mutate()} -- not {@code getDefaultOptions()}, which is a + * separate default method nobody in the request-building path actually calls. Confirmed by + * disassembling both: {@link ChatModel#getOptions()}'s default body is a bare + * {@code ChatOptions.builder().build()}, a plain {@link ChatOptions} that is not a + * {@link ToolCallingChatOptions}. Since {@code ToolCallingAdvisor.adviseCall} starts with an + * {@code instanceof ToolCallingChatOptions} check on that exact object and falls straight + * through to the underlying model with no tool loop at all when it fails, leaving this method's + * default in place silently turns every tool call in this repository into a no-op -- confirmed + * the hard way, by a first version of this class that overrode {@code getDefaultOptions()} + * instead and watched every test below get back an empty answer after exactly one model call. + * Real providers (OpenAI, Anthropic) return their own {@code ToolCallingChatOptions} + * implementation from {@code getOptions()} for the same reason. + */ + @Override + public ChatOptions getOptions() { + return ToolCallingChatOptions.builder().build(); + } + + public int callCount() { + return this.capturedPrompts.size(); + } + + public static final class Builder { + + private final Deque script = new ArrayDeque<>(); + + private Builder() { + } + + /** Queues a plain-text final answer with no tool calls -- ends the tool-calling loop. */ + public Builder thenRespond(String text) { + this.script.add(new ChatResponse(List.of(new Generation(new AssistantMessage(text))))); + return this; + } + + /** Queues an assistant turn that calls one or more tools, continuing the loop. */ + public Builder thenCallTools(AssistantMessage.ToolCall... toolCalls) { + AssistantMessage message = AssistantMessage.builder() + .content("") + .toolCalls(List.of(toolCalls)) + .build(); + this.script.add(new ChatResponse(List.of(new Generation(message)))); + return this; + } + + public ScriptedToolModel build() { + return new ScriptedToolModel(new ArrayDeque<>(this.script)); + } + } +} diff --git a/chat-memory/src/test/java/com/ankurm/chatmemory/support/Services.java b/chat-memory/src/test/java/com/ankurm/chatmemory/support/Services.java new file mode 100644 index 0000000..515df23 --- /dev/null +++ b/chat-memory/src/test/java/com/ankurm/chatmemory/support/Services.java @@ -0,0 +1,29 @@ +package com.ankurm.chatmemory.support; + +import java.net.InetSocketAddress; +import java.net.Socket; + +import org.junit.jupiter.api.Assumptions; + +/** Skips a test class (with a pointer to the script that fixes it) when a backing service is not running. */ +public final class Services { + + public static final int PG_PORT = 5440; + public static final int REDIS_STACK_PORT = 6390; + public static final int REDIS_PLAIN_PORT = 6391; + + private Services() { + } + + public static void require(String name, int port) { + boolean up; + try (Socket s = new Socket()) { + s.connect(new InetSocketAddress("127.0.0.1", port), 500); + up = true; + } + catch (Exception e) { + up = false; + } + Assumptions.assumeTrue(up, name + " is not listening on 127.0.0.1:" + port + " -- run scripts/services-up.sh first"); + } +} diff --git a/chat-memory/src/test/java/com/ankurm/chatmemory/support/Show.java b/chat-memory/src/test/java/com/ankurm/chatmemory/support/Show.java new file mode 100644 index 0000000..b4c894f --- /dev/null +++ b/chat-memory/src/test/java/com/ankurm/chatmemory/support/Show.java @@ -0,0 +1,31 @@ +package com.ankurm.chatmemory.support; + +import java.util.List; + +import org.springframework.ai.chat.messages.Message; + +/** One-line renderings of messages so transcripts stay short and diff cleanly. */ +public final class Show { + + private Show() { + } + + public static String one(Message m) { + return String.format("%-9s %s", m.getMessageType(), m.getText()); + } + + public static String inline(List messages) { + return messages.stream().map(Show::brief).reduce((a, b) -> a + " | " + b).orElse("(empty)"); + } + + private static String brief(Message m) { + char kind = m.getMessageType().name().charAt(0); + if (m instanceof org.springframework.ai.chat.messages.AssistantMessage a && a.hasToolCalls()) { + return kind + ":[tool call " + a.getToolCalls().getFirst().name() + "]"; + } + if (m instanceof org.springframework.ai.chat.messages.ToolResponseMessage r) { + return kind + ":[tool result " + r.getResponses().getFirst().responseData() + "]"; + } + return kind + ":" + m.getText(); + } +} diff --git a/chat-memory/src/test/java/com/ankurm/chatmemory/support/TestModelConfig.java b/chat-memory/src/test/java/com/ankurm/chatmemory/support/TestModelConfig.java new file mode 100644 index 0000000..e32b070 --- /dev/null +++ b/chat-memory/src/test/java/com/ankurm/chatmemory/support/TestModelConfig.java @@ -0,0 +1,14 @@ +package com.ankurm.chatmemory.support; + +import org.springframework.boot.test.context.TestConfiguration; +import org.springframework.context.annotation.Bean; + +/** Puts one shared {@link ScriptedChatModel} in the context so a test can inspect what it was sent. */ +@TestConfiguration(proxyBeanMethods = false) +public class TestModelConfig { + + @Bean + ScriptedChatModel scriptedChatModel() { + return new ScriptedChatModel(); + } +} diff --git a/chat-memory/src/test/java/com/ankurm/chatmemory/support/Transcript.java b/chat-memory/src/test/java/com/ankurm/chatmemory/support/Transcript.java new file mode 100644 index 0000000..adc3487 --- /dev/null +++ b/chat-memory/src/test/java/com/ankurm/chatmemory/support/Transcript.java @@ -0,0 +1,47 @@ +package com.ankurm.chatmemory.support; + +import java.io.IOException; +import java.io.PrintWriter; +import java.io.StringWriter; +import java.nio.file.Files; +import java.nio.file.Path; + +/** + * Writes a numbered transcript under {@code output/} (repository root, not {@code docs/}) and + * echoes it to the console. Every console block quoted in the article comes out of one of these + * files verbatim. + */ +public final class Transcript implements AutoCloseable { + + private final Path path; + private final StringWriter buffer = new StringWriter(); + private final PrintWriter out = new PrintWriter(buffer); + + public Transcript(String fileName, String title) { + this.path = Path.of("output", fileName); + out.println("# " + title); + out.println(); + } + + public Transcript line(String format, Object... args) { + out.println(args.length == 0 ? format : String.format(format, args)); + return this; + } + + public Transcript blank() { + out.println(); + return this; + } + + @Override + public void close() { + out.flush(); + try { + Files.createDirectories(path.getParent()); + Files.writeString(path, buffer.toString()); + } catch (IOException e) { + throw new IllegalStateException("could not write " + path, e); + } + System.out.print(buffer); + } +}