Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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;

Expand Down Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -19,17 +19,22 @@
import java.sql.Connection;
import java.sql.DatabaseMetaData;
import java.sql.SQLException;
import java.util.List;

import javax.sql.DataSource;

import org.junit.jupiter.api.Test;

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;

/**
Expand Down Expand Up @@ -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);
}

}
Original file line number Diff line number Diff line change
Expand Up @@ -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.
*
Expand All @@ -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);
}
}

}
}