Fixed tool calling history
This commit is contained in:
1 parent
dd43033888
commit
524ed8ce53
11 files changed
+471
-55
No files matched your search
@@ -6,6 +6,7 @@ import iris.db.IrisDatabase
|
||||
|
||||
actual fun appDataDir(): String = AndroidEnv.context.filesDir.absolutePath
|
||||
|
||||
// The driver creates the schema in the open-helper callback (v1: no
|
||||
// migrations yet; add .sqm files + `Schema.migrate` in onUpgrade later).
|
||||
// The schema-aware driver creates the schema on first run and applies
|
||||
// pending .sqm migrations on existing DBs (e.g. v1 -> v2: the `tool`
|
||||
// table), stamping PRAGMA user_version along the way.
|
||||
actual fun createCacheDriver(): SqlDriver = AndroidSqliteDriver(IrisDatabase.Schema, AndroidEnv.context, "iris_cache.db")
|
||||
@@ -16,9 +16,9 @@ import iris.protocol.IrisJson
|
||||
* the controller) so every reconciled frame survives a process death.
|
||||
* - [metaGet] / [metaPut] hold small UI state (last-viewed lane).
|
||||
*
|
||||
* Only [MessageItem]s are persisted — tool cards, live streaming state and
|
||||
* [MessageItem]s and [ToolItem]s are persisted; live streaming state and
|
||||
* local system notices are ephemeral. Rows are JSON payloads keyed by
|
||||
* (lane, id), so the schema does not drift with [MessageItem] fields.
|
||||
* (lane, id), so the schema does not drift with the model fields.
|
||||
*
|
||||
* All access is synchronized: the debounced save collectors run on the
|
||||
* controller scope while [dispose] may flush from the UI thread.
|
||||
@@ -32,27 +32,70 @@ class ChatDb(
|
||||
|
||||
// ── messages ──────────────────────────────────────────────────────────
|
||||
|
||||
/** All persisted lanes (lane key -> messages ordered by ts). */
|
||||
fun loadLanes(): Map<String, List<MessageItem>> =
|
||||
/** All persisted lanes (lane key -> items: messages ordered by ts with
|
||||
* tool cards interleaved at their anchored position). */
|
||||
fun loadLanes(): Map<String, List<ChatItem>> =
|
||||
synchronized(lock) {
|
||||
val lanes = linkedMapOf<String, MutableList<MessageItem>>()
|
||||
val messages = linkedMapOf<String, MutableList<MessageItem>>()
|
||||
for (row in db.cacheQueries.allMessages().executeAsList()) {
|
||||
val item = decodeMessage(row.payload) ?: continue
|
||||
lanes.getOrPut(row.lane) { mutableListOf() }.add(item)
|
||||
messages.getOrPut(row.lane) { mutableListOf() }.add(item)
|
||||
}
|
||||
val tools = linkedMapOf<String, MutableList<ToolItem>>()
|
||||
for (row in db.cacheQueries.allTools().executeAsList()) {
|
||||
val item = decodeTool(row.payload) ?: continue
|
||||
tools.getOrPut(row.lane) { mutableListOf() }.add(item)
|
||||
}
|
||||
(messages.keys + tools.keys).distinct().associateWith { lane ->
|
||||
val items: MutableList<ChatItem> = messages[lane].orEmpty().toMutableList()
|
||||
// Insert each tool card after its anchor message (the message
|
||||
// it followed live). Cards sharing an anchor keep their seq
|
||||
// order; a card whose anchor is gone (deleted message) falls
|
||||
// to the end of the lane. The lastPos cache assumes anchors
|
||||
// are monotonically non-decreasing per lane (true for the
|
||||
// onToolStart anchor rule: last non-streaming message) — a
|
||||
// tool anchored to an EARLIER message processed after a
|
||||
// later-anchored one would be misplaced.
|
||||
val lastPos = mutableMapOf<String, Int>()
|
||||
for (tool in tools[lane].orEmpty()) {
|
||||
val anchor = tool.anchorId
|
||||
val pos =
|
||||
if (anchor == null) {
|
||||
-1
|
||||
} else {
|
||||
lastPos.getOrPut(anchor) { items.indexOfFirst { it.id == anchor } }
|
||||
}
|
||||
if (pos >= 0) {
|
||||
items.add(pos + 1, tool)
|
||||
if (anchor != null) lastPos[anchor] = pos + 1
|
||||
} else {
|
||||
items.add(tool)
|
||||
}
|
||||
}
|
||||
items
|
||||
}
|
||||
lanes.mapValues { it.value.toList() }
|
||||
}
|
||||
|
||||
/** Replace the whole message cache with [lanes] (atomic snapshot). Tool
|
||||
* cards and local system notices are skipped (ephemeral). */
|
||||
/** Replace the whole message + tool cache with [lanes] (atomic snapshot).
|
||||
* Local system notices are skipped (ephemeral). */
|
||||
fun saveLanes(lanes: Map<String, List<ChatItem>>) {
|
||||
synchronized(lock) {
|
||||
db.transaction {
|
||||
db.cacheQueries.clearMessages()
|
||||
db.cacheQueries.clearTools()
|
||||
for ((lane, items) in lanes) {
|
||||
var toolSeq = 0
|
||||
for (item in items) {
|
||||
if (item is MessageItem && !item.isSystem) {
|
||||
db.cacheQueries.upsertMessage(lane, item.id, item.ts, json.encodeToString(item))
|
||||
when (item) {
|
||||
is MessageItem -> {
|
||||
if (!item.isSystem) {
|
||||
db.cacheQueries.upsertMessage(lane, item.id, item.ts, json.encodeToString(item))
|
||||
}
|
||||
}
|
||||
|
||||
is ToolItem -> {
|
||||
db.cacheQueries.upsertTool(lane, item.id, toolSeq++.toLong(), json.encodeToString(item))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -109,6 +152,7 @@ class ChatDb(
|
||||
synchronized(lock) {
|
||||
db.transaction {
|
||||
db.cacheQueries.clearMessages()
|
||||
db.cacheQueries.clearTools()
|
||||
db.cacheQueries.clearChannels()
|
||||
db.cacheQueries.clearMeta()
|
||||
}
|
||||
@@ -122,6 +166,13 @@ class ChatDb(
|
||||
null
|
||||
}
|
||||
|
||||
private fun decodeTool(payload: String): ToolItem? =
|
||||
try {
|
||||
json.decodeFromString<ToolItem>(payload).sanitizeForRestore()
|
||||
} catch (_: Exception) {
|
||||
null
|
||||
}
|
||||
|
||||
/** A restored message is never mid-flight: a streaming bubble is
|
||||
* finalized (the sync replay finalizes it for real on reconnect) and a
|
||||
* pending send becomes failed (tap to retry) — the gateway never
|
||||
@@ -132,4 +183,9 @@ class ChatDb(
|
||||
pending = false,
|
||||
status = if (status == MsgStatus.Pending) MsgStatus.Failed else status,
|
||||
)
|
||||
|
||||
/** A restored tool card is never mid-flight: an open card (the process
|
||||
* died before tool.end) is closed as interrupted, mirroring
|
||||
* [ChatStore.finalizeInterrupted]. */
|
||||
private fun ToolItem.sanitizeForRestore(): ToolItem = if (done) this else copy(done = true, ok = false)
|
||||
}
|
||||
@@ -90,7 +90,13 @@ data class MediaItem(
|
||||
val localPath: String? = null,
|
||||
)
|
||||
|
||||
/** A structured tool-activity card (spinner until [done]). */
|
||||
/** A structured tool-activity card (spinner until [done]).
|
||||
* [anchorId] is the id of the message this card follows in the lane (the
|
||||
* last non-streaming message when the tool started) — persisted with the
|
||||
* card so a restart restores it in its correct position (user message →
|
||||
* tool card → answer) instead of dropping it or appending it at the end.
|
||||
* [Serializable]: persisted as a JSON payload in the local cache (ChatDb). */
|
||||
@Serializable
|
||||
data class ToolItem(
|
||||
override val id: String,
|
||||
val index: Int,
|
||||
@@ -102,6 +108,7 @@ data class ToolItem(
|
||||
val ok: Boolean = true,
|
||||
val duration: Double? = null,
|
||||
val outputPreview: String? = null,
|
||||
val anchorId: String? = null,
|
||||
) : ChatItem
|
||||
|
||||
class ChatStore {
|
||||
@@ -114,7 +121,6 @@ class ChatStore {
|
||||
val currentLane: StateFlow<String> = _currentLane.asStateFlow()
|
||||
|
||||
private var localSeq = 0
|
||||
private var toolSeq = 0
|
||||
|
||||
/** When false, `message.start`/`message.update` frames are ignored and each
|
||||
* reply materializes as a single final message on `message.stop`
|
||||
@@ -169,8 +175,11 @@ class ChatStore {
|
||||
lane: String,
|
||||
media: List<MediaItem> = emptyList(),
|
||||
): String {
|
||||
localSeq++
|
||||
val id = "local_$localSeq"
|
||||
// Process-unique id: the in-memory seq resets on every ChatStore
|
||||
// creation, and failed sends are persisted — a restart would
|
||||
// otherwise re-mint local_1 and collide with the restored bubble
|
||||
// (duplicate list key).
|
||||
val id = randomId("local")
|
||||
updateLane(lane) {
|
||||
it + MessageItem(id = id, role = ROLE_USER, text = text, ts = 0, pending = true, status = MsgStatus.Pending, media = media)
|
||||
}
|
||||
@@ -418,9 +427,17 @@ class ChatStore {
|
||||
frame: Frame,
|
||||
) {
|
||||
val p = frame.payloadAs<ToolStartPayload>() ?: return
|
||||
toolSeq++
|
||||
val id = "tool_$toolSeq"
|
||||
// Process-unique id: the in-memory seq resets on every ChatStore
|
||||
// creation, and cards are persisted — a restart would otherwise
|
||||
// re-mint tool_1 and collide with the restored card (duplicate list
|
||||
// key, upsert overwrite). onToolProgress/onToolEnd match by index,
|
||||
// so non-sequential ids are safe.
|
||||
val id = randomId("tool")
|
||||
updateLane(lane) { list ->
|
||||
// Anchor the card to the message it follows: the last non-streaming
|
||||
// message (a live streaming bubble is the answer that arrives AFTER
|
||||
// the tool, so it is skipped). Restored with the card on restart.
|
||||
val anchorId = list.lastOrNull { it is MessageItem && !it.streaming }?.id
|
||||
list +
|
||||
ToolItem(
|
||||
id = id,
|
||||
@@ -428,6 +445,7 @@ class ChatStore {
|
||||
name = p.name,
|
||||
preview = p.preview,
|
||||
args = p.args,
|
||||
anchorId = anchorId,
|
||||
)
|
||||
}
|
||||
}
|
||||
@@ -682,9 +700,11 @@ class ChatStore {
|
||||
* Load a history page into [lane] (oldest → newest). The history is the
|
||||
* authoritative full list of final messages for the lane; it replaces the
|
||||
* lane's final messages and preserves non-final items (tool cards, live
|
||||
* streaming bubbles, commentary) that are not part of the history. Used to
|
||||
* restore the view on first open of a chat / after a process death, where
|
||||
* the in-memory store is empty and the `sync` delta does not cover older
|
||||
* streaming bubbles, commentary) that are not part of the history. Tool
|
||||
* cards sort at their anchor message's position, so a history refresh
|
||||
* keeps them between the user message and the answer. Used to restore the
|
||||
* view on first open of a chat / after a process death, where the
|
||||
* in-memory store is empty and the `sync` delta does not cover older
|
||||
* messages.
|
||||
*/
|
||||
fun loadHistory(
|
||||
@@ -697,8 +717,28 @@ class ChatStore {
|
||||
list.filter { item ->
|
||||
item !is MessageItem || item.id !in historyIds
|
||||
}
|
||||
// ts of every item in the current lane: ts-less items (commentary,
|
||||
// tool cards) inherit the ts of the item before them, so a tool
|
||||
// card anchored to a commentary still sorts at the right place.
|
||||
val tsOf = mutableMapOf<String, Long>()
|
||||
var lastTs = 0L
|
||||
for (item in list) {
|
||||
val t = (item as? MessageItem)?.ts?.takeIf { it > 0 } ?: lastTs
|
||||
if (t > 0) lastTs = t
|
||||
tsOf[item.id] = t
|
||||
}
|
||||
// History is authoritative for the ts of its messages.
|
||||
for (m in messages) {
|
||||
if (m.ts > 0) tsOf[m.id] = m.ts
|
||||
}
|
||||
(messages + preserved).sortedBy { item ->
|
||||
(item as? MessageItem)?.ts?.takeIf { it > 0 } ?: Long.MAX_VALUE
|
||||
when (item) {
|
||||
is MessageItem -> item.ts.takeIf { it > 0 } ?: Long.MAX_VALUE
|
||||
|
||||
// A resolved ts of 0 means the anchor itself is ts-less
|
||||
// (lane start) — sort with it (end) instead of to the top.
|
||||
is ToolItem -> tsOf[item.anchorId]?.takeIf { it > 0 } ?: Long.MAX_VALUE
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -710,7 +750,7 @@ class ChatStore {
|
||||
* delta and the `history` refresh reconcile the cache on connect.
|
||||
* No-op when [lanes] is empty (first launch).
|
||||
*/
|
||||
fun loadFromCache(lanes: Map<String, List<MessageItem>>) {
|
||||
fun loadFromCache(lanes: Map<String, List<ChatItem>>) {
|
||||
if (lanes.isEmpty()) return
|
||||
_lanes.value = lanes
|
||||
}
|
||||
|
||||
@@ -465,13 +465,24 @@ class IrisController(
|
||||
client.events.collect { frame ->
|
||||
try {
|
||||
when (frame.type) {
|
||||
TYPE_TOOL_START,
|
||||
TYPE_TOOL_PROGRESS,
|
||||
TYPE_TOOL_END,
|
||||
-> {
|
||||
// Tool cards are restored from the local cache
|
||||
// (anchored to their message); they are not part
|
||||
// of `history`. A sync replay (frame carries a
|
||||
// cursor) would create duplicate cards appended
|
||||
// AFTER the lane's restored messages — drop
|
||||
// them. Live tool frames (no cursor) flow through
|
||||
// as usual.
|
||||
if (frame.cursor == null) chat.onFrame(frame)
|
||||
}
|
||||
|
||||
TYPE_MESSAGE,
|
||||
TYPE_MESSAGE_START,
|
||||
TYPE_MESSAGE_UPDATE,
|
||||
TYPE_MESSAGE_STOP,
|
||||
TYPE_TOOL_START,
|
||||
TYPE_TOOL_PROGRESS,
|
||||
TYPE_TOOL_END,
|
||||
TYPE_COMMENTARY,
|
||||
TYPE_MEDIA_OFFER,
|
||||
-> {
|
||||
|
||||
@@ -11,6 +11,14 @@ CREATE TABLE message (
|
||||
PRIMARY KEY (lane, id)
|
||||
);
|
||||
|
||||
CREATE TABLE tool (
|
||||
lane TEXT NOT NULL, -- lane key: chatId or chatId::threadId
|
||||
id TEXT NOT NULL, -- local tool card id (tool_N)
|
||||
seq INTEGER NOT NULL, -- card order within the lane (lane position)
|
||||
payload TEXT NOT NULL, -- serialized ToolItem (carries its anchor_id)
|
||||
PRIMARY KEY (lane, id)
|
||||
);
|
||||
|
||||
CREATE TABLE channel (
|
||||
chat_id TEXT NOT NULL PRIMARY KEY,
|
||||
payload TEXT NOT NULL -- serialized ChannelInfo
|
||||
@@ -33,6 +41,18 @@ VALUES (?, ?, ?, ?);
|
||||
clearMessages:
|
||||
DELETE FROM message;
|
||||
|
||||
allTools:
|
||||
SELECT lane, seq, payload
|
||||
FROM tool
|
||||
ORDER BY lane, seq;
|
||||
|
||||
upsertTool:
|
||||
INSERT OR REPLACE INTO tool (lane, id, seq, payload)
|
||||
VALUES (?, ?, ?, ?);
|
||||
|
||||
clearTools:
|
||||
DELETE FROM tool;
|
||||
|
||||
allChannels:
|
||||
SELECT chat_id, payload
|
||||
FROM channel;
|
||||
|
||||
@@ -0,0 +1,9 @@
|
||||
-- v1 -> v2: persist tool cards (M-cache: tool cards survive a restart,
|
||||
-- anchored to the message they follow — see ChatDb.loadLanes).
|
||||
CREATE TABLE tool (
|
||||
lane TEXT NOT NULL,
|
||||
id TEXT NOT NULL,
|
||||
seq INTEGER NOT NULL,
|
||||
payload TEXT NOT NULL,
|
||||
PRIMARY KEY (lane, id)
|
||||
);
|
||||
@@ -1,5 +1,9 @@
|
||||
package iris.data
|
||||
|
||||
import iris.protocol.Frame
|
||||
import iris.protocol.IrisJson
|
||||
import iris.protocol.TYPE_TOOL_START
|
||||
import iris.protocol.ToolStartPayload
|
||||
import kotlin.test.Test
|
||||
import kotlin.test.assertEquals
|
||||
|
||||
@@ -10,13 +14,14 @@ class ChatStoreCacheTest {
|
||||
store.loadFromCache(
|
||||
mapOf(
|
||||
"android:default" to
|
||||
listOf(
|
||||
listOf<ChatItem>(
|
||||
MessageItem(id = "m1", role = "user", text = "hi", ts = 1),
|
||||
ToolItem(id = "tool_1", index = 0, name = "bash", done = true, anchorId = "m1"),
|
||||
MessageItem(id = "m2", role = "assistant", text = "hello", ts = 2),
|
||||
),
|
||||
),
|
||||
)
|
||||
assertEquals(listOf("m1", "m2"), store.lanes.value["android:default"]!!.map { it.id })
|
||||
assertEquals(listOf("m1", "tool_1", "m2"), store.lanes.value["android:default"]!!.map { it.id })
|
||||
}
|
||||
|
||||
@Test
|
||||
@@ -26,4 +31,83 @@ class ChatStoreCacheTest {
|
||||
store.loadFromCache(emptyMap())
|
||||
assertEquals(1, store.lanes.value["android:default"]!!.size)
|
||||
}
|
||||
|
||||
@Test
|
||||
fun loadHistoryKeepsToolCardsAtAnchor() {
|
||||
val store = ChatStore()
|
||||
// Lane as restored from the cache: user message, tool card anchored
|
||||
// to it, and the final answer.
|
||||
store.loadFromCache(
|
||||
mapOf(
|
||||
"android:default" to
|
||||
listOf<ChatItem>(
|
||||
MessageItem(id = "m1", role = "user", text = "count", ts = 100),
|
||||
ToolItem(id = "tool_1", index = 0, name = "terminal", done = true, anchorId = "m1"),
|
||||
MessageItem(id = "m2", role = "assistant", text = "16", ts = 200),
|
||||
),
|
||||
),
|
||||
)
|
||||
// A history refresh (authoritative messages, no tool cards) must keep
|
||||
// the tool card between the user message and the answer — not push
|
||||
// it to the end.
|
||||
store.loadHistory(
|
||||
"android:default",
|
||||
listOf(
|
||||
MessageItem(id = "m1", role = "user", text = "count", ts = 100),
|
||||
MessageItem(id = "m2", role = "assistant", text = "16", ts = 200),
|
||||
),
|
||||
)
|
||||
assertEquals(listOf("m1", "tool_1", "m2"), store.lanes.value["android:default"]!!.map { it.id })
|
||||
}
|
||||
|
||||
@Test
|
||||
fun loadHistoryDedupesAndKeepsUnanchoredToolsLast() {
|
||||
val store = ChatStore()
|
||||
store.loadFromCache(
|
||||
mapOf(
|
||||
"android:default" to
|
||||
listOf<ChatItem>(
|
||||
MessageItem(id = "m1", role = "user", text = "count", ts = 100),
|
||||
ToolItem(id = "tool_1", index = 0, name = "bash", done = true),
|
||||
),
|
||||
),
|
||||
)
|
||||
store.loadHistory(
|
||||
"android:default",
|
||||
listOf(MessageItem(id = "m1", role = "user", text = "count", ts = 100)),
|
||||
)
|
||||
// A card whose anchor is unknown falls to the end (degenerate case).
|
||||
assertEquals(listOf("m1", "tool_1"), store.lanes.value["android:default"]!!.map { it.id })
|
||||
}
|
||||
|
||||
@Test
|
||||
fun toolStartIdsNeverCollideWithRestoredCards() {
|
||||
val store = ChatStore()
|
||||
// Lane restored from the cache after a restart, holding a persisted
|
||||
// card minted by the PREVIOUS process.
|
||||
store.loadFromCache(
|
||||
mapOf(
|
||||
"android:default" to
|
||||
listOf<ChatItem>(
|
||||
MessageItem(id = "m1", role = "user", text = "hi", ts = 1),
|
||||
ToolItem(id = "tool_1", index = 0, name = "bash", done = true, anchorId = "m1"),
|
||||
),
|
||||
),
|
||||
)
|
||||
// A new turn must not re-mint an id the restored lane already holds
|
||||
// (duplicate LazyColumn key / upsert overwrite under the same PK).
|
||||
store.onFrame(
|
||||
Frame(
|
||||
type = TYPE_TOOL_START,
|
||||
chatId = "android",
|
||||
payload =
|
||||
IrisJson.instance.encodeToJsonElement(
|
||||
ToolStartPayload.serializer(),
|
||||
ToolStartPayload(index = 0, name = "terminal", preview = "ls"),
|
||||
),
|
||||
),
|
||||
)
|
||||
val ids = store.lanes.value["android:default"]!!.map { it.id }
|
||||
assertEquals(ids.size, ids.toSet().size)
|
||||
}
|
||||
}
|
||||
@@ -4,13 +4,41 @@ import app.cash.sqldelight.db.SqlDriver
|
||||
import app.cash.sqldelight.driver.jdbc.sqlite.JdbcSqliteDriver
|
||||
import iris.db.IrisDatabase
|
||||
import java.io.File
|
||||
import java.sql.DriverManager
|
||||
import java.util.Properties
|
||||
|
||||
actual fun appDataDir(): String = File(System.getProperty("user.home"), ".iris").apply { mkdirs() }.absolutePath
|
||||
|
||||
actual fun createCacheDriver(): SqlDriver {
|
||||
val driver = JdbcSqliteDriver("jdbc:sqlite:${File(appDataDir(), "iris_cache.db").absolutePath}")
|
||||
// v1: create the schema (no migrations yet; add .sqm files +
|
||||
// `Schema.migrate` when the schema changes).
|
||||
IrisDatabase.Schema.create(driver)
|
||||
return driver
|
||||
actual fun createCacheDriver(): SqlDriver = createCacheDriver(File(appDataDir(), "iris_cache.db"))
|
||||
|
||||
/** Schema-aware driver for the cache DB at [file]: creates the schema on
|
||||
* first run, applies pending .sqm migrations on existing DBs (e.g. v1 ->
|
||||
* v2: the `tool` table), and persists the schema version (PRAGMA
|
||||
* user_version). */
|
||||
internal fun createCacheDriver(file: File): SqlDriver {
|
||||
stampLegacyV1(file)
|
||||
return JdbcSqliteDriver("jdbc:sqlite:${file.absolutePath}", Properties(), IrisDatabase.Schema)
|
||||
}
|
||||
|
||||
/** One-time shim for pre-existing desktop cache DBs: the old code called
|
||||
* `Schema.create()` directly, which never stamped PRAGMA user_version — so
|
||||
* those files hold the v1 schema at user_version 0, which the schema-aware
|
||||
* driver would treat as a fresh DB and crash on (`CREATE TABLE message` on
|
||||
* an existing table). Stamp them as v1 so the 1 -> 2 migration runs. */
|
||||
internal fun stampLegacyV1(file: File) {
|
||||
if (!file.exists()) return
|
||||
DriverManager.getConnection("jdbc:sqlite:${file.absolutePath}").use { conn ->
|
||||
val userVersion =
|
||||
conn.createStatement().use { st ->
|
||||
st.executeQuery("PRAGMA user_version").use { rs -> if (rs.next()) rs.getInt(1) else 0 }
|
||||
}
|
||||
if (userVersion != 0) return
|
||||
val hasMessageTable =
|
||||
conn.createStatement().use { st ->
|
||||
st
|
||||
.executeQuery("SELECT 1 FROM sqlite_master WHERE type = 'table' AND name = 'message'")
|
||||
.use { rs -> rs.next() }
|
||||
}
|
||||
if (hasMessageTable) conn.createStatement().use { st -> st.execute("PRAGMA user_version = 1") }
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,97 @@
|
||||
package iris.data
|
||||
|
||||
import iris.platform.createCacheDriver
|
||||
import java.io.File
|
||||
import java.nio.file.Files
|
||||
import java.sql.DriverManager
|
||||
import kotlin.test.Test
|
||||
import kotlin.test.assertEquals
|
||||
|
||||
/** Migration path for pre-existing desktop cache DBs: the schema-aware
|
||||
* driver plus the legacy user_version-0 shim (DesktopStorage). */
|
||||
class CacheMigrationTest {
|
||||
private val v1Schema =
|
||||
"""
|
||||
CREATE TABLE message (
|
||||
lane TEXT NOT NULL,
|
||||
id TEXT NOT NULL,
|
||||
ts INTEGER NOT NULL,
|
||||
payload TEXT NOT NULL,
|
||||
PRIMARY KEY (lane, id)
|
||||
);
|
||||
CREATE TABLE channel (
|
||||
chat_id TEXT NOT NULL,
|
||||
payload TEXT NOT NULL,
|
||||
PRIMARY KEY (chat_id)
|
||||
);
|
||||
CREATE TABLE meta (
|
||||
key TEXT NOT NULL,
|
||||
value TEXT NOT NULL,
|
||||
PRIMARY KEY (key)
|
||||
);
|
||||
""".trimIndent()
|
||||
|
||||
private fun tempDb(name: String): File = Files.createTempDirectory("iris_cache_test").toFile().let { File(it, name) }
|
||||
|
||||
private fun File.queryInt(sql: String): Int =
|
||||
DriverManager.getConnection("jdbc:sqlite:$absolutePath").use { conn ->
|
||||
conn.createStatement().use { st ->
|
||||
st.executeQuery(sql).use { rs -> if (rs.next()) rs.getInt(1) else -1 }
|
||||
}
|
||||
}
|
||||
|
||||
private fun File.writeV1(userVersion: Int) {
|
||||
DriverManager.getConnection("jdbc:sqlite:$absolutePath").use { conn ->
|
||||
conn.createStatement().use { st ->
|
||||
v1Schema
|
||||
.split(";")
|
||||
.map { it.trim() }
|
||||
.filter { it.isNotEmpty() }
|
||||
.forEach { st.execute(it) }
|
||||
st.execute("INSERT INTO message (lane, id, ts, payload) VALUES ('l', 'm1', 1, '{}')")
|
||||
// (row simulates a legacy DB with data; must survive migration)
|
||||
if (userVersion > 0) st.execute("PRAGMA user_version = $userVersion")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
fun legacyV1DbWithZeroUserVersionMigrates() {
|
||||
val file = tempDb("legacy0.db")
|
||||
try {
|
||||
// The old desktop code called Schema.create() directly, which
|
||||
// never stamped user_version: v1 tables at user_version 0.
|
||||
file.writeV1(userVersion = 0)
|
||||
// Must migrate (1 -> 2 via the shim), not crash on
|
||||
// `CREATE TABLE message` against the existing table.
|
||||
createCacheDriver(file).use { }
|
||||
assertEquals(1, file.queryInt("SELECT COUNT(*) FROM message"))
|
||||
assertEquals(1, file.queryInt("SELECT 1 FROM sqlite_master WHERE type = 'table' AND name = 'tool'"))
|
||||
} finally {
|
||||
file.parentFile?.deleteRecursively()
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
fun legacyV1DbWithStampedVersionMigrates() {
|
||||
val file = tempDb("legacy1.db")
|
||||
try {
|
||||
file.writeV1(userVersion = 1)
|
||||
createCacheDriver(file).use { }
|
||||
assertEquals(1, file.queryInt("SELECT 1 FROM sqlite_master WHERE type = 'table' AND name = 'tool'"))
|
||||
} finally {
|
||||
file.parentFile?.deleteRecursively()
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
fun freshDbCreatesSchema() {
|
||||
val file = tempDb("fresh.db")
|
||||
try {
|
||||
createCacheDriver(file).use { }
|
||||
assertEquals(1, file.queryInt("SELECT 1 FROM sqlite_master WHERE type = 'table' AND name = 'tool'"))
|
||||
} finally {
|
||||
file.parentFile?.deleteRecursively()
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -6,7 +6,9 @@ import iris.protocol.ChannelInfo
|
||||
import iris.protocol.RuntimeMeta
|
||||
import kotlin.test.Test
|
||||
import kotlin.test.assertEquals
|
||||
import kotlin.test.assertFalse
|
||||
import kotlin.test.assertNull
|
||||
import kotlin.test.assertTrue
|
||||
|
||||
/** Cache DB round-trip tests (JDBC in-memory SQLite; shared by both host targets). */
|
||||
class ChatDbTest {
|
||||
@@ -42,22 +44,80 @@ class ChatDbTest {
|
||||
db.saveLanes(
|
||||
mapOf(
|
||||
"android:default" to
|
||||
listOf(
|
||||
listOf<ChatItem>(
|
||||
msg("m1", ts = 100),
|
||||
ToolItem(id = "tool_1", index = 0, name = "bash", anchorId = "m1"),
|
||||
msg("m2", role = "assistant", text = "hi", ts = 200),
|
||||
ToolItem(id = "tool_1", index = 0, name = "bash"),
|
||||
),
|
||||
"android:default::thr_1" to listOf(msg("m3", ts = 300)),
|
||||
"android:default::thr_1" to listOf<ChatItem>(msg("m3", ts = 300)),
|
||||
),
|
||||
)
|
||||
val loaded = db.loadLanes()
|
||||
assertEquals(setOf("android:default", "android:default::thr_1"), loaded.keys)
|
||||
// Tool cards are ephemeral — not persisted.
|
||||
assertEquals(listOf("m1", "m2"), loaded["android:default"]!!.map { it.id })
|
||||
// Tool cards are persisted and restored at their anchored position
|
||||
// (after the message they follow, before the answer).
|
||||
assertEquals(listOf("m1", "tool_1", "m2"), loaded["android:default"]!!.map { it.id })
|
||||
assertEquals(listOf("m3"), loaded["android:default::thr_1"]!!.map { it.id })
|
||||
// Ordered by ts.
|
||||
assertEquals(100L, loaded["android:default"]!![0].ts)
|
||||
assertEquals(200L, loaded["android:default"]!![1].ts)
|
||||
// Messages ordered by ts.
|
||||
assertEquals(100L, (loaded["android:default"]!![0] as MessageItem).ts)
|
||||
assertEquals(200L, (loaded["android:default"]!![2] as MessageItem).ts)
|
||||
}
|
||||
|
||||
@Test
|
||||
fun toolCardsShareAnchorKeepOrder() {
|
||||
val db = newDb()
|
||||
db.saveLanes(
|
||||
mapOf(
|
||||
"android:default" to
|
||||
listOf<ChatItem>(
|
||||
msg("m1", ts = 100),
|
||||
ToolItem(id = "tool_1", index = 0, name = "search_files", anchorId = "m1"),
|
||||
ToolItem(id = "tool_2", index = 1, name = "terminal", anchorId = "m1"),
|
||||
msg("m2", role = "assistant", ts = 200),
|
||||
),
|
||||
),
|
||||
)
|
||||
assertEquals(listOf("m1", "tool_1", "tool_2", "m2"), db.loadLanes()["android:default"]!!.map { it.id })
|
||||
}
|
||||
|
||||
@Test
|
||||
fun toolCardWithMissingAnchorFallsToEnd() {
|
||||
val db = newDb()
|
||||
db.saveLanes(
|
||||
mapOf(
|
||||
"android:default" to
|
||||
listOf<ChatItem>(
|
||||
msg("m1", ts = 100),
|
||||
ToolItem(id = "tool_1", index = 0, name = "bash", anchorId = "deleted"),
|
||||
msg("m2", role = "assistant", ts = 200),
|
||||
),
|
||||
),
|
||||
)
|
||||
assertEquals(listOf("m1", "m2", "tool_1"), db.loadLanes()["android:default"]!!.map { it.id })
|
||||
}
|
||||
|
||||
@Test
|
||||
fun restoreClosesOpenToolCard() {
|
||||
val db = newDb()
|
||||
db.saveLanes(
|
||||
mapOf(
|
||||
"android:default" to
|
||||
listOf<ChatItem>(
|
||||
msg("m1", ts = 100),
|
||||
ToolItem(id = "tool_1", index = 0, name = "bash", done = false, anchorId = "m1"),
|
||||
ToolItem(id = "tool_2", index = 1, name = "ls", done = true, ok = true, anchorId = "m1"),
|
||||
),
|
||||
),
|
||||
)
|
||||
val lane = db.loadLanes()["android:default"]!!
|
||||
// The process died before tool.end — the open card is closed as
|
||||
// interrupted, not left spinning.
|
||||
val open = lane.first { it.id == "tool_1" } as ToolItem
|
||||
assertTrue(open.done)
|
||||
assertFalse(open.ok)
|
||||
// A completed card is untouched.
|
||||
val done = lane.first { it.id == "tool_2" } as ToolItem
|
||||
assertTrue(done.ok)
|
||||
}
|
||||
|
||||
@Test
|
||||
@@ -82,21 +142,22 @@ class ChatDbTest {
|
||||
),
|
||||
)
|
||||
val lane = db.loadLanes()["android:default"]!!
|
||||
val msgItem = { id: String -> lane.first { it.id == id } as MessageItem }
|
||||
// A pending send becomes failed (tap to retry); the gateway never
|
||||
// acknowledged it before the process died.
|
||||
assertEquals(MsgStatus.Failed, lane.first { it.id == "p1" }.status)
|
||||
assertEquals(false, lane.first { it.id == "p1" }.pending)
|
||||
assertEquals(MsgStatus.Failed, msgItem("p1").status)
|
||||
assertEquals(false, msgItem("p1").pending)
|
||||
// A streaming bubble is restored as finalized.
|
||||
assertEquals(false, lane.first { it.id == "s1" }.streaming)
|
||||
assertEquals(false, msgItem("s1").streaming)
|
||||
// Read status is preserved.
|
||||
assertEquals(MsgStatus.Read, lane.first { it.id == "r1" }.status)
|
||||
assertEquals(MsgStatus.Read, msgItem("r1").status)
|
||||
}
|
||||
|
||||
@Test
|
||||
fun systemMessagesAreNotPersisted() {
|
||||
val db = newDb()
|
||||
db.saveLanes(mapOf("android:default" to listOf(msg("sys_1", role = "system", isSystem = true))))
|
||||
assertEquals(emptyMap<String, List<MessageItem>>(), db.loadLanes())
|
||||
db.saveLanes(mapOf("android:default" to listOf<ChatItem>(msg("sys_1", role = "system", isSystem = true))))
|
||||
assertTrue(db.loadLanes().isEmpty())
|
||||
}
|
||||
|
||||
@Test
|
||||
@@ -121,8 +182,8 @@ class ChatDbTest {
|
||||
),
|
||||
),
|
||||
)
|
||||
db.saveLanes(mapOf("android:default" to listOf(item)))
|
||||
val loaded = db.loadLanes()["android:default"]!!.first()
|
||||
db.saveLanes(mapOf("android:default" to listOf<ChatItem>(item)))
|
||||
val loaded = db.loadLanes()["android:default"]!!.first() as MessageItem
|
||||
assertEquals("gpt", loaded.runtime?.model)
|
||||
assertEquals("/tmp/a.png", loaded.media.first().localPath)
|
||||
}
|
||||
@@ -150,10 +211,10 @@ class ChatDbTest {
|
||||
assertNull(db.metaGet("last_lane"))
|
||||
db.metaPut("last_lane", "android:chan_1")
|
||||
assertEquals("android:chan_1", db.metaGet("last_lane"))
|
||||
db.saveLanes(mapOf("android:default" to listOf(msg("m1"))))
|
||||
db.saveLanes(mapOf("android:default" to listOf<ChatItem>(msg("m1"))))
|
||||
db.saveChannels(listOf(ChannelInfo(chatId = "android:default", name = "General")))
|
||||
db.clearAll()
|
||||
assertEquals(emptyMap<String, List<MessageItem>>(), db.loadLanes())
|
||||
assertTrue(db.loadLanes().isEmpty())
|
||||
assertEquals(emptyList<ChannelInfo>(), db.loadChannels())
|
||||
assertNull(db.metaGet("last_lane"))
|
||||
}
|
||||
|
||||
Reference in new issue
Block a user