Skip to content
Closed
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
56 changes: 52 additions & 4 deletions packages/app/src/pages/session.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -94,6 +94,7 @@ import { formatServerError, isLocalSessionNotFoundError, isSessionNotFoundError
import { legacySessionHref, requireServerKey, sessionHref } from "@/utils/session-route"
import { useUsageExceededDialogs } from "./session/usage-exceeded-dialogs"
import { createSessionOwnership } from "./session/session-ownership"
import { restoreSessionScrollPosition, shouldResumeSessionAutoScroll } from "./session/session-scroll"
import { createSessionLineage } from "./session/session-lineage"

type FollowupItem = FollowupDraft & { id: string }
Expand Down Expand Up @@ -1401,12 +1402,19 @@ export default function Page() {
working: () => true,
overflowAnchor: "none",
})
const shouldResumeAutoScroll = () =>
shouldResumeSessionAutoScroll({
locationHash: location.hash,
messageId: store.messageId,
pendingMessage: ui.pendingMessage,
savedScroll: view().scroll("session"),
})
createEffect(
on(
() => params.id,
(id, previous) => {
if (!id || !previous || id === previous) return
if (location.hash || store.messageId || ui.pendingMessage) return
if (!shouldResumeAutoScroll()) return
autoScroll.resume()
},
),
Expand Down Expand Up @@ -1444,6 +1452,43 @@ export default function Page() {
})
}

let sessionScrollRestoreFrame: number | undefined
const persistSessionScroll = (el: HTMLDivElement) => {
view().setScroll("session", {
x: el.scrollLeft,
y: el.scrollTop,
})
}

const restoreSessionScroll = () => {
const el = scroller
if (!el) return

const next = restoreSessionScrollPosition({
savedScroll: view().scroll("session"),
clientWidth: el.clientWidth,
clientHeight: el.clientHeight,
scrollWidth: el.scrollWidth,
scrollHeight: el.scrollHeight,
})
if (!next) return

if (el.scrollLeft !== next.x) el.scrollLeft = next.x
if (el.scrollTop !== next.y) el.scrollTop = next.y
if (next.awayFromBottom) autoScroll.pause()

scheduleScrollState(el)
}

const queueSessionScrollRestore = () => {
if (sessionScrollRestoreFrame !== undefined) return

sessionScrollRestoreFrame = requestAnimationFrame(() => {
sessionScrollRestoreFrame = undefined
restoreSessionScroll()
})
}

const resumeScroll = () => {
setStore("messageId", undefined)
autoScroll.resume()
Expand Down Expand Up @@ -1474,6 +1519,7 @@ export default function Page() {
autoScroll.scrollRef(el)
if (!el) return
scheduleScrollState(el)
queueSessionScrollRestore()
fill()
}

Expand All @@ -1486,6 +1532,7 @@ export default function Page() {
() => {
const el = scroller
if (el) scheduleScrollState(el)
queueSessionScrollRestore()
fill()
},
)
Expand Down Expand Up @@ -1528,6 +1575,7 @@ export default function Page() {
}

onCleanup(() => {
if (sessionScrollRestoreFrame !== undefined) cancelAnimationFrame(sessionScrollRestoreFrame)
if (historyContinuationFrame !== undefined) cancelAnimationFrame(historyContinuationFrame)
})

Expand Down Expand Up @@ -2043,23 +2091,23 @@ export default function Page() {
scroll={ui.scroll}
onResumeScroll={resumeScroll}
setScrollRef={setScrollRef}
onPersistScrollPosition={persistSessionScroll}
onScheduleScrollState={scheduleScrollState}
onAutoScrollHandleScroll={autoScroll.handleScroll}
onMarkScrollGesture={markScrollGesture}
hasScrollGesture={hasScrollGesture}
onUserScroll={markUserScroll}
onHistoryScroll={onHistoryScroll}
onAutoScrollInteraction={autoScroll.handleInteraction}
shouldAnchorBottom={() =>
!location.hash && !store.messageId && !ui.pendingMessage && !autoScroll.userScrolled()
}
shouldAnchorBottom={() => shouldResumeAutoScroll() && !autoScroll.userScrolled()}
centered={centered()}
setContentRef={(el) => {
content = el
autoScroll.contentRef(el)

const root = scroller
if (root) scheduleScrollState(root)
queueSessionScrollRestore()
}}
userMessages={visibleUserMessages()}
setHistoryAnchor={(handlers) => {
Expand Down
42 changes: 42 additions & 0 deletions packages/app/src/pages/session/session-scroll.test.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,42 @@
import { describe, expect, test } from "bun:test"
import { restoreSessionScrollPosition, shouldResumeSessionAutoScroll } from "./session-scroll"

describe("session scroll restoration", () => {
test("does not resume auto-scroll when a session scroll position was saved", () => {
expect(
shouldResumeSessionAutoScroll({
locationHash: "",
messageId: undefined,
pendingMessage: undefined,
savedScroll: { x: 0, y: 0 },
}),
).toBe(false)
})

test("resumes auto-scroll when there is no saved session scroll position", () => {
expect(
shouldResumeSessionAutoScroll({
locationHash: "",
messageId: undefined,
pendingMessage: undefined,
savedScroll: undefined,
}),
).toBe(true)
})

test("restores a saved position and marks it away from the bottom", () => {
expect(
restoreSessionScrollPosition({
savedScroll: { x: 25, y: 300 },
clientWidth: 400,
clientHeight: 500,
scrollWidth: 900,
scrollHeight: 1200,
}),
).toEqual({
x: 25,
y: 300,
awayFromBottom: true,
})
})
})
39 changes: 39 additions & 0 deletions packages/app/src/pages/session/session-scroll.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,39 @@
import type { SessionScroll } from "@/context/layout-scroll"

const DEFAULT_BOTTOM_THRESHOLD = 10

const clamp = (value: number, min: number, max: number) => Math.max(min, Math.min(value, max))

export function shouldResumeSessionAutoScroll(input: {
locationHash: string
messageId?: string
pendingMessage?: string
savedScroll?: SessionScroll
}) {
if (input.locationHash) return false
if (input.messageId || input.pendingMessage) return false
return input.savedScroll === undefined
}

export function restoreSessionScrollPosition(input: {
savedScroll?: SessionScroll
clientWidth: number
clientHeight: number
scrollWidth: number
scrollHeight: number
bottomThreshold?: number
}) {
const saved = input.savedScroll
if (!saved) return undefined

const maxX = Math.max(0, input.scrollWidth - input.clientWidth)
const maxY = Math.max(0, input.scrollHeight - input.clientHeight)
const x = clamp(saved.x, 0, maxX)
const y = clamp(saved.y, 0, maxY)

return {
x,
y,
awayFromBottom: maxY - y > (input.bottomThreshold ?? DEFAULT_BOTTOM_THRESHOLD),
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -239,6 +239,7 @@ export function MessageTimeline(props: {
scroll: { overflow: boolean; bottom: boolean; jump: boolean }
onResumeScroll: () => void
setScrollRef: (el: HTMLDivElement | undefined) => void
onPersistScrollPosition: (el: HTMLDivElement) => void
onScheduleScrollState: (el: HTMLDivElement) => void
onAutoScrollHandleScroll: () => void
onMarkScrollGesture: (target?: EventTarget | null) => void
Expand Down Expand Up @@ -613,6 +614,7 @@ export function MessageTimeline(props: {

const handleListScroll = (event: Event & { currentTarget: HTMLDivElement }) => {
if (prependLoading) updatePrependAnchor()
props.onPersistScrollPosition(event.currentTarget)
props.onScheduleScrollState(event.currentTarget)
props.onHistoryScroll()
if (!props.hasScrollGesture()) return
Expand Down
Loading