diff --git a/packages/app/src/pages/session.tsx b/packages/app/src/pages/session.tsx index 4c796899c471..f2eea897988b 100644 --- a/packages/app/src/pages/session.tsx +++ b/packages/app/src/pages/session.tsx @@ -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 } @@ -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() }, ), @@ -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() @@ -1474,6 +1519,7 @@ export default function Page() { autoScroll.scrollRef(el) if (!el) return scheduleScrollState(el) + queueSessionScrollRestore() fill() } @@ -1486,6 +1532,7 @@ export default function Page() { () => { const el = scroller if (el) scheduleScrollState(el) + queueSessionScrollRestore() fill() }, ) @@ -1528,6 +1575,7 @@ export default function Page() { } onCleanup(() => { + if (sessionScrollRestoreFrame !== undefined) cancelAnimationFrame(sessionScrollRestoreFrame) if (historyContinuationFrame !== undefined) cancelAnimationFrame(historyContinuationFrame) }) @@ -2043,6 +2091,7 @@ export default function Page() { scroll={ui.scroll} onResumeScroll={resumeScroll} setScrollRef={setScrollRef} + onPersistScrollPosition={persistSessionScroll} onScheduleScrollState={scheduleScrollState} onAutoScrollHandleScroll={autoScroll.handleScroll} onMarkScrollGesture={markScrollGesture} @@ -2050,9 +2099,7 @@ export default function Page() { 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 @@ -2060,6 +2107,7 @@ export default function Page() { const root = scroller if (root) scheduleScrollState(root) + queueSessionScrollRestore() }} userMessages={visibleUserMessages()} setHistoryAnchor={(handlers) => { diff --git a/packages/app/src/pages/session/session-scroll.test.ts b/packages/app/src/pages/session/session-scroll.test.ts new file mode 100644 index 000000000000..220eebee37a2 --- /dev/null +++ b/packages/app/src/pages/session/session-scroll.test.ts @@ -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, + }) + }) +}) diff --git a/packages/app/src/pages/session/session-scroll.ts b/packages/app/src/pages/session/session-scroll.ts new file mode 100644 index 000000000000..8166cb2d22d0 --- /dev/null +++ b/packages/app/src/pages/session/session-scroll.ts @@ -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), + } +} diff --git a/packages/app/src/pages/session/timeline/message-timeline.tsx b/packages/app/src/pages/session/timeline/message-timeline.tsx index 2f3421a7507f..54a31d88db4c 100644 --- a/packages/app/src/pages/session/timeline/message-timeline.tsx +++ b/packages/app/src/pages/session/timeline/message-timeline.tsx @@ -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 @@ -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