|
|
@@ -18,7 +18,7 @@ import { message as cleanMessage } from "@/utils/diffs"
|
|
|
import { sessionNotFoundError } from "@/utils/server-errors"
|
|
|
import { rootSession } from "@/utils/session-route"
|
|
|
import { normalizeSessionInfo } from "@/utils/session"
|
|
|
-import { normalizeSessionMessages } from "@/utils/session-message"
|
|
|
+import { compareMessages, messageKey, normalizeSessionMessages } from "@/utils/session-message"
|
|
|
import { dropSessionCaches, pickSessionCacheEvictions, SESSION_CACHE_LIMIT } from "./global-sync/session-cache"
|
|
|
import { createV2SessionReducer, type V2SessionReduction } from "./server-session-v2-reducer"
|
|
|
import type { ServerApi } from "@/utils/server"
|
|
|
@@ -26,7 +26,6 @@ import type { ServerApi } from "@/utils/server"
|
|
|
type MessageApi = ServerApi["message"]
|
|
|
|
|
|
const cmp = (a: string, b: string) => (a < b ? -1 : a > b ? 1 : 0)
|
|
|
-const cmpMessage = (a: Message, b: Message) => a.time.created - b.time.created || cmp(a.id, b.id)
|
|
|
const SKIP_PARTS = new Set(["patch", "step-start", "step-finish"])
|
|
|
const initialMessagePageSize = 20
|
|
|
const historyMessagePageSize = 200
|
|
|
@@ -64,7 +63,7 @@ type MessagePage = {
|
|
|
function legacyMessageSource(items: { info: Message; parts: Part[] }[]): SessionMessageInfo[] {
|
|
|
return items
|
|
|
.slice()
|
|
|
- .sort((a, b) => cmp(a.info.id, b.info.id))
|
|
|
+ .sort((a, b) => compareMessages(a.info, b.info))
|
|
|
.map((item) => {
|
|
|
if (item.info.role === "user") {
|
|
|
return {
|
|
|
@@ -111,17 +110,16 @@ function mergeOptimisticPage(page: MessagePage, items: OptimisticItem[]) {
|
|
|
const part = new Map(page.part.map((item) => [item.id, item.part]))
|
|
|
const observed: { messageID: string; parts: Part[] }[] = []
|
|
|
for (const item of items) {
|
|
|
- const result = Binary.search(session, item.message.id, (message) => message.id)
|
|
|
- if (!result.found) session.splice(result.index, 0, item.message)
|
|
|
+ const result = Binary.search(session, messageKey(item.message), messageKey)
|
|
|
+ const found = result.found
|
|
|
+ if (!found) session.splice(result.index, 0, item.message)
|
|
|
const current = part.get(item.message.id)
|
|
|
- const confirmed = result.found
|
|
|
- ? item.parts.filter((part) => Binary.search(current ?? [], part.id, (value) => value.id).found)
|
|
|
- : []
|
|
|
- if (result.found) observed.push({ messageID: item.message.id, parts: confirmed })
|
|
|
+ const confirmed = found ? item.parts.filter((part) => current?.some((value) => value.id === part.id)) : []
|
|
|
+ if (found) observed.push({ messageID: item.message.id, parts: confirmed })
|
|
|
part.set(
|
|
|
item.message.id,
|
|
|
merge(
|
|
|
- result.found ? (current ?? []) : merge(item.confirmedParts ?? [], current ?? []),
|
|
|
+ found ? (current ?? []) : merge(item.confirmedParts ?? [], current ?? []),
|
|
|
item.parts.filter((part) => !confirmed.includes(part)),
|
|
|
),
|
|
|
)
|
|
|
@@ -158,6 +156,7 @@ function reconcileFetched<T extends { id: string }>(
|
|
|
retained?: ReadonlySet<string>
|
|
|
removed?: ReadonlySet<string>
|
|
|
preserveUnfetched?: boolean | ((item: T) => boolean)
|
|
|
+ compare?: (a: T, b: T) => number
|
|
|
} = {},
|
|
|
) {
|
|
|
const result = new Map(fetched.map((item) => [item.id, item]))
|
|
|
@@ -180,7 +179,8 @@ function reconcileFetched<T extends { id: string }>(
|
|
|
if (!item) result.delete(id)
|
|
|
}
|
|
|
for (const id of options.removed ?? emptyIDs) result.delete(id)
|
|
|
- return [...result.values()].sort((a, b) => cmp(a.id, b.id))
|
|
|
+ const items = [...result.values()]
|
|
|
+ return options.compare ? items.sort(options.compare) : items
|
|
|
}
|
|
|
|
|
|
type ServerSessionOptions = { retry?: typeof retry; protocol?: Promise<"v1" | "v2"> }
|
|
|
@@ -413,8 +413,7 @@ export function createServerSession(
|
|
|
if (!load) return
|
|
|
// A part event keeps an existing parent when the fetched page omits it without overriding fetched metadata.
|
|
|
const messages = data.message[sessionID]
|
|
|
- if (messages && Binary.search(messages, messageID, (message) => message.id).found)
|
|
|
- load.retainedMessages.add(messageID)
|
|
|
+ if (messages?.some((message) => message.id === messageID)) load.retainedMessages.add(messageID)
|
|
|
const parts = load.touchedParts.get(messageID)
|
|
|
if (parts) {
|
|
|
parts.add(partID)
|
|
|
@@ -437,16 +436,14 @@ export function createServerSession(
|
|
|
load.touchedParts.set(messageID, new Set(parts))
|
|
|
load.carriedDeltaParts.set(messageID, new Set(parts))
|
|
|
const messages = data.message[sessionID]
|
|
|
- if (messages && Binary.search(messages, messageID, (message) => message.id).found)
|
|
|
- load.retainedMessages.add(messageID)
|
|
|
+ if (messages?.some((message) => message.id === messageID)) load.retainedMessages.add(messageID)
|
|
|
}
|
|
|
for (const [messageID, parts] of load.removedParts) {
|
|
|
const touched = load.touchedParts.get(messageID) ?? new Set<string>()
|
|
|
parts.forEach((partID) => touched.add(partID))
|
|
|
load.touchedParts.set(messageID, touched)
|
|
|
const messages = data.message[sessionID]
|
|
|
- if (messages && Binary.search(messages, messageID, (message) => message.id).found)
|
|
|
- load.retainedMessages.add(messageID)
|
|
|
+ if (messages?.some((message) => message.id === messageID)) load.retainedMessages.add(messageID)
|
|
|
}
|
|
|
for (const [messageID, parts] of load.optimisticParts) {
|
|
|
load.removedMessages.delete(messageID)
|
|
|
@@ -555,7 +552,7 @@ export function createServerSession(
|
|
|
const source = pages.flatMap((page) => page.data).toReversed()
|
|
|
const normalized = normalizeSessionMessages(sessionID, source)
|
|
|
return {
|
|
|
- session: normalized.messages.sort((a, b) => cmp(a.id, b.id)),
|
|
|
+ session: normalized.messages.sort(compareMessages),
|
|
|
part: [...normalized.parts.entries()]
|
|
|
.map(([id, part]) => ({ id, part: part.sort((a, b) => cmp(a.id, b.id)) }))
|
|
|
.sort((a, b) => cmp(a.id, b.id)),
|
|
|
@@ -572,7 +569,7 @@ export function createServerSession(
|
|
|
})
|
|
|
const items = (response.data ?? []).filter((item) => !!item?.info?.id)
|
|
|
return {
|
|
|
- session: items.map((item) => cleanMessage(item.info)).sort((a, b) => cmp(a.id, b.id)),
|
|
|
+ session: items.map((item) => cleanMessage(item.info)).sort(compareMessages),
|
|
|
part: items.map((item) => ({
|
|
|
id: item.info.id,
|
|
|
part: item.parts.filter((part) => !!part?.id).sort((a, b) => cmp(a.id, b.id)),
|
|
|
@@ -696,7 +693,7 @@ export function createServerSession(
|
|
|
const normalized = normalizeSessionMessages(sessionID, source)
|
|
|
return {
|
|
|
...page,
|
|
|
- session: normalized.messages.sort((a, b) => cmp(a.id, b.id)),
|
|
|
+ session: normalized.messages.sort(compareMessages),
|
|
|
part: [...normalized.parts.entries()]
|
|
|
.map(([id, part]) => ({ id, part: part.sort((a, b) => cmp(a.id, b.id)) }))
|
|
|
.sort((a, b) => cmp(a.id, b.id)),
|
|
|
@@ -713,6 +710,7 @@ export function createServerSession(
|
|
|
retained: load?.retainedMessages,
|
|
|
removed: load?.removedMessages,
|
|
|
preserveUnfetched,
|
|
|
+ compare: compareMessages,
|
|
|
})
|
|
|
batch(() => {
|
|
|
if (source) setData("session_message", sessionID, reconcile(source))
|
|
|
@@ -754,7 +752,7 @@ export function createServerSession(
|
|
|
try {
|
|
|
const page = await fetchMessages(sessionID, limit, before, () => resetMessageLoad(sessionID, load))
|
|
|
const first = page.session.reduce<Message | undefined>(
|
|
|
- (oldest, message) => (!oldest || cmpMessage(message, oldest) < 0 ? message : oldest),
|
|
|
+ (oldest, message) => (!oldest || compareMessages(message, oldest) < 0 ? message : oldest),
|
|
|
undefined,
|
|
|
)
|
|
|
if (generations.get(sessionID) !== active) return
|
|
|
@@ -804,14 +802,15 @@ export function createServerSession(
|
|
|
session: merge(
|
|
|
page.session,
|
|
|
parents.map((parent) => parent.message),
|
|
|
- ),
|
|
|
+ ).sort(compareMessages),
|
|
|
part: merge(
|
|
|
page.part,
|
|
|
parents.map((parent) => ({ id: parent.message.id, part: parent.parts })),
|
|
|
),
|
|
|
}
|
|
|
const preserveUnfetched =
|
|
|
- mode === "prepend" || (!result.complete && (!first || ((message: Message) => cmpMessage(message, first) < 0)))
|
|
|
+ mode === "prepend" ||
|
|
|
+ (!result.complete && (!first || ((message: Message) => compareMessages(message, first) < 0)))
|
|
|
applyMessagePage(
|
|
|
sessionID,
|
|
|
result,
|
|
|
@@ -928,7 +927,7 @@ export function createServerSession(
|
|
|
.message({ sessionID, messageID })
|
|
|
.then((message) => {
|
|
|
const current = data.session_message[sessionID] ?? []
|
|
|
- const messages = [...current.filter((item) => item.id !== message.id), message].sort((a, b) => cmp(a.id, b.id))
|
|
|
+ const messages = [...current.filter((item) => item.id !== message.id), message].sort(compareMessages)
|
|
|
projectV2({ sessionID, messages, touched: [message.id] })
|
|
|
})
|
|
|
.catch(() => {})
|
|
|
@@ -1051,7 +1050,7 @@ export function createServerSession(
|
|
|
setData("message", info.sessionID, [info])
|
|
|
return
|
|
|
}
|
|
|
- const result = Binary.search(messages, info.id, (message) => message.id)
|
|
|
+ const result = Binary.search(messages, messageKey(info), messageKey)
|
|
|
if (result.found) setData("message", info.sessionID, result.index, reconcile(info))
|
|
|
if (!result.found)
|
|
|
setData("message", info.sessionID, (value = []) => {
|
|
|
@@ -1084,8 +1083,8 @@ export function createServerSession(
|
|
|
produce((draft) => {
|
|
|
const messages = draft.message[props.sessionID]
|
|
|
if (messages) {
|
|
|
- const result = Binary.search(messages, props.messageID, (message) => message.id)
|
|
|
- if (result.found) messages.splice(result.index, 1)
|
|
|
+ const index = messages.findIndex((message) => message.id === props.messageID)
|
|
|
+ if (index >= 0) messages.splice(index, 1)
|
|
|
}
|
|
|
deleteMessageParts(draft, props.messageID)
|
|
|
}),
|
|
|
@@ -1097,7 +1096,7 @@ export function createServerSession(
|
|
|
if (SKIP_PARTS.has(part.type)) return
|
|
|
const messages = data.message[part.sessionID]
|
|
|
const load = messageLoads.get(part.sessionID)
|
|
|
- const missing = !messages || !Binary.search(messages, part.messageID, (message) => message.id).found
|
|
|
+ const missing = !messages?.some((message) => message.id === part.messageID)
|
|
|
// Outside a page load, accepting a part without its ordered parent event would create an unbounded orphan.
|
|
|
if (
|
|
|
missing &&
|
|
|
@@ -1341,7 +1340,7 @@ export function createServerSession(
|
|
|
if (items) items.set(input.message.id, { ...input, parts, confirmedParts: [] })
|
|
|
if (!items)
|
|
|
optimistic.set(input.sessionID, new Map([[input.message.id, { ...input, parts, confirmedParts: [] }]]))
|
|
|
- setData("message", input.sessionID, (messages = []) => merge(messages, [input.message]))
|
|
|
+ setData("message", input.sessionID, (messages = []) => merge(messages, [input.message]).sort(compareMessages))
|
|
|
setData(
|
|
|
"part_text_accum_delta",
|
|
|
produce((draft) => {
|