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 @@ -26,6 +26,7 @@
import java.util.stream.Collectors;

import org.neo4j.driver.Session;
import org.neo4j.driver.SimpleQueryRunner;
import org.neo4j.driver.Transaction;
import org.neo4j.driver.TransactionContext;

Expand Down Expand Up @@ -124,12 +125,12 @@ else if (msgType.equals(MessageType.TOOL.getValue())) {

@Override
public void saveAll(String conversationId, List<Message> messages) {
// First delete existing messages for this conversation
deleteByConversationId(conversationId);

// Then add the new messages
// Replace the conversation in a single transaction. Deleting through a
// separate, already committed transaction loses the previous messages when
// the write that follows it fails.
try (Session s = this.config.getDriver().session()) {
s.executeWriteWithoutResult(tx -> {
deleteConversation(tx, conversationId);
for (Message m : messages) {
addMessageToTransaction(tx, conversationId, m);
}
Expand All @@ -139,6 +140,19 @@ public void saveAll(String conversationId, List<Message> messages) {

@Override
public void deleteByConversationId(String conversationId) {
try (Session s = this.config.getDriver().session()) {
try (Transaction t = s.beginTransaction()) {
deleteConversation(t, conversationId);
t.commit();
}
}
}

public Neo4jChatMemoryRepositoryConfig getConfig() {
return this.config;
}

private void deleteConversation(SimpleQueryRunner queryRunner, String conversationId) {
String deleteMessagesStatement = """
MATCH (s:$($sessionLabel) {id:$conversationId})-[r:HAS_MESSAGE]->(m:$($messageLabel))
OPTIONAL MATCH (m)-[:HAS_METADATA]->(metadata:$($metadataLabel))
Expand All @@ -158,17 +172,8 @@ OPTIONAL MATCH (m)-[:HAS_TOOL_CALL]->(tc:$($toolCallLabel))
this.config.getMetadataLabel(), "mediaLabel", this.config.getMediaLabel(), "toolResponseLabel",
this.config.getToolResponseLabel(), "toolCallLabel", this.config.getToolCallLabel());

try (Session s = this.config.getDriver().session()) {
try (Transaction t = s.beginTransaction()) {
t.run(deleteMessagesStatement, params);
t.run(deleteConversationStatement, params);
t.commit();
}
}
}

public Neo4jChatMemoryRepositoryConfig getConfig() {
return this.config;
queryRunner.run(deleteMessagesStatement, params);
queryRunner.run(deleteConversationStatement, params);
}

private Message buildToolMessage(org.neo4j.driver.Record record) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -429,6 +429,30 @@ void saveAssistantMessageWithOptionalMetadataFails() {
.hasMessageContaining("Unable to convert java.util.Optional");
}

@Test
void saveAllKeepsTheExistingConversationWhenAWriteFails() {
var conversationId = UUID.randomUUID().toString();

this.chatMemoryRepository.saveAll(conversationId,
List.of(new UserMessage("First message"), new AssistantMessage("Second message")));
assertThat(this.chatMemoryRepository.findByConversationId(conversationId)).hasSize(2);

// Optional-typed metadata cannot be serialized, so the replacement write
// fails on its second message. See
// saveAssistantMessageWithOptionalMetadataFails.
AssistantMessage unstorable = AssistantMessage.builder()
.content("Replacement that cannot be stored")
.properties(Map.of("refusal", Optional.of("I cannot answer that")))
.build();

assertThatThrownBy(() -> this.chatMemoryRepository.saveAll(conversationId,
List.of(new UserMessage("Replacement message"), unstorable)))
.isInstanceOf(org.neo4j.driver.exceptions.ClientException.class);

assertThat(this.chatMemoryRepository.findByConversationId(conversationId)).extracting(Message::getText)
.containsExactly("First message", "Second message");
}

private Message createMessageByType(String content, MessageType messageType) {
return switch (messageType) {
case ASSISTANT -> new AssistantMessage(content);
Expand Down