Skip to content
Merged
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 @@ -342,4 +342,183 @@ class ChatRepositoryCheckpointInstrumentedTest {
)
assertEquals(otherResponseGroupId, grownMessages.last().responseGroupId)
}
}

@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" })
}
}
15 changes: 11 additions & 4 deletions app/src/main/java/com/zhousl/aether/data/ChatRepository.kt
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
}
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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<ChatAgentMessageRefEntity>

@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")
Expand All @@ -66,45 +66,52 @@ 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<ChatMessageEntity>

@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<ChatMessageEntity>

@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<String>): List<ChatSessionMessageStatsEntity>

@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<String>): Flow<List<ChatSessionMessageStatsEntity>>

@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<ChatMessageSummaryEntity>

@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<List<ChatMessageSummaryEntity>>

@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<String>): List<ChatMessageSummaryEntity>
Expand Down Expand Up @@ -188,6 +195,60 @@ interface ChatHistoryDao {
@Upsert
suspend fun upsertWorkspaceFileRefs(refs: List<ChatWorkspaceFileRefEntity>)

/**
* 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<ChatMessageEntity>,
workspaceFileRefs: List<ChatWorkspaceFileRefEntity> = 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)

Expand All @@ -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)

Expand Down Expand Up @@ -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()

Expand Down
Original file line number Diff line number Diff line change
@@ -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<ChatMessageEntity>,
): List<ChatMessageEntity> {
if (messages.isEmpty()) return emptyList()
val seenIds = HashSet<String>(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<ChatMessageEntity>,
incoming: List<ChatMessageEntity>,
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 }
}
Loading
Loading