diff --git a/memory-repositories/spring-ai-model-chat-memory-repository-jdbc/src/main/java/org/springframework/ai/chat/memory/repository/jdbc/JdbcChatMemoryRepository.java b/memory-repositories/spring-ai-model-chat-memory-repository-jdbc/src/main/java/org/springframework/ai/chat/memory/repository/jdbc/JdbcChatMemoryRepository.java index 32e2918499..7c1292ba8d 100644 --- a/memory-repositories/spring-ai-model-chat-memory-repository-jdbc/src/main/java/org/springframework/ai/chat/memory/repository/jdbc/JdbcChatMemoryRepository.java +++ b/memory-repositories/spring-ai-model-chat-memory-repository-jdbc/src/main/java/org/springframework/ai/chat/memory/repository/jdbc/JdbcChatMemoryRepository.java @@ -43,6 +43,7 @@ import org.springframework.jdbc.core.RowMapper; import org.springframework.jdbc.datasource.DataSourceTransactionManager; import org.springframework.transaction.PlatformTransactionManager; +import org.springframework.transaction.TransactionDefinition; import org.springframework.transaction.support.TransactionTemplate; import org.springframework.util.Assert; @@ -81,11 +82,22 @@ private JdbcChatMemoryRepository(JdbcTemplate jdbcTemplate, JdbcChatMemoryReposi Assert.notNull(dialect, "dialect cannot be null"); this.jdbcTemplate = jdbcTemplate; this.dialect = dialect; + boolean usingDefaultTransactionManager = txManager == null; + PlatformTransactionManager effectiveTransactionManager; if (txManager == null) { Assert.state(jdbcTemplate.getDataSource() != null, "jdbcTemplate dataSource cannot be null"); - txManager = new DataSourceTransactionManager(jdbcTemplate.getDataSource()); + effectiveTransactionManager = new DataSourceTransactionManager(jdbcTemplate.getDataSource()); + } + else { + effectiveTransactionManager = txManager; + } + this.transactionTemplate = new TransactionTemplate(effectiveTransactionManager); + if (usingDefaultTransactionManager && dialect instanceof MysqlChatMemoryRepositoryDialect) { + // Under MySQL and MariaDB's default REPEATABLE READ isolation, concurrent + // deletes for different missing conversation IDs can acquire overlapping gap + // locks and deadlock when saveAll() subsequently inserts the messages. + this.transactionTemplate.setIsolationLevel(TransactionDefinition.ISOLATION_READ_COMMITTED); } - this.transactionTemplate = new TransactionTemplate(txManager); } @Override diff --git a/memory-repositories/spring-ai-model-chat-memory-repository-jdbc/src/test/java/org/springframework/ai/chat/memory/repository/jdbc/JdbcChatMemoryRepositoryBuilderTests.java b/memory-repositories/spring-ai-model-chat-memory-repository-jdbc/src/test/java/org/springframework/ai/chat/memory/repository/jdbc/JdbcChatMemoryRepositoryBuilderTests.java index 6e60bbf40a..04da567e0c 100644 --- a/memory-repositories/spring-ai-model-chat-memory-repository-jdbc/src/test/java/org/springframework/ai/chat/memory/repository/jdbc/JdbcChatMemoryRepositoryBuilderTests.java +++ b/memory-repositories/spring-ai-model-chat-memory-repository-jdbc/src/test/java/org/springframework/ai/chat/memory/repository/jdbc/JdbcChatMemoryRepositoryBuilderTests.java @@ -19,6 +19,7 @@ import java.sql.Connection; import java.sql.DatabaseMetaData; import java.sql.SQLException; +import java.util.List; import javax.sql.DataSource; @@ -26,10 +27,14 @@ import org.springframework.jdbc.core.JdbcTemplate; import org.springframework.transaction.PlatformTransactionManager; +import org.springframework.transaction.TransactionDefinition; +import org.springframework.transaction.TransactionStatus; import static org.assertj.core.api.Assertions.assertThat; import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.mockito.ArgumentMatchers.any; import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.verify; import static org.mockito.Mockito.when; /** @@ -254,4 +259,28 @@ void repositoryShouldUseProvidedJdbcTemplate() throws SQLException { assertThat(repository).extracting("jdbcTemplate").isSameAs(jdbcTemplate); } + @Test + void explicitTransactionManagerKeepsDefaultIsolation() { + DataSource dataSource = mock(DataSource.class); + JdbcTemplate jdbcTemplate = mock(JdbcTemplate.class); + PlatformTransactionManager transactionManager = mock(PlatformTransactionManager.class); + TransactionStatus transactionStatus = mock(TransactionStatus.class); + when(jdbcTemplate.getDataSource()).thenReturn(dataSource); + when(transactionManager.getTransaction(any())).thenAnswer(invocation -> { + TransactionDefinition definition = invocation.getArgument(0); + assertThat(definition.getIsolationLevel()).isEqualTo(TransactionDefinition.ISOLATION_DEFAULT); + return transactionStatus; + }); + + JdbcChatMemoryRepository repository = JdbcChatMemoryRepository.builder() + .jdbcTemplate(jdbcTemplate) + .dialect(new MysqlChatMemoryRepositoryDialect()) + .transactionManager(transactionManager) + .build(); + + repository.saveAll("conversation-id", List.of()); + + verify(transactionManager).commit(transactionStatus); + } + } diff --git a/memory-repositories/spring-ai-model-chat-memory-repository-jdbc/src/test/java/org/springframework/ai/chat/memory/repository/jdbc/JdbcChatMemoryRepositoryMysqlIT.java b/memory-repositories/spring-ai-model-chat-memory-repository-jdbc/src/test/java/org/springframework/ai/chat/memory/repository/jdbc/JdbcChatMemoryRepositoryMysqlIT.java index 482a335a1d..b49802ffdd 100644 --- a/memory-repositories/spring-ai-model-chat-memory-repository-jdbc/src/test/java/org/springframework/ai/chat/memory/repository/jdbc/JdbcChatMemoryRepositoryMysqlIT.java +++ b/memory-repositories/spring-ai-model-chat-memory-repository-jdbc/src/test/java/org/springframework/ai/chat/memory/repository/jdbc/JdbcChatMemoryRepositoryMysqlIT.java @@ -16,10 +16,26 @@ package org.springframework.ai.chat.memory.repository.jdbc; +import java.util.List; +import java.util.Objects; +import java.util.concurrent.BrokenBarrierException; +import java.util.concurrent.CyclicBarrier; +import java.util.concurrent.Executors; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.TimeoutException; + +import javax.sql.DataSource; + +import org.junit.jupiter.api.Test; + +import org.springframework.ai.chat.messages.UserMessage; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.jdbc.core.JdbcTemplate; import org.springframework.test.context.TestPropertySource; import org.springframework.test.context.jdbc.Sql; +import static org.assertj.core.api.Assertions.assertThat; + /** * Integration tests for {@link JdbcChatMemoryRepository} with MySQL. * @@ -34,4 +50,63 @@ @Sql(scripts = "classpath:org/springframework/ai/chat/memory/repository/jdbc/schema-mysql.sql") class JdbcChatMemoryRepositoryMysqlIT extends AbstractJdbcChatMemoryRepositoryIT { + @Test + void savesDifferentConversationsConcurrently() throws Exception { + this.jdbcTemplate.update("DELETE FROM SPRING_AI_CHAT_MEMORY"); + var dataSource = Objects.requireNonNull(this.jdbcTemplate.getDataSource()); + var conversationIds = List.of("conversation-1", "conversation-2", "conversation-3", "conversation-4"); + var repository = JdbcChatMemoryRepository.builder() + .jdbcTemplate(new DeleteBarrierJdbcTemplate(dataSource, conversationIds.size())) + .build(); + var executor = Executors.newFixedThreadPool(conversationIds.size()); + + try { + var saves = conversationIds.stream() + .map(conversationId -> executor + .submit(() -> repository.saveAll(conversationId, List.of(new UserMessage(conversationId))))) + .toList(); + for (var save : saves) { + save.get(10, TimeUnit.SECONDS); + } + } + finally { + executor.shutdownNow(); + } + + assertThat(conversationIds) + .allSatisfy(conversationId -> assertThat(repository.findByConversationId(conversationId)).hasSize(1)); + } + + private static final class DeleteBarrierJdbcTemplate extends JdbcTemplate { + + private final CyclicBarrier barrier; + + private DeleteBarrierJdbcTemplate(DataSource dataSource, int concurrentSaves) { + super(dataSource); + this.barrier = new CyclicBarrier(concurrentSaves); + } + + @Override + public int update(String sql, Object... args) { + int updatedRows = super.update(sql, args); + if (sql.startsWith("DELETE FROM SPRING_AI_CHAT_MEMORY")) { + awaitDeleteBarrier(); + } + return updatedRows; + } + + private void awaitDeleteBarrier() { + try { + this.barrier.await(10, TimeUnit.SECONDS); + } + catch (InterruptedException ex) { + Thread.currentThread().interrupt(); + throw new IllegalStateException("Interrupted while waiting for concurrent deletes", ex); + } + catch (BrokenBarrierException | TimeoutException ex) { + throw new IllegalStateException("Concurrent deletes did not reach the barrier", ex); + } + } + + } }