From 8b6f9ea86de92ab3819b1594dc4c2d102c6f4d73 Mon Sep 17 00:00:00 2001 From: Chumor <185651145+Chumor@users.noreply.github.com> Date: Sat, 19 Sep 2026 10:07:00 +0800 Subject: [PATCH] fix(history): preserve unchanged messages during sync --- ...hatRepositoryCheckpointInstrumentedTest.kt | 181 +++++++++++++++++- .../com/zhousl/aether/data/ChatRepository.kt | 15 +- .../aether/data/chatdb/ChatHistoryDao.kt | 85 +++++++- .../data/chatdb/ChatHistoryMessageSync.kt | 35 ++++ .../data/chatdb/SharedChatHistoryStore.kt | 41 ++-- .../data/chatdb/ChatHistoryMessageSyncTest.kt | 96 ++++++++++ 6 files changed, 420 insertions(+), 33 deletions(-) create mode 100644 shared/src/commonMain/kotlin/com/zhousl/aether/data/chatdb/ChatHistoryMessageSync.kt create mode 100644 shared/src/commonTest/kotlin/com/zhousl/aether/data/chatdb/ChatHistoryMessageSyncTest.kt diff --git a/app/src/androidTest/java/com/zhousl/aether/data/ChatRepositoryCheckpointInstrumentedTest.kt b/app/src/androidTest/java/com/zhousl/aether/data/ChatRepositoryCheckpointInstrumentedTest.kt index 7de6446f..04815ecd 100644 --- a/app/src/androidTest/java/com/zhousl/aether/data/ChatRepositoryCheckpointInstrumentedTest.kt +++ b/app/src/androidTest/java/com/zhousl/aether/data/ChatRepositoryCheckpointInstrumentedTest.kt @@ -342,4 +342,183 @@ class ChatRepositoryCheckpointInstrumentedTest { ) assertEquals(otherResponseGroupId, grownMessages.last().responseGroupId) } -} \ No newline at end of file + + @Test + fun snapshotAfterCompletedCheckpointKeepsOverflowAndClearsParkedRows() = runBlocking { + val sessionId = "session-checkpoint-snapshot" + val responseGroupId = "agent-group-target" + val before = ChatMessage( + id = "user-before", + author = MessageAuthor.User, + text = "before", + ) + val checkpoint = ChatMessage( + id = "agent-target", + author = MessageAuthor.Agent, + text = "partial", + isIncomplete = true, + responseGroupId = responseGroupId, + ) + val overflow = ChatMessage( + id = "user-after", + author = MessageAuthor.User, + text = "keep overflow", + ) + val session = ChatSession( + id = sessionId, + title = "Snapshot after checkpoint", + preview = "keep overflow", + messages = listOf(before, checkpoint, overflow), + ) + repository.updateChatState( + sessions = listOf(session), + currentSessionId = sessionId, + ) + repository.upsertAssistantResponseCheckpoints( + checkpoints = listOf( + AssistantResponseCheckpoint( + target = AssistantResponseCheckpointTarget( + sessionId = sessionId, + responseGroupId = responseGroupId, + ), + fromPosition = 1, + messages = listOf( + checkpoint.copy(id = "agent-target-complete", text = "complete", isIncomplete = false), + ), + ), + ), + ) + assertEquals(0, database.chatHistoryDao().getStoredParkedMessageCount(sessionId)) + + val restored = repository.getSessionWithMessages(sessionId)?.messages.orEmpty() + repository.updateChatState( + sessions = listOf( + session.copy( + preview = "complete", + messages = restored, + ), + ), + currentSessionId = sessionId, + ) + + assertEquals(0, database.chatHistoryDao().getStoredParkedMessageCount(sessionId)) + assertEquals( + listOf("user-before", "agent-target-complete", "user-after"), + repository.getSessionWithMessages(sessionId)?.messages.orEmpty().map { it.id }, + ) + assertEquals( + listOf(0, 1, 2), + database.chatHistoryDao().getMessageSummariesForSession(sessionId).map { it.position }, + ) + } + + @Test + fun snapshotCleansLeftoverParkedWithoutRewritingUnchangedActiveMessages() = runBlocking { + val sessionId = "session-parked-leftover" + val before = ChatMessage( + id = "user-before", + author = MessageAuthor.User, + text = "before", + ) + val agent = ChatMessage( + id = "agent-1", + author = MessageAuthor.Agent, + text = "answer", + ) + val overflow = ChatMessage( + id = "user-after", + author = MessageAuthor.User, + text = "overflow", + ) + repository.updateChatState( + sessions = listOf( + ChatSession( + id = sessionId, + title = "Parked leftover", + preview = "overflow", + messages = listOf(before, agent, overflow), + ), + ), + currentSessionId = sessionId, + ) + database.chatHistoryDao().parkMessagesFromPositionOutsideResponseGroup( + sessionId = sessionId, + responseGroupId = "missing-group", + fromPosition = 2, + toPosition = 3, + ) + assertTrue(database.chatHistoryDao().getStoredParkedMessageCount(sessionId) > 0) + + repository.updateChatState( + sessions = listOf( + ChatSession( + id = sessionId, + title = "Parked leftover", + preview = "answer", + messages = listOf(before, agent), + ), + ), + currentSessionId = sessionId, + ) + + assertEquals(0, database.chatHistoryDao().getStoredParkedMessageCount(sessionId)) + assertEquals( + listOf("user-before", "agent-1"), + repository.getSessionWithMessages(sessionId)?.messages.orEmpty().map { it.id }, + ) + } + + @Test + fun snapshotRemovesDeletedSuffixAgentRefsAndKeepsPrefixRefs() = runBlocking { + val sessionId = "session-agent-refs" + val user = ChatMessage( + id = "user-1", + author = MessageAuthor.User, + text = "hello", + ) + val agent = ChatMessage( + id = "agent-1", + author = MessageAuthor.Agent, + text = "answer", + ) + val followUp = ChatMessage( + id = "user-2", + author = MessageAuthor.User, + text = "more", + ) + repository.updateChatState( + sessions = listOf( + ChatSession( + id = sessionId, + title = "Refs", + preview = "more", + messages = listOf(user, agent, followUp), + ), + ), + currentSessionId = sessionId, + ) + repository.upsertAgentMessageRefs(sessionId, listOf("user-1"), listOf("entry-user")) + repository.upsertAgentMessageRefs(sessionId, listOf("agent-1"), listOf("entry-agent")) + repository.upsertAgentMessageRefs(sessionId, listOf("user-2"), listOf("entry-follow-up")) + + repository.updateChatState( + sessions = listOf( + ChatSession( + id = sessionId, + title = "Refs", + preview = "answer", + messages = listOf(user, agent), + ), + ), + currentSessionId = sessionId, + ) + + val remaining = database.chatHistoryDao().getAgentMessageRefs(sessionId) + assertEquals(setOf("user-1", "agent-1"), remaining.map { it.aetherMessageId }.toSet()) + assertEquals( + listOf("entry-user"), + remaining.filter { it.aetherMessageId == "user-1" }.map { it.piEntryId }, + ) + assertTrue(remaining.none { it.aetherMessageId == "user-2" }) + } +} diff --git a/app/src/main/java/com/zhousl/aether/data/ChatRepository.kt b/app/src/main/java/com/zhousl/aether/data/ChatRepository.kt index c05187de..59a91746 100644 --- a/app/src/main/java/com/zhousl/aether/data/ChatRepository.kt +++ b/app/src/main/java/com/zhousl/aether/data/ChatRepository.kt @@ -659,11 +659,18 @@ class ChatRepository( syncedMessages.isEmpty() && existingMessageCount > 0 if (!isMetadataOnlySnapshot) { - chatHistoryDao.deleteWorkspaceFileRefsForSession(session.id) - chatHistoryDao.deleteMessagesForSession(session.id) - chatHistoryDao.upsertMessagesChunked( + val messageEntities = syncedMessages.mapIndexed { index, message -> + ChatMessageEntityMapper.toEntity( + sessionId = session.id, + position = index, + message = message, + ) + } + val workspaceFileRefs = syncedMessages.toWorkspaceFileRefs(session.id) + chatHistoryDao.syncMessagesForSession( sessionId = session.id, - messages = syncedMessages, + messages = messageEntities, + workspaceFileRefs = workspaceFileRefs, ) } } diff --git a/shared/src/commonMain/kotlin/com/zhousl/aether/data/chatdb/ChatHistoryDao.kt b/shared/src/commonMain/kotlin/com/zhousl/aether/data/chatdb/ChatHistoryDao.kt index 85673d24..d4f1259a 100644 --- a/shared/src/commonMain/kotlin/com/zhousl/aether/data/chatdb/ChatHistoryDao.kt +++ b/shared/src/commonMain/kotlin/com/zhousl/aether/data/chatdb/ChatHistoryDao.kt @@ -56,7 +56,7 @@ interface ChatHistoryDao { @Query("SELECT * FROM chat_agent_message_refs WHERE chatSessionId = :sessionId AND aetherMessageId = :messageId ORDER BY ordinal") suspend fun getAgentMessageRefs(sessionId: String, messageId: String): List - @Query("SELECT COUNT(*) FROM chat_messages WHERE sessionId = :sessionId") + @Query("SELECT COUNT(*) FROM chat_messages WHERE sessionId = :sessionId AND position >= 0") suspend fun getMessageCountForSession(sessionId: String): Int @Query("SELECT COUNT(*) FROM chat_messages WHERE sessionId = :sessionId AND responseGroupId = :responseGroupId AND position >= :fromPosition") @@ -66,13 +66,20 @@ interface ChatHistoryDao { fromPosition: Int, ): Int - @Query("SELECT * FROM chat_messages WHERE sessionId = :sessionId ORDER BY position ASC") + @Query("SELECT * FROM chat_messages WHERE sessionId = :sessionId AND position >= 0 ORDER BY position ASC") suspend fun getMessagesForSession(sessionId: String): List + @Query("SELECT * FROM chat_messages WHERE sessionId = :sessionId AND position >= :fromPosition AND position < :toPosition ORDER BY position ASC") + suspend fun getMessagesForSessionInPositionRange( + sessionId: String, + fromPosition: Int, + toPosition: Int, + ): List + @Query(""" SELECT sessionId, COUNT(*) AS messageCount, MAX(COALESCE(createdAtMillis, 0)) AS lastMessageAtMillis FROM chat_messages - WHERE sessionId IN (:sessionIds) + WHERE sessionId IN (:sessionIds) AND position >= 0 GROUP BY sessionId """) suspend fun getMessageStatsForSessions(sessionIds: List): List @@ -80,7 +87,7 @@ interface ChatHistoryDao { @Query(""" SELECT sessionId, COUNT(*) AS messageCount, MAX(COALESCE(createdAtMillis, 0)) AS lastMessageAtMillis FROM chat_messages - WHERE sessionId IN (:sessionIds) + WHERE sessionId IN (:sessionIds) AND position >= 0 GROUP BY sessionId """) fun observeMessageStatsForSessions(sessionIds: List): Flow> @@ -88,7 +95,7 @@ interface ChatHistoryDao { @Query(""" SELECT sessionId, id, position, author, text, createdAtMillis, responseGroupId, displayKind, messageSchemaVersion, length(messageJson) AS messageJsonLength, isIncomplete FROM chat_messages - WHERE hasUsageStatistics = 1 + WHERE hasUsageStatistics = 1 AND position >= 0 ORDER BY sessionId ASC, position ASC """) suspend fun getUsageStatisticsMessageSummaries(): List @@ -96,7 +103,7 @@ interface ChatHistoryDao { @Query(""" SELECT sessionId, id, position, author, text, createdAtMillis, responseGroupId, displayKind, messageSchemaVersion, length(messageJson) AS messageJsonLength, isIncomplete FROM chat_messages - WHERE sessionId = :sessionId + WHERE sessionId = :sessionId AND position >= 0 ORDER BY position ASC """) fun observeMessageSummariesForSession(sessionId: String): Flow> @@ -104,7 +111,7 @@ interface ChatHistoryDao { @Query(""" SELECT sessionId, id, position, author, text, createdAtMillis, responseGroupId, displayKind, messageSchemaVersion, length(messageJson) AS messageJsonLength, isIncomplete FROM chat_messages - WHERE sessionId IN (:sessionIds) + WHERE sessionId IN (:sessionIds) AND position >= 0 ORDER BY sessionId ASC, position ASC """) suspend fun getMessageSummariesForSessions(sessionIds: List): List @@ -188,6 +195,60 @@ interface ChatHistoryDao { @Upsert suspend fun upsertWorkspaceFileRefs(refs: List) + /** + * Synchronizes the active history from its first changed message. + * + * Compare before clearing parked checkpoint rows so unchanged history and its references remain intact. + */ + @Transaction + suspend fun syncMessagesForSession( + sessionId: String, + messages: List, + workspaceFileRefs: List = emptyList(), + ) { + val canonical = canonicalActiveChatMessages(sessionId, messages) + val existingCount = getMessageCountForSession(sessionId) + var firstChangedPosition: Int? = null + for (startPosition in canonical.indices step ChatHistoryMessageSyncChunkSize) { + val endPosition = minOf(startPosition + ChatHistoryMessageSyncChunkSize, canonical.size) + val existingBatch = getMessagesForSessionInPositionRange(sessionId, startPosition, endPosition) + val incomingBatch = canonical.subList(startPosition, endPosition) + firstChangedPosition = firstChangedMessagePosition(existingBatch, incomingBatch, startPosition) + if (firstChangedPosition != null) break + } + val syncFromPosition = firstChangedPosition ?: canonical.size.takeIf { existingCount > canonical.size } + if (syncFromPosition == null) { + if (getStoredParkedMessageCount(sessionId) > 0) { + deleteParkedMessagesForSession(sessionId) + deleteOrphanedAgentMessageRefs(sessionId) + deleteWorkspaceFileRefsForInactiveMessages(sessionId) + } + return + } + + val changedMessages = canonical.subList(syncFromPosition, canonical.size) + deleteWorkspaceFileRefsFromPosition(sessionId, syncFromPosition) + deleteMessagesFromPosition(sessionId, syncFromPosition) + deleteParkedMessagesForSession(sessionId) + changedMessages + .chunked(ChatHistoryMessageSyncChunkSize) + .forEach { batch -> upsertMessages(batch) } + deleteOrphanedAgentMessageRefs(sessionId) + deleteWorkspaceFileRefsForInactiveMessages(sessionId) + val changedMessageIds = changedMessages.asSequence().map(ChatMessageEntity::id).toSet() + workspaceFileRefs + .asSequence() + .filter { it.messageId in changedMessageIds } + .chunked(ChatHistoryWorkspaceRefSyncChunkSize) + .forEach { batch -> upsertWorkspaceFileRefs(batch.toList()) } + } + + @Query("SELECT COUNT(*) FROM chat_messages WHERE sessionId = :sessionId AND position < 0") + suspend fun getStoredParkedMessageCount(sessionId: String): Int + + @Query("DELETE FROM chat_messages WHERE sessionId = :sessionId AND position < 0") + suspend fun deleteParkedMessagesForSession(sessionId: String) + @Query("DELETE FROM chat_workspace_file_refs WHERE sessionId = :sessionId AND messageId = :messageId") suspend fun deleteWorkspaceFileRefsForMessage(sessionId: String, messageId: String) @@ -201,6 +262,9 @@ interface ChatHistoryDao { @Query("DELETE FROM chat_workspace_file_refs WHERE sessionId = :sessionId AND messageId IN (SELECT id FROM chat_messages WHERE sessionId = :sessionId AND position >= :fromPosition)") suspend fun deleteWorkspaceFileRefsFromPosition(sessionId: String, fromPosition: Int) + @Query("DELETE FROM chat_workspace_file_refs WHERE sessionId = :sessionId AND messageId NOT IN (SELECT id FROM chat_messages WHERE sessionId = :sessionId AND position >= 0)") + suspend fun deleteWorkspaceFileRefsForInactiveMessages(sessionId: String) + @Query("DELETE FROM chat_workspace_file_refs WHERE sessionId = :sessionId") suspend fun deleteWorkspaceFileRefsForSession(sessionId: String) @@ -279,6 +343,13 @@ interface ChatHistoryDao { @Query("DELETE FROM chat_agent_message_refs WHERE chatSessionId = :sessionId") suspend fun deleteAgentMessageRefs(sessionId: String) + @Query(""" + DELETE FROM chat_agent_message_refs + WHERE chatSessionId = :sessionId + AND aetherMessageId NOT IN (SELECT id FROM chat_messages WHERE sessionId = :sessionId) + """) + suspend fun deleteOrphanedAgentMessageRefs(sessionId: String) + @Query("DELETE FROM chat_agent_sessions") suspend fun deleteAllAgentSessions() diff --git a/shared/src/commonMain/kotlin/com/zhousl/aether/data/chatdb/ChatHistoryMessageSync.kt b/shared/src/commonMain/kotlin/com/zhousl/aether/data/chatdb/ChatHistoryMessageSync.kt new file mode 100644 index 00000000..51bd9306 --- /dev/null +++ b/shared/src/commonMain/kotlin/com/zhousl/aether/data/chatdb/ChatHistoryMessageSync.kt @@ -0,0 +1,35 @@ +package com.zhousl.aether.data.chatdb + +internal const val ChatHistoryMessageSyncChunkSize = 8 +internal const val ChatHistoryWorkspaceRefSyncChunkSize = 32 + +fun canonicalActiveChatMessages( + sessionId: String, + messages: List, +): List { + if (messages.isEmpty()) return emptyList() + val seenIds = HashSet(messages.size) + return messages.mapIndexed { index, message -> + require(message.sessionId == sessionId) { + "Message ${message.id} belongs to ${message.sessionId}, not $sessionId." + } + require(message.id.isNotBlank()) { "Chat message IDs cannot be blank." } + require(seenIds.add(message.id)) { + "Duplicate chat message id ${message.id} in session $sessionId." + } + if (message.position == index) message else message.copy(position = index) + } +} + +/** Finds the first differing position in a dense incoming batch starting at [startPosition]. */ +fun firstChangedMessagePosition( + existing: List, + incoming: List, + startPosition: Int, +): Int? { + require(startPosition >= 0) { "startPosition must be non-negative." } + return incoming.indices.firstOrNull { index -> + val stored = existing.getOrNull(index) + stored?.position != startPosition + index || stored != incoming[index] + }?.let { startPosition + it } +} diff --git a/shared/src/commonMain/kotlin/com/zhousl/aether/data/chatdb/SharedChatHistoryStore.kt b/shared/src/commonMain/kotlin/com/zhousl/aether/data/chatdb/SharedChatHistoryStore.kt index c225d1be..7ce895d4 100644 --- a/shared/src/commonMain/kotlin/com/zhousl/aether/data/chatdb/SharedChatHistoryStore.kt +++ b/shared/src/commonMain/kotlin/com/zhousl/aether/data/chatdb/SharedChatHistoryStore.kt @@ -470,28 +470,27 @@ class SharedChatHistoryStore( sortOrder = -platformSortOrder(), ) ) - dao.deleteMessagesForSession(sessionId) - dao.deleteWorkspaceFileRefsForSession(sessionId) - dao.upsertMessages( - messages.mapIndexed { index, message -> - val json = message.toJsonObject() - ChatMessageEntity( - sessionId = sessionId, - id = message.id, - position = index, - messageJson = json.toString(), - author = if (message.fromUser) "User" else "Agent", - text = message.text, - createdAtMillis = message.createdAtMillis, - responseGroupId = message.responseGroupId.ifBlank { null }, - displayKind = message.displayKind.name, - hasUsageStatistics = message.usage != null, - isIncomplete = false, - ) - } + val messageEntities = messages.mapIndexed { index, message -> + val json = message.toJsonObject() + ChatMessageEntity( + sessionId = sessionId, + id = message.id, + position = index, + messageJson = json.toString(), + author = if (message.fromUser) "User" else "Agent", + text = message.text, + createdAtMillis = message.createdAtMillis, + responseGroupId = message.responseGroupId.ifBlank { null }, + displayKind = message.displayKind.name, + hasUsageStatistics = message.usage != null, + isIncomplete = false, + ) + } + dao.syncMessagesForSession( + sessionId = sessionId, + messages = messageEntities, + workspaceFileRefs = messages.toWorkspaceFileRefs(sessionId), ) - val workspaceFileRefs = messages.toWorkspaceFileRefs(sessionId) - if (workspaceFileRefs.isNotEmpty()) dao.upsertWorkspaceFileRefs(workspaceFileRefs) dao.upsertMeta( ChatStateMetaEntity( currentSessionId = sessionId, diff --git a/shared/src/commonTest/kotlin/com/zhousl/aether/data/chatdb/ChatHistoryMessageSyncTest.kt b/shared/src/commonTest/kotlin/com/zhousl/aether/data/chatdb/ChatHistoryMessageSyncTest.kt new file mode 100644 index 00000000..19b52bc8 --- /dev/null +++ b/shared/src/commonTest/kotlin/com/zhousl/aether/data/chatdb/ChatHistoryMessageSyncTest.kt @@ -0,0 +1,96 @@ +package com.zhousl.aether.data.chatdb + +import kotlin.test.Test +import kotlin.test.assertEquals +import kotlin.test.assertFailsWith +import kotlin.test.assertNull +import kotlin.test.assertTrue + +class ChatHistoryMessageSyncTest { + @Test + fun duplicateMessageIdsAreRejectedBeforeRoomWrites() { + val error = assertFailsWith { + canonicalActiveChatMessages( + sessionId = "session", + messages = listOf( + testMessage("session", "user-1", 0), + testMessage("session", "user-1", 1), + ), + ) + } + assertTrue(error.message.orEmpty().contains("user-1")) + } + + @Test + fun positionsAreNormalizedToADenseActivePrefix() { + val canonical = canonicalActiveChatMessages( + sessionId = "session", + messages = listOf( + testMessage("session", "user-1", 4), + testMessage("session", "agent-1", 9), + ), + ) + assertEquals(listOf(0, 1), canonical.map(ChatMessageEntity::position)) + assertEquals(listOf("user-1", "agent-1"), canonical.map(ChatMessageEntity::id)) + } + + @Test + fun findsFirstChangedMessageInBatchAndKeepsEarlierRowsUntouched() { + val existing = listOf( + testMessage("session", "agent-1", 8, "partial"), + testMessage("session", "user-2", 9, "follow-up"), + ) + val incoming = listOf( + testMessage("session", "agent-1", 8, "complete"), + testMessage("session", "user-2", 9, "follow-up"), + ) + + assertEquals(8, firstChangedMessagePosition(existing, incoming, startPosition = 8)) + assertNull(firstChangedMessagePosition(existing.drop(1), incoming.drop(1), startPosition = 9)) + } + + @Test + fun missingPositionIsTreatedAsFirstChange() { + val existing = listOf(testMessage("session", "user-2", 2)) + val incoming = listOf( + testMessage("session", "user-1", 1), + testMessage("session", "user-2", 2), + ) + + assertEquals(1, firstChangedMessagePosition(existing, incoming, startPosition = 1)) + } + + @Test + fun completedTurnOnlyRewritesTheChangedSuffix() { + val checkpoint = listOf( + testMessage("session", "user-1", 0, "before"), + testMessage("session", "agent-partial", 1, "partial"), + testMessage("session", "user-after", 2, "keep"), + ) + val completed = canonicalActiveChatMessages( + sessionId = "session", + messages = listOf( + testMessage("session", "user-1", 0, "before"), + testMessage("session", "agent-complete", 1, "complete"), + testMessage("session", "user-after", 2, "keep"), + ), + ) + + assertNull(firstChangedMessagePosition(checkpoint.take(1), completed.take(1), startPosition = 0)) + assertEquals(1, firstChangedMessagePosition(checkpoint.drop(1), completed.drop(1), startPosition = 1)) + } +} + +private fun testMessage( + sessionId: String, + id: String, + position: Int, + text: String = id, +): ChatMessageEntity = ChatMessageEntity( + sessionId = sessionId, + id = id, + position = position, + messageJson = """{"id":"$id","text":"$text"}""", + author = "User", + text = text, +)