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
6 changes: 6 additions & 0 deletions .changeset/cloud-session-import.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,6 @@
---
"@kilocode/cli": patch
"@kilocode/kilo-gateway": patch
---

Fix Cloud Agent session imports in installed CLI builds and prevent malformed exports or write failures from leaving partial imports.
336 changes: 261 additions & 75 deletions packages/kilo-gateway/src/cloud-sessions.ts
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import { z } from "zod"
import { buildKiloHeaders } from "./headers.js"

export const UUID_RE = /^[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}$/i
Expand All @@ -6,6 +7,147 @@ export interface DrizzleDb {
insert(table: object): { values(data: object): { onConflictDoNothing(): { run(): void } } }
}

export interface PrepareDeps {
Instance: {
readonly directory: string
readonly project: { readonly id: string }
}
readonly workspaceID?: string
readonly path?: string
Identifier: {
ascending(prefix: "session" | "message" | "part", given?: string): string
descending(prefix: "session" | "message" | "part", given?: string): string
}
}

export class SessionImportValidationError extends Error {}

const fileSchema = z
.object({
id: z.string(),
sessionID: z.string(),
messageID: z.string(),
type: z.literal("file"),
mime: z.string(),
url: z.string(),
})
.passthrough()

const stateSchema = z
.object({
status: z.literal("completed"),
attachments: z.array(fileSchema).optional(),
})
.passthrough()

const partSchema = z
.object({
id: z.string(),
sessionID: z.string(),
messageID: z.string(),
type: z.string(),
tail_start_id: z.unknown().optional(),
state: z.unknown().optional(),
})
.passthrough()

const messageSchema = z.object({
info: z
.object({
id: z.string(),
sessionID: z.string(),
parentID: z.string().optional(),
role: z.enum(["user", "assistant"]),
time: z.object({ created: z.number().finite() }).passthrough(),
})
.passthrough(),
parts: z.array(partSchema),
})

const exportSchema = z
.object({
info: z
.object({
id: z.string(),
time: z
.object({
created: z.number().optional(),
updated: z.number().optional(),
compacting: z.number().optional(),
archived: z.number().optional(),
})
.passthrough()
.optional(),
})
.passthrough(),
messages: z.array(messageSchema),
})
.superRefine((data, ctx) => {
const ids = new Set<string>()
const pids = new Set<string>()
const parents = new Map<string, string | undefined>()

for (const msg of data.messages) {
if (msg.info.sessionID !== data.info.id) ctx.addIssue({ code: "custom", message: "Invalid message info" })
if (ids.has(msg.info.id)) ctx.addIssue({ code: "custom", message: "Duplicate message ID" })
ids.add(msg.info.id)
parents.set(msg.info.id, msg.info.parentID)

for (const part of msg.parts) {
if (part.sessionID !== data.info.id || part.messageID !== msg.info.id)
ctx.addIssue({ code: "custom", message: "Invalid message part" })
if (part.type === "compaction" && part.tail_start_id !== undefined && typeof part.tail_start_id !== "string")
ctx.addIssue({ code: "custom", message: "Invalid compaction tail" })
if (pids.has(part.id)) ctx.addIssue({ code: "custom", message: "Duplicate part ID" })
pids.add(part.id)

if (part.type !== "tool") continue
const status = z.object({ status: z.unknown() }).safeParse(part.state)
if (!status.success || status.data.status !== "completed") continue
const state = stateSchema.safeParse(part.state)
if (!state.success) {
ctx.addIssue({ code: "custom", message: "Invalid tool attachments" })
continue
}
for (const file of state.data.attachments ?? []) {
if (file.sessionID !== data.info.id || file.messageID !== msg.info.id)
ctx.addIssue({ code: "custom", message: "Invalid tool attachment" })
if (pids.has(file.id)) ctx.addIssue({ code: "custom", message: "Duplicate part ID" })
pids.add(file.id)
}
}
}

for (const msg of data.messages) {
const parent = msg.info.parentID
if (parent !== undefined && !ids.has(parent))
ctx.addIssue({ code: "custom", message: "Dangling message parent" })

const seen = new Set([msg.info.id])
let current = parent
while (current !== undefined) {
if (seen.has(current)) {
ctx.addIssue({ code: "custom", message: "Circular message parent" })
break
}
seen.add(current)
current = parents.get(current)
}

for (const part of msg.parts) {
if (part.type !== "compaction" || typeof part.tail_start_id !== "string") continue
if (!ids.has(part.tail_start_id)) ctx.addIssue({ code: "custom", message: "Dangling compaction tail" })
}
}
})

function completed(part: z.infer<typeof partSchema>) {
if (part.type !== "tool") return
const result = stateSchema.safeParse(part.state)
if (!result.success) return
return result.data
}

const INGEST_BASE = process.env.KILO_SESSION_INGEST_URL ?? "https://ingest.kilosessions.ai"
const TIMEOUT = 30_000

Expand Down Expand Up @@ -56,106 +198,150 @@ export async function fetchCloudSessionForImport(token: string, sessionId: strin
return { ok: true, data }
}

export interface ImportDeps {
export interface ImportDeps extends PrepareDeps {
Database: {
transaction<T>(callback: (db: DrizzleDb) => T): T
effect(fn: () => void | Promise<unknown>): void
}
Instance: {
readonly directory: string
readonly project: { readonly id: string }
}
SessionTable: object
MessageTable: object
PartTable: object
SessionToRow: (info: any) => Record<string, unknown>
Bus: { publish(event: { type: string; properties: unknown }, payload: unknown): void | Promise<unknown> }
SessionCreatedEvent: { type: string; properties: unknown }
Identifier: {
ascending(prefix: "session" | "message" | "part", given?: string): string
descending(prefix: "session" | "message" | "part", given?: string): string
}
}

export function importSessionToDb(data: any, deps: ImportDeps) {
const {
Database,
Instance,
SessionTable,
MessageTable,
PartTable,
SessionToRow,
Bus,
SessionCreatedEvent,
Identifier,
} = deps

const localSessionID = Identifier.descending("session")
const msgMap = new Map<string, string>()
const projectID = Instance.project.id
export function prepareSessionImport(data: unknown, deps: PrepareDeps) {
const parsed = exportSchema.safeParse(data)
if (!parsed.success)
throw new SessionImportValidationError(parsed.error.issues[0]?.message ?? "Invalid session export")
const source = parsed.data

const sessionID = deps.Identifier.descending("session")
const ids = new Map<string, string>()
const pids = new Map<string, string>()
for (const msg of source.messages) {
ids.set(msg.info.id, deps.Identifier.ascending("message"))
for (const part of msg.parts) {
pids.set(part.id, deps.Identifier.ascending("part"))
for (const file of completed(part)?.attachments ?? []) {
pids.set(file.id, deps.Identifier.ascending("part"))
}
}
}

const now = Date.now()
const time = {
created: data.info.time?.created ?? now,
created: source.info.time?.created ?? now,
updated: now,
...(data.info.time?.compacting !== undefined && { compacting: data.info.time.compacting }),
...(data.info.time?.archived !== undefined && { archived: data.info.time.archived }),
...(source.info.time?.compacting !== undefined && { compacting: source.info.time.compacting }),
...(source.info.time?.archived !== undefined && { archived: source.info.time.archived }),
}

const info = {
...data.info,
id: localSessionID,
projectID,
slug: data.info.slug,
directory: Instance.directory,
version: data.info.version,
const info: Record<string, unknown> & {
id: string
projectID: string
directory: string
time: typeof time
} = {
...source.info,
id: sessionID,
projectID: deps.Instance.project.id,
slug: source.info.slug,
directory: deps.Instance.directory,
version: source.info.version,
time,
}
delete info.workspaceID
delete info.path
if (deps.workspaceID !== undefined) info.workspaceID = deps.workspaceID
if (deps.path !== undefined) info.path = deps.path
delete info.parentID
delete info.share
delete info.revert
delete info.permission

Database.transaction((db) => {
db.insert(SessionTable)
.values(SessionToRow(info as Record<string, unknown>))
.onConflictDoNothing()
.run()

const messages = Array.isArray(data.messages) ? data.messages : []
for (const msg of messages.filter((m: any) => m.info)) {
const msgID = Identifier.ascending("message")
msgMap.set(msg.info.id, msgID)
msg.info.id = msgID
msg.info.sessionID = localSessionID
if (msg.info.parentID) msg.info.parentID = msgMap.get(msg.info.parentID) ?? msg.info.parentID

db.insert(MessageTable)
.values({
id: msgID,
session_id: localSessionID,
time_created: msg.info.time?.created ?? Date.now(),
data: msg.info,
})
.onConflictDoNothing()
.run()
const messages: Array<{
id: string
session_id: string
time_created: number
data: Record<string, unknown>
}> = []
const parts: Array<{
id: string
message_id: string
session_id: string
data: Record<string, unknown>
}> = []
for (const msg of source.messages) {
const id = ids.get(msg.info.id)!
const parentID = msg.info.parentID === undefined ? undefined : ids.get(msg.info.parentID)!
const next = {
...msg.info,
id,
sessionID,
...(parentID ? { parentID } : {}),
}
messages.push({ id, session_id: sessionID, time_created: msg.info.time.created, data: next })

for (const part of msg.parts ?? []) {
const partID = Identifier.ascending("part")
part.id = partID
part.messageID = msgID
part.sessionID = localSessionID

db.insert(PartTable)
.values({
id: partID,
message_id: msgID,
session_id: localSessionID,
data: part,
})
.onConflictDoNothing()
.run()
for (const part of msg.parts) {
const partID = pids.get(part.id)!
const tail =
part.type === "compaction" && typeof part.tail_start_id === "string"
? ids.get(part.tail_start_id)!
: undefined
const data: Record<string, unknown> = {
...part,
id: partID,
messageID: id,
sessionID,
...(tail ? { tail_start_id: tail } : {}),
}
const state = completed(part)
if (state?.attachments) {
data.state = {
...state,
attachments: state.attachments.map((file) => {
const fileID = pids.get(file.id)!
return { ...file, id: fileID, messageID: id, sessionID }
}),
}
}
parts.push({
id: partID,
message_id: id,
session_id: sessionID,
data,
})
}
}

return { info, messages, parts }
}

export function importSessionToDb(data: unknown, deps: ImportDeps) {
const prepared = prepareSessionImport(data, deps)

deps.Database.transaction((db) => {
db.insert(deps.SessionTable).values(deps.SessionToRow(prepared.info)).onConflictDoNothing().run()

for (const row of prepared.messages) {
const { id: _, sessionID: __, ...data } = row.data
db.insert(deps.MessageTable)
.values({ id: row.id, session_id: row.session_id, time_created: row.time_created, data })
.onConflictDoNothing()
.run()
}
for (const row of prepared.parts) {
const { id: _, messageID: __, sessionID: ___, ...data } = row.data
db.insert(deps.PartTable)
.values({ id: row.id, message_id: row.message_id, session_id: row.session_id, data })
.onConflictDoNothing()
.run()
}

Database.effect(() => Bus.publish(SessionCreatedEvent, { info }))
deps.Database.effect(() => deps.Bus.publish(deps.SessionCreatedEvent, { info: prepared.info }))
})

return info
return prepared.info
}
Loading
Loading