Merge pull request #9937 from Kilo-Org/recondite-crane

fix(jetbrains): auto-scroll
This commit is contained in:
Kirill Kalishev
2026-05-06 09:50:40 -04:00
committed by GitHub
16 changed files with 971 additions and 583 deletions
@@ -0,0 +1,5 @@
---
"@kilocode/kilo-jetbrains": patch
---
Open existing JetBrains sessions scrolled to the latest message after history loads.
+5
View File
@@ -0,0 +1,5 @@
---
"@kilocode/kilo-jetbrains": patch
---
Keep the JetBrains chat transcript pinned to bottom reliably while responses stream.
@@ -0,0 +1,5 @@
---
"@kilocode/kilo-jetbrains": patch
---
Reduce redundant JetBrains chat scroll layout work while keeping bottom-follow behavior.
@@ -0,0 +1,264 @@
package ai.kilocode.client.session
import ai.kilocode.client.plugin.KiloBundle
import ai.kilocode.client.session.ui.SessionMessageListPanel
import ai.kilocode.client.session.ui.SessionRootPanel
import ai.kilocode.client.session.ui.SessionStyle
import ai.kilocode.client.session.ui.SessionStyleTarget
import ai.kilocode.client.ui.UiStyle
import com.intellij.openapi.application.ApplicationManager
import com.intellij.openapi.util.IconLoader
import com.intellij.ui.components.JBLabel
import com.intellij.ui.components.JBScrollPane
import com.intellij.ui.icons.CachedImageIcon
import com.intellij.ui.svg.SvgAttributePatcher
import com.intellij.util.SVGLoader
import com.intellij.util.ui.JBUI
import java.awt.Color
import java.awt.Cursor
import java.awt.Point
import java.awt.Rectangle
import java.awt.event.MouseAdapter
import java.awt.event.MouseEvent
import javax.swing.Icon
import javax.swing.JComponent
import javax.swing.JPanel
import javax.swing.JScrollBar
private const val ICON_DIGEST = 0x5c011b0bb17L
private const val OPAQUE_ALPHA = 255
internal class SessionScroll(
private val root: SessionRootPanel,
private val host: JPanel,
private val messages: SessionMessageListPanel,
body: JPanel,
) {
companion object {
private val ICON = IconLoader.getIcon("/icons/scroll-bottom.svg", SessionScroll::class.java)
private const val THRESHOLD = 32
private const val OPEN_PASSES = 12
}
val component = JBScrollPane(body).apply {
border = JBUI.Borders.empty()
verticalScrollBarPolicy = JBScrollPane.VERTICAL_SCROLLBAR_AS_NEEDED
horizontalScrollBarPolicy = JBScrollPane.HORIZONTAL_SCROLLBAR_NEVER
}
internal val bar: JScrollBar get() = component.verticalScrollBar
internal val jump: JBLabel
val view: JComponent? get() = component.viewport.view as? JComponent
private var style = SessionStyle.current()
private var tail = true
private var auto = false
private var opening = false
private var stable = -1
private var seq = 0
init {
jump = JBLabel(patchedIcon(ICON)).apply {
cursor = Cursor.getPredefinedCursor(Cursor.HAND_CURSOR)
toolTipText = KiloBundle.message("session.scroll.bottom")
isVisible = false
addMouseListener(object : MouseAdapter() {
override fun mouseClicked(e: MouseEvent) {
jumpBottom()
}
})
}
component.verticalScrollBar.addAdjustmentListener { onScroll() }
root.addOverlay(jump) { _, child ->
val size = child.preferredSize
val gap = JBUI.scale(UiStyle.Space.PAD)
Rectangle(
host.x + host.width - size.width - gap,
host.y + host.height - size.height - gap,
size.width,
size.height,
)
}
}
fun show(panel: JPanel) {
if (component.viewport.view === panel) return
(panel as? SessionStyleTarget)?.applyStyle(style)
component.viewport.setView(panel)
component.repaint()
updateJump()
}
fun atBottom(): Boolean {
val bar = component.verticalScrollBar
return when {
component.viewport.view !== messages -> tail
bar.maximum <= bar.visibleAmount -> true
else -> bar.value + bar.visibleAmount >= bar.maximum - JBUI.scale(THRESHOLD)
}
}
fun followBottom(follow: Boolean) {
if (!follow) {
seq++
updateJump()
return
}
tail = true
auto = true
show(messages)
auto = false
followPass(++seq, 2)
}
fun openBottom(done: () -> Unit) {
opening = true
stable = -1
tail = true
auto = true
show(messages)
auto = false
val id = ++seq
ApplicationManager.getApplication().invokeLater {
openPass(id, OPEN_PASSES, done)
}
}
fun refresh() {
updateJump()
}
fun applyStyle(style: SessionStyle) {
this.style = style
jump.icon = patchedIcon(ICON)
messages.applyStyle(style)
val view = component.viewport.view
if (view !== messages) (view as? SessionStyleTarget)?.applyStyle(style)
refresh()
}
private fun jumpBottom() {
opening = false
stable = -1
tail = true
auto = true
show(messages)
auto = false
followPass(++seq, 2)
}
private fun followPass(id: Int, remaining: Int) {
if (id != seq || !tail) return
auto = true
try {
layoutScroll()
scrollToBottom()
updateJump()
} finally {
auto = false
}
if (remaining <= 0) return
ApplicationManager.getApplication().invokeLater {
followPass(id, remaining - 1)
}
}
private fun openPass(id: Int, remaining: Int, done: () -> Unit) {
if (id != seq) {
opening = false
stable = -1
return
}
auto = true
val prev = bottom()
try {
tail = true
layoutScroll()
scrollToBottom()
updateJump()
} finally {
auto = false
}
if (remaining <= 0) {
opening = false
stable = -1
done()
return
}
val next = bottom()
val left = if (next == prev && next == stable) remaining - 1 else OPEN_PASSES
stable = next
ApplicationManager.getApplication().invokeLater {
openPass(id, left, done)
}
}
private fun layoutScroll() {
root.validate()
}
private fun scrollToBottom() {
val view = component.viewport.view ?: return
val y = (view.height - component.viewport.extentSize.height).coerceAtLeast(0)
component.viewport.viewPosition = Point(0, y)
(view as? JComponent)?.scrollRectToVisible(Rectangle(0, view.height.coerceAtLeast(1) - 1, 1, 1))
val bar = component.verticalScrollBar
bar.value = bottom()
}
private fun bottom(): Int {
val bar = component.verticalScrollBar
return (bar.maximum - bar.visibleAmount).coerceAtLeast(bar.minimum)
}
private fun onScroll() {
if (auto || opening) {
updateJump()
return
}
if (component.viewport.view === messages) {
tail = atBottom()
if (!tail) seq++
}
updateJump()
}
private fun updateJump() {
val visible = component.viewport.view === messages && !atBottom()
if (jump.isVisible == visible) return
jump.isVisible = visible
root.overlay.revalidate()
root.overlay.repaint()
}
}
private fun patchedIcon(icon: Icon): Icon {
val cached = icon as? CachedImageIcon ?: return icon
return cached.createWithPatcher(object : SVGLoader.SvgElementColorPatcherProvider, SvgAttributePatcher {
override fun digest(): LongArray {
val bg = JBUI.CurrentTheme.Button.defaultButtonColorStart().rgb.toLong()
val fg = JBUI.CurrentTheme.Button.defaultButtonForeground().rgb.toLong()
return longArrayOf(bg, fg, ICON_DIGEST)
}
override fun attributeForPath(path: String) = this
override fun patchColors(attributes: MutableMap<String, String>) {
when (attributes["id"]) {
"ScrollButton.Background" ->
set(attributes, "fill", JBUI.CurrentTheme.Button.defaultButtonColorStart())
"ScrollButton.Foreground" ->
set(attributes, "stroke", JBUI.CurrentTheme.Button.defaultButtonForeground())
}
}
private fun set(attributes: MutableMap<String, String>, key: String, color: Color) {
if (!attributes.containsKey(key) || attributes[key] == "none") return
attributes[key] = "rgb(${color.red},${color.green},${color.blue})"
if (color.alpha != OPAQUE_ALPHA) {
attributes["$key-opacity"] = "${color.alpha / OPAQUE_ALPHA.toFloat()}"
}
}
})
}
@@ -30,23 +30,13 @@ import com.intellij.openapi.editor.colors.EditorColorsListener
import com.intellij.openapi.editor.colors.EditorColorsManager
import com.intellij.openapi.Disposable
import com.intellij.openapi.project.Project
import com.intellij.openapi.util.IconLoader
import com.intellij.openapi.util.registry.Registry
import com.intellij.ui.icons.CachedImageIcon
import com.intellij.ui.svg.SvgAttributePatcher
import com.intellij.ui.components.JBLabel
import com.intellij.ui.components.JBScrollPane
import com.intellij.util.SVGLoader
import com.intellij.util.ui.Centerizer
import com.intellij.util.ui.JBUI
import kotlinx.coroutines.CoroutineScope
import java.awt.BorderLayout
import java.awt.Color
import java.awt.Rectangle
import javax.swing.BoxLayout
import javax.swing.BoxLayout.Y_AXIS
import javax.swing.Icon
import javax.swing.JButton
import javax.swing.JComponent
import javax.swing.JPanel
@@ -93,11 +83,13 @@ class SessionUi private constructor(
companion object {
private val LOG = KiloLog.create(SessionUi::class.java)
private val SCROLL_ICON = IconLoader.getIcon("/icons/scroll-bottom.svg", SessionUi::class.java)
}
private val project = project
private val app = app
private var opening = id != null
private var pending = false
private var loaded: Boolean? = null
private val flushMs =
Registry.intValue("kilo.session.flushMs", EVENT_FLUSH_MS.toInt())
.takeIf { it > 0 }
@@ -110,8 +102,9 @@ class SessionUi private constructor(
condense = Registry.`is`("kilo.session.condense", true),
displayMs = displayMs,
open = open,
beforeUpdate = ::atBottom,
afterUpdate = ::followBottom,
beforeUpdate = { if (opening) false else scroll.atBottom() },
afterUpdate = { if (!opening) scroll.followBottom(it) },
loaded = ::onHistoryLoaded,
)
@@ -125,9 +118,7 @@ class SessionUi private constructor(
private lateinit var messageBody: SessionMessageListPanel
private lateinit var scroll: JBScrollPane
private lateinit var jump: JButton
internal lateinit var scroll: SessionScroll
private lateinit var question: QuestionPanel
private lateinit var permission: PermissionPanel
@@ -142,7 +133,19 @@ class SessionUi private constructor(
bindUi()
bindStyle()
applyStyle(style)
showBody(if (loading) progressBody else blankBody)
onStateChanged(controller.model.state)
scroll.show(startBody())
loaded?.let(::finishOpen)
}
override fun addNotify() {
super.addNotify()
resumeOpen()
}
override fun doLayout() {
super.doLayout()
resumeOpen()
}
internal val blank: Boolean get() = controller.blank
@@ -172,22 +175,7 @@ class SessionUi private constructor(
}
messageBody = SessionMessageListPanel(controller.model, this)
scroll = JBScrollPane(blankBody).apply {
border = JBUI.Borders.empty()
verticalScrollBarPolicy = JBScrollPane.VERTICAL_SCROLLBAR_AS_NEEDED
horizontalScrollBarPolicy = JBScrollPane.HORIZONTAL_SCROLLBAR_NEVER
}
jump = JButton(patchedIcon(SCROLL_ICON)).apply {
border = JBUI.Borders.empty()
isContentAreaFilled = false
isBorderPainted = false
isFocusPainted = false
isFocusable = false
isOpaque = false
toolTipText = KiloBundle.message("session.scroll.bottom")
isVisible = false
addActionListener { jumpBottom() }
}
scroll = SessionScroll(root, sessionContent, messageBody, blankBody)
question = QuestionPanel(controller)
permission = PermissionPanel(controller)
connection = ConnectionPanel(this, controller)
@@ -198,7 +186,7 @@ class SessionUi private constructor(
onAbort = { controller.abort() },
)
sessionContent.add(scroll, BorderLayout.CENTER)
sessionContent.add(scroll.component, BorderLayout.CENTER)
root.content.add(sessionContent, BorderLayout.CENTER)
// Dock panels stay in normal flow so each visible state takes layout space
// above the prompt.
@@ -209,17 +197,6 @@ class SessionUi private constructor(
add(connection)
add(prompt)
}, BorderLayout.SOUTH)
root.addOverlay(jump) { _, child ->
val size = child.preferredSize
val gap = JBUI.scale(12)
Rectangle(
sessionContent.x + sessionContent.width - size.width - gap,
sessionContent.y + sessionContent.height - size.height - gap,
size.width,
size.height,
)
}
add(root, BorderLayout.CENTER)
}
@@ -230,7 +207,6 @@ class SessionUi private constructor(
prompt.onReset = { controller.clearModelOverride() }
prompt.model.favorites = { app.favorites.value }
prompt.model.onFavoriteToggle = { item -> app.toggleModelFavorite(item.provider, item.id) }
scroll.verticalScrollBar.addAdjustmentListener { updateJump() }
controller.addListener(this) { event ->
when (event) {
@@ -264,16 +240,16 @@ class SessionUi private constructor(
}
is SessionControllerEvent.ViewChanged.ShowProgress -> {
showBody(progressBody)
scroll.show(progressBody)
}
is SessionControllerEvent.ViewChanged.ShowRecents -> {
val panel = EmptySessionPanel(this, controller, event.recents)
showBody(panel)
scroll.show(panel)
}
is SessionControllerEvent.ViewChanged.ShowSession -> {
showBody(messageBody)
scroll.show(messageBody)
}
is SessionControllerEvent.AppChanged,
@@ -322,10 +298,45 @@ class SessionUi private constructor(
})
}
private fun startBody(): JPanel {
if (controller.model.showSession) return messageBody
if (loading) return progressBody
return blankBody
}
private fun onHistoryLoaded(show: Boolean) {
loaded = show
if (!this::scroll.isInitialized) return
finishOpen(show)
}
private fun finishOpen(show: Boolean) {
loaded = show
if (!opening) return
if (!show) {
pending = false
opening = false
return
}
pending = true
resumeOpen()
}
private fun resumeOpen() {
if (!pending || !opening || !this::scroll.isInitialized) return
if (width <= 0 || height <= 0) return
pending = false
scroll.openBottom {
opening = false
}
}
private fun sendPrompt(text: String) {
if (text.isBlank()) return
LOG.debug {
"${ChatLogSummary.prompt(text)} agent=${controller.model.agent ?: "none"} model=${controller.model.model ?: "none"} ready=${controller.ready}"
val agent = controller.model.agent ?: "none"
val model = controller.model.model ?: "none"
"${ChatLogSummary.prompt(text)} agent=$agent model=$model ready=${controller.ready}"
}
controller.prompt(text)
prompt.clear()
@@ -352,106 +363,21 @@ class SessionUi private constructor(
refresh()
}
internal fun atBottom(): Boolean {
val bar = scroll.verticalScrollBar
if (bar.maximum <= bar.visibleAmount) return true
return bar.value + bar.visibleAmount >= bar.maximum - JBUI.scale(32)
}
internal fun followBottom(follow: Boolean) {
if (!follow) return
showBody(messageBody)
scrollToBottom()
updateJump()
ApplicationManager.getApplication().invokeLater {
scroll.viewport.view?.revalidate()
scroll.viewport.view?.doLayout()
scroll.revalidate()
scroll.doLayout()
scrollToBottom()
updateJump()
}
}
private fun jumpBottom() {
showBody(messageBody)
scrollToBottom()
updateJump()
ApplicationManager.getApplication().invokeLater {
scroll.viewport.view?.revalidate()
scroll.viewport.view?.doLayout()
scroll.revalidate()
scroll.doLayout()
scrollToBottom()
updateJump()
}
}
private fun scrollToBottom() {
val bar = scroll.verticalScrollBar
bar.value = bar.maximum
}
private fun updateJump() {
val visible = scroll.viewport.view === messageBody && !atBottom()
if (jump.isVisible == visible) return
jump.isVisible = visible
root.overlay.revalidate()
root.overlay.repaint()
}
private fun refresh() {
updateJump()
scroll.refresh()
root.revalidate()
root.repaint()
}
private fun showBody(panel: JPanel) {
if (scroll.viewport.view === panel) return
(panel as? SessionStyleTarget)?.applyStyle(style)
scroll.viewport.setView(panel)
scroll.revalidate()
scroll.repaint()
updateJump()
}
override fun applyStyle(style: SessionStyle) {
this.style = style
jump.icon = patchedIcon(SCROLL_ICON)
loadingLabel.font = style.uiFont
messageBody.applyStyle(style)
prompt.applyStyle(style)
(scroll.viewport.view as? SessionStyleTarget)?.applyStyle(style)
scroll.applyStyle(style)
refresh()
}
override fun dispose() {}
}
private fun patchedIcon(icon: Icon): Icon {
val cached = icon as? CachedImageIcon ?: return icon
return cached.createWithPatcher(object : SVGLoader.SvgElementColorPatcherProvider, SvgAttributePatcher {
override fun digest(): LongArray {
val bg = JBUI.CurrentTheme.Button.defaultButtonColorStart().rgb.toLong()
val fg = JBUI.CurrentTheme.Button.defaultButtonForeground().rgb.toLong()
return longArrayOf(bg, fg, 0x5c011b0bb17L)
}
override fun attributeForPath(path: String) = this
override fun patchColors(attributes: MutableMap<String, String>) {
when (attributes["id"]) {
"ScrollButton.Background" -> set(attributes, "fill", JBUI.CurrentTheme.Button.defaultButtonColorStart())
"ScrollButton.Foreground" -> set(attributes, "stroke", JBUI.CurrentTheme.Button.defaultButtonForeground())
}
}
private fun set(attributes: MutableMap<String, String>, key: String, color: Color) {
if (!attributes.containsKey(key) || attributes[key] == "none") return
attributes[key] = "rgb(${color.red},${color.green},${color.blue})"
if (color.alpha != 255) attributes["$key-opacity"] = "${color.alpha / 255f}"
}
})
}
private fun variantTitle(value: String): String = value.replaceFirstChar { it.titlecase() }
@@ -50,14 +50,20 @@ class SessionMessageListPanel(
is SessionModelEvent.TurnUpdated -> onTurnUpdated(event.turn)
is SessionModelEvent.TurnRemoved -> onTurnRemoved(event.id)
is SessionModelEvent.ContentAdded ->
is SessionModelEvent.ContentAdded -> {
msgToView[event.messageId]?.upsertPart(event.content)
refresh()
}
is SessionModelEvent.ContentUpdated ->
is SessionModelEvent.ContentUpdated -> {
msgToView[event.messageId]?.upsertPart(event.content)
refresh()
}
is SessionModelEvent.ContentRemoved ->
is SessionModelEvent.ContentRemoved -> {
msgToView[event.messageId]?.removePart(event.contentId)
refresh()
}
is SessionModelEvent.ContentDelta -> {
// Use the full current content from the model rather than
@@ -67,6 +73,7 @@ class SessionMessageListPanel(
// on first appendDelta and fires both events in sequence).
val content = model.content(event.messageId, event.contentId)
if (content != null) msgToView[event.messageId]?.upsertPart(content)
refresh()
}
is SessionModelEvent.HistoryLoaded -> rebuild()
@@ -71,6 +71,7 @@ class SessionController(
private val open: (SessionDto) -> Unit = {},
private val beforeUpdate: () -> Boolean = { false },
private val afterUpdate: (Boolean) -> Unit = {},
private val loaded: (Boolean) -> Unit = {},
) : Disposable {
companion object {
@@ -395,16 +396,18 @@ class SessionController(
}
}
recoverPending(id)
edt {
if (!model.isEmpty()) {
showMessages()
return@edt
}
refreshRecents(force = true)
runEdt {
val show = !model.isEmpty()
if (show) showMessages()
if (!show) refreshRecents(force = true)
loaded(show)
}
} catch (e: Exception) {
LOG.warn("${ChatLogSummary.sid(id)} kind=history dir=${ChatLogSummary.dir(directory)} failed message=${e.message}", e)
edt { refreshRecents(force = true) }
edt {
refreshRecents(force = true)
loaded(false)
}
} finally {
edt {
if (historyState != state) return@edt
@@ -54,6 +54,7 @@ class MessageView(
val existing = parts[content.id]
if (existing != null) {
existing.update(content)
refresh()
return
}
val view = ViewFactory.create(content)
@@ -61,8 +62,7 @@ class MessageView(
parts[content.id] = view
add(view)
syncBorder()
revalidate()
repaint()
refresh()
}
/** Remove the renderer for [contentId] if present. */
@@ -70,8 +70,7 @@ class MessageView(
val view = parts.remove(contentId) ?: return
remove(view)
syncBorder()
revalidate()
repaint()
refresh()
}
private fun syncBorder() {
@@ -81,7 +80,9 @@ class MessageView(
/** Append a streaming delta to the renderer for [contentId]. */
fun appendDelta(contentId: String, delta: String) {
parts[contentId]?.appendDelta(delta)
val part = parts[contentId] ?: return
part.appendDelta(delta)
refresh()
}
/** Look up a renderer by part id. */
@@ -96,6 +97,10 @@ class MessageView(
override fun applyStyle(style: SessionStyle) {
this.style = style
for (view in parts.values) view.applyStyle(style)
refresh()
}
private fun refresh() {
revalidate()
repaint()
}
@@ -28,10 +28,13 @@ class TextView(text: Text) : PartView() {
override fun update(content: Content) {
if (content !is Text) return
md.set(content.content.toString())
refresh()
}
override fun appendDelta(delta: String) {
if (delta.isEmpty()) return
md.append(delta)
refresh()
}
/** Current markdown source — used by tests to assert rendered content. */
@@ -42,6 +45,10 @@ class TextView(text: Text) : PartView() {
if (md.font != style.transcriptFont) md.font = style.transcriptFont
if (md.codeFont != style.editorFamily) md.codeFont = style.editorFamily
if (!changed) return
refresh()
}
private fun refresh() {
revalidate()
repaint()
}
@@ -2,6 +2,7 @@ package ai.kilocode.client.ui
import com.intellij.openapi.editor.colors.EditorColorsManager
import com.intellij.ui.JBColor
import com.intellij.ui.RoundedLineBorder
import com.intellij.ui.components.JBLabel
import com.intellij.util.ui.JBUI
import com.intellij.util.ui.UIUtil
@@ -1,4 +1,4 @@
<svg width="56" height="56" viewBox="0 0 56 56" fill="none" xmlns="http://www.w3.org/2000/svg">
<circle id="ScrollButton.Background" cx="28" cy="28" r="24" fill="#384F6B"/>
<path id="ScrollButton.Foreground" d="M28 17V37M28 37L17 26M28 37L39 26" stroke="#FFFFFF" stroke-width="3" stroke-linecap="round" stroke-linejoin="round"/>
<svg width="40" height="40" viewBox="0 0 40 40" fill="none" xmlns="http://www.w3.org/2000/svg">
<circle id="ScrollButton.Background" cx="20" cy="20" r="18" fill="#384F6B"/>
<path id="ScrollButton.Foreground" d="M20 12V27M20 27L14 21M20 27L26 21" stroke="#FFFFFF" stroke-width="1.5" stroke-linecap="round" stroke-linejoin="round"/>
</svg>

Before

Width:  |  Height:  |  Size: 340 B

After

Width:  |  Height:  |  Size: 342 B

@@ -1,4 +1,4 @@
<svg width="56" height="56" viewBox="0 0 56 56" fill="none" xmlns="http://www.w3.org/2000/svg">
<circle id="ScrollButton.Background" cx="28" cy="28" r="24" fill="#233143"/>
<path id="ScrollButton.Foreground" d="M28 17V37M28 37L17 26M28 37L39 26" stroke="#FFFFFF" stroke-width="3" stroke-linecap="round" stroke-linejoin="round"/>
<svg width="40" height="40" viewBox="0 0 40 40" fill="none" xmlns="http://www.w3.org/2000/svg">
<circle id="ScrollButton.Background" cx="20" cy="20" r="18" fill="#233143"/>
<path id="ScrollButton.Foreground" d="M20 12V27M20 27L14 21M20 27L26 21" stroke="#FFFFFF" stroke-width="1.5" stroke-linecap="round" stroke-linejoin="round"/>
</svg>

Before

Width:  |  Height:  |  Size: 340 B

After

Width:  |  Height:  |  Size: 342 B

@@ -0,0 +1,331 @@
package ai.kilocode.client.session
import ai.kilocode.client.session.ui.SessionMessageListPanel
import ai.kilocode.rpc.dto.ChatEventDto
import ai.kilocode.rpc.dto.PermissionRequestDto
import ai.kilocode.rpc.dto.SessionStatusDto
import com.intellij.util.ui.JBUI
import kotlinx.coroutines.CompletableDeferred
@Suppress("UnstableApiUsage")
class SessionScrollTest : SessionUiTestBase() {
fun `test session update follows when transcript is at bottom`() {
showMessages()
fillTranscript(24)
val bar = scrollBar()
setBottom(bar)
emit(ChatEventDto.MessageUpdated("ses_test", message("tail")))
drainScroll()
assertBottom(bar)
}
fun `test session update follows when transcript is near bottom threshold`() {
showMessages()
fillTranscript(24)
val bar = scrollBar()
val threshold = JBUI.scale(32)
if (bottom(bar) <= threshold) {
fillTranscript(24, start = 24)
}
setValue(bar, bottom(bar) - threshold + 1)
emit(ChatEventDto.MessageUpdated("ses_test", message("tail")))
drainScroll()
assertBottom(bar)
}
fun `test session update preserves position outside bottom threshold`() {
showMessages()
fillTranscript(24)
val bar = scrollBar()
val threshold = JBUI.scale(32)
setValue(bar, bottom(bar) - threshold - 8)
val value = bar.value
emit(ChatEventDto.MessageUpdated("ses_test", message("tail")))
drainScroll()
assertEquals(value, bar.value)
}
fun `test session update preserves middle scroll position`() {
showMessages()
fillTranscript(24)
val bar = scrollBar()
setValue(bar, bottom(bar) / 2)
val value = bar.value
emit(ChatEventDto.MessageUpdated("ses_test", message("tail")))
drainScroll()
assertEquals(value, bar.value)
}
fun `test user scroll between updates disables following`() {
showMessages()
fillTranscript(24)
val bar = scrollBar()
setBottom(bar)
emit(ChatEventDto.MessageUpdated("ses_test", message("tail1")))
drainScroll()
assertBottom(bar)
setValue(bar, bottom(bar) / 2)
val value = bar.value
emit(ChatEventDto.MessageUpdated("ses_test", message("tail2")))
drainScroll()
assertEquals(value, bar.value)
}
fun `test user scroll cancels pending follow`() {
showMessages()
fillTranscript(24)
val bar = scrollBar()
setBottom(bar)
emit(ChatEventDto.MessageUpdated("ses_test", message("tail_pending")), flush = false)
forceFlushWithoutDispatch()
setValue(bar, bottom(bar) / 2)
val value = bar.value
drainScroll()
assertEquals(value, bar.value)
}
fun `test stale follow does not override later non follow`() {
showMessages()
fillTranscript(24)
val bar = scrollBar()
setBottom(bar)
emit(ChatEventDto.MessageUpdated("ses_test", message("tail_stale1")), flush = false)
forceFlushWithoutDispatch()
setValue(bar, bottom(bar) / 2)
val value = bar.value
emit(ChatEventDto.MessageUpdated("ses_test", message("tail_stale2")))
drainScroll()
assertEquals(value, bar.value)
}
fun `test user returning to bottom between updates resumes following`() {
showMessages()
fillTranscript(24)
val bar = scrollBar()
setValue(bar, bottom(bar) / 2)
val value = bar.value
emit(ChatEventDto.MessageUpdated("ses_test", message("tail1")))
drainScroll()
assertEquals(value, bar.value)
setBottom(bar)
emit(ChatEventDto.MessageUpdated("ses_test", message("tail2")))
drainScroll()
assertBottom(bar)
}
fun `test part delta follows bottom after height growth`() {
showMessages()
fillTranscript(24)
val bar = scrollBar()
val id = "stream_bottom"
emit(ChatEventDto.MessageUpdated("ses_test", message(id)), flush = false)
emit(ChatEventDto.PartUpdated("ses_test", part("stream_part", id, "text", "start\n")), flush = false)
forceFlush()
setBottom(bar)
repeat(40) { i ->
emit(ChatEventDto.PartDelta("ses_test", id, "stream_part", "text", "line $i\n"), flush = false)
}
forceFlush()
drainScroll()
assertBottom(bar)
assertFalse(jumpButton().isVisible)
}
fun `test part delta preserves middle scroll position`() {
showMessages()
fillTranscript(24)
val bar = scrollBar()
val id = "stream_middle"
emit(ChatEventDto.MessageUpdated("ses_test", message(id)), flush = false)
emit(ChatEventDto.PartUpdated("ses_test", part("stream_part", id, "text", "start\n")), flush = false)
forceFlush()
setValue(bar, bottom(bar) / 2)
val value = bar.value
repeat(40) { i ->
emit(ChatEventDto.PartDelta("ses_test", id, "stream_part", "text", "line $i\n"), flush = false)
}
forceFlush()
drainScroll()
assertEquals(value, bar.value)
}
fun `test batched update samples scroll once before model changes`() {
showMessages()
fillTranscript(24)
val bar = scrollBar()
setValue(bar, bottom(bar) / 2)
val value = bar.value
emit(ChatEventDto.MessageUpdated("ses_test", message("batch")), flush = false)
emit(ChatEventDto.PartUpdated("ses_test", part("part", "batch", "text", "hello")), flush = false)
forceFlush()
drainScroll()
assertEquals(value, bar.value)
}
fun `test state changes do not force scroll when user is in middle`() {
showMessages()
fillTranscript(24)
val bar = scrollBar()
setValue(bar, bottom(bar) / 2)
val value = bar.value
emit(ChatEventDto.TurnOpen("ses_test"))
drainScroll()
assertEquals(value, bar.value)
}
fun `test scroll button appears only when transcript is away from bottom`() {
showMessages()
fillTranscript(24)
val button = jumpButton()
val bar = scrollBar()
setBottom(bar)
drainScroll()
assertFalse(button.isVisible)
setValue(bar, bottom(bar) / 2)
drainScroll()
assertTrue(button.isVisible)
setBottom(bar)
drainScroll()
assertFalse(button.isVisible)
}
fun `test scroll button scrolls transcript to bottom`() {
showMessages()
fillTranscript(24)
val button = jumpButton()
val bar = scrollBar()
setValue(bar, bottom(bar) / 2)
drainScroll()
assertTrue(button.isVisible)
click(button)
drainScroll()
assertBottom(bar)
assertFalse(button.isVisible)
}
fun `test scroll button remains hidden outside transcript body`() {
val button = jumpButton()
settle()
layout()
assertFalse(button.isVisible)
}
fun `test history load follows initially empty transcript`() {
rpc.history.addAll(history(24))
ui = newUi(id = "ses_test")
settle()
drainScroll()
assertBottom(scrollBar())
}
fun `test recovered state after history preserves user scroll position`() {
rpc.history.addAll(history(24))
rpc.statuses.value = mapOf("ses_test" to SessionStatusDto("busy"))
ui = newUi(id = "ses_test")
settle()
drainScroll()
val bar = scrollBar()
assertBottom(bar)
setValue(bar, bottom(bar) / 2)
val value = bar.value
emit(ChatEventDto.TurnOpen("ses_test"))
drainScroll()
assertEquals(value, bar.value)
}
fun `test existing session scrolls after recovered dock layout`() {
rpc.history.addAll(history(24))
rpc.pendingPermissionList.add(PermissionRequestDto("perm_pending", "ses_test", "edit", listOf("*.kt")))
ui = newUi(id = "ses_test")
settle()
drainScroll()
assertBottom(scrollBar())
}
fun `test replayed event during existing session open cannot cancel initial bottom`() {
val gate = CompletableDeferred<Unit>()
rpc.historyGate = gate
rpc.history.addAll(history(24))
ui = newUi(id = "ses_test")
emit(ChatEventDto.MessageUpdated("ses_test", message("replay")), flush = false)
gate.complete(Unit)
settle()
drainScroll()
assertBottom(scrollBar())
}
fun `test existing session waits for panel layout before initial bottom scroll`() {
rpc.history.addAll(history(24))
ui = newUi(id = "ses_test")
ui.setSize(0, 0)
settle()
ui.setSize(800, 600)
drainScroll()
assertBottom(scrollBar())
}
fun `test existing session scroll waits through deferred transcript revalidation`() {
rpc.history.addAll(history(24))
ui = newUi(id = "ses_test")
settle()
scrollView()?.preferredSize
com.intellij.openapi.application.ApplicationManager.getApplication().invokeLater {
scrollView()?.revalidate()
}
drainScroll()
assertBottom(scrollBar())
}
fun `test scroll owns the session viewport`() {
settle()
assertSame(scrollComponent(), scrollView()?.parent?.parent)
assertFalse(scrollView() is SessionMessageListPanel)
}
}
@@ -1,9 +1,5 @@
package ai.kilocode.client.session
import ai.kilocode.client.app.KiloAppService
import ai.kilocode.client.app.KiloSessionService
import ai.kilocode.client.app.KiloWorkspaceService
import ai.kilocode.client.app.Workspace
import ai.kilocode.client.session.model.Permission
import ai.kilocode.client.session.model.PermissionMeta
import ai.kilocode.client.session.model.Question
@@ -17,75 +13,11 @@ import ai.kilocode.client.session.ui.prompt.PromptPanel
import ai.kilocode.client.session.ui.QuestionPanel
import ai.kilocode.client.session.ui.SessionMessageListPanel
import ai.kilocode.client.session.ui.SessionRootPanel
import ai.kilocode.client.session.update.SessionController
import ai.kilocode.client.session.update.SessionControllerEvent
import ai.kilocode.client.testing.FakeAppRpcApi
import ai.kilocode.client.testing.FakeSessionRpcApi
import ai.kilocode.client.testing.FakeWorkspaceRpcApi
import ai.kilocode.rpc.dto.KiloAppStateDto
import ai.kilocode.rpc.dto.KiloAppStatusDto
import ai.kilocode.rpc.dto.KiloWorkspaceStateDto
import ai.kilocode.rpc.dto.KiloWorkspaceStatusDto
import ai.kilocode.rpc.dto.MessageDto
import ai.kilocode.rpc.dto.MessageTimeDto
import ai.kilocode.rpc.dto.MessageWithPartsDto
import ai.kilocode.rpc.dto.PartDto
import ai.kilocode.rpc.dto.SessionDto
import ai.kilocode.rpc.dto.SessionTimeDto
import ai.kilocode.rpc.dto.ChatEventDto
import com.intellij.testFramework.fixtures.BasePlatformTestCase
import com.intellij.util.ui.JBUI
import com.intellij.ui.components.JBScrollPane
import kotlinx.coroutines.CoroutineScope
import kotlinx.coroutines.SupervisorJob
import kotlinx.coroutines.cancel
import kotlinx.coroutines.delay
import kotlinx.coroutines.runBlocking
import javax.swing.JButton
import javax.swing.JScrollBar
import javax.swing.JLayeredPane
@Suppress("UnstableApiUsage")
class SessionUiLayoutTest : BasePlatformTestCase() {
private lateinit var scope: CoroutineScope
private lateinit var sessions: KiloSessionService
private lateinit var app: KiloAppService
private lateinit var workspaces: KiloWorkspaceService
private lateinit var rpc: FakeSessionRpcApi
private lateinit var workspace: Workspace
private lateinit var ui: SessionUi
override fun setUp() {
super.setUp()
scope = CoroutineScope(SupervisorJob())
rpc = FakeSessionRpcApi()
val appRpc = FakeAppRpcApi().also {
it.state.value = KiloAppStateDto(KiloAppStatusDto.READY)
}
val workspaceRpc = FakeWorkspaceRpcApi().also {
it.state.value = KiloWorkspaceStateDto(status = KiloWorkspaceStatusDto.READY)
}
sessions = KiloSessionService(project, scope, rpc)
app = KiloAppService(scope, appRpc)
workspaces = KiloWorkspaceService(scope, workspaceRpc)
workspace = workspaces.workspace("/test")
ui = SessionUi(project, workspace, sessions, app, scope, displayMs = 0).apply {
setSize(800, 600)
}
layout()
}
override fun tearDown() {
try {
scope.cancel()
} finally {
super.tearDown()
}
}
class SessionUiLayoutTest : SessionUiTestBase() {
fun `test root contains content and overlay layers`() {
val root = find<SessionRootPanel>(ui)
@@ -172,43 +104,37 @@ class SessionUiLayoutTest : BasePlatformTestCase() {
fun `test empty and message bodies share the same scroll pane`() {
settle()
val scroll = find<JBScrollPane>(ui)
val pane = scrollComponent()
val empty = find<EmptySessionPanel>(ui)
assertSame(empty, scroll.viewport.view)
assertSame(empty, scrollView())
com.intellij.openapi.application.ApplicationManager.getApplication().invokeAndWait {
controller().prompt("hello")
}
layout()
assertSame(scroll, find<SessionMessageListPanel>(ui).parent.parent)
assertSame(find<SessionMessageListPanel>(ui), scroll.viewport.view)
assertSame(pane, find<SessionMessageListPanel>(ui).parent.parent)
assertSame(find<SessionMessageListPanel>(ui), scrollView())
}
fun `test new session starts with loading body`() {
ui = SessionUi(project, workspace, sessions, app, scope, displayMs = 1_000).apply {
setSize(800, 600)
}
ui = newUi(displayMs = 1_000)
assertFalse(find<JBScrollPane>(ui).viewport.view is EmptySessionPanel)
assertFalse(scrollView() is EmptySessionPanel)
}
fun `test action-created new session starts blank`() {
ui = SessionUi(project, workspace, sessions, app, scope, displayMs = 1_000, loading = false).apply {
setSize(800, 600)
}
ui = newUi(displayMs = 1_000, loading = false)
assertFalse(find<JBScrollPane>(ui).viewport.view is EmptySessionPanel)
assertFalse(find<JBScrollPane>(ui).viewport.view is SessionMessageListPanel)
assertFalse(scrollView() is EmptySessionPanel)
assertFalse(scrollView() is SessionMessageListPanel)
}
fun `test clicking recent session calls opener`() {
val opened = mutableListOf<String>()
rpc.recent.add(session("ses_1"))
ui = SessionUi(project, workspace, sessions, app, scope, displayMs = 0, open = { opened.add(it.id) }).apply {
setSize(800, 600)
}
ui = newUi(open = { opened.add(it.id) })
settle()
layout()
@@ -218,346 +144,46 @@ class SessionUiLayoutTest : BasePlatformTestCase() {
}
fun `test existing session id loads history and shows message body`() {
rpc.history.add(MessageWithPartsDto(message("msg1"), emptyList()))
rpc.history.addAll(history(1))
ui = SessionUi(project, workspace, sessions, app, scope, id = "ses_test", displayMs = 0).apply {
setSize(800, 600)
}
ui = newUi(id = "ses_test")
settle()
assertSame(find<SessionMessageListPanel>(ui), find<JBScrollPane>(ui).viewport.view)
assertSame(find<SessionMessageListPanel>(ui), scrollView())
}
fun `test new session keeps loading body before recents delay`() {
rpc.recentGate = kotlinx.coroutines.CompletableDeferred()
ui = SessionUi(project, workspace, sessions, app, scope, displayMs = 1_000).apply {
setSize(800, 600)
}
ui = newUi(displayMs = 1_000)
settleShort(100)
assertFalse(find<JBScrollPane>(ui).viewport.view is EmptySessionPanel)
assertFalse(scrollView() is EmptySessionPanel)
}
fun `test slow recents switch to loading body only after progress event`() {
rpc.recentGate = kotlinx.coroutines.CompletableDeferred()
rpc.recent.add(session("ses_1"))
ui = SessionUi(project, workspace, sessions, app, scope, displayMs = 50).apply {
setSize(800, 600)
}
ui = newUi(displayMs = 50)
settleShort(20)
assertFalse(find<JBScrollPane>(ui).viewport.view is EmptySessionPanel)
assertFalse(scrollView() is EmptySessionPanel)
settleShort(80)
assertFalse(find<JBScrollPane>(ui).viewport.view is EmptySessionPanel)
assertFalse(scrollView() is EmptySessionPanel)
rpc.recentGate!!.complete(Unit)
settle()
val panel = find<EmptySessionPanel>(ui)
assertSame(panel, find<JBScrollPane>(ui).viewport.view)
assertSame(panel, scrollView())
assertEquals(1, panel.recentCount())
}
fun `test session update follows when transcript is at bottom`() {
showMessages()
fillTranscript(24)
val bar = scrollBar()
setBottom(bar)
emit(ChatEventDto.MessageUpdated("ses_test", message("tail")))
drainScroll()
assertBottom(bar)
}
fun `test session update follows when transcript is near bottom threshold`() {
showMessages()
fillTranscript(24)
val bar = scrollBar()
val threshold = JBUI.scale(32)
if (bottom(bar) <= threshold) {
fillTranscript(24, start = 24)
}
setValue(bar, bottom(bar) - threshold + 1)
emit(ChatEventDto.MessageUpdated("ses_test", message("tail")))
drainScroll()
assertBottom(bar)
}
fun `test session update preserves position outside bottom threshold`() {
showMessages()
fillTranscript(24)
val bar = scrollBar()
val threshold = JBUI.scale(32)
setValue(bar, bottom(bar) - threshold - 8)
val value = bar.value
emit(ChatEventDto.MessageUpdated("ses_test", message("tail")))
drainScroll()
assertEquals(value, bar.value)
}
fun `test session update preserves middle scroll position`() {
showMessages()
fillTranscript(24)
val bar = scrollBar()
setValue(bar, bottom(bar) / 2)
val value = bar.value
emit(ChatEventDto.MessageUpdated("ses_test", message("tail")))
drainScroll()
assertEquals(value, bar.value)
}
fun `test user scroll between updates disables following`() {
showMessages()
fillTranscript(24)
val bar = scrollBar()
setBottom(bar)
emit(ChatEventDto.MessageUpdated("ses_test", message("tail1")))
drainScroll()
assertBottom(bar)
setValue(bar, bottom(bar) / 2)
val value = bar.value
emit(ChatEventDto.MessageUpdated("ses_test", message("tail2")))
drainScroll()
assertEquals(value, bar.value)
}
fun `test user returning to bottom between updates resumes following`() {
showMessages()
fillTranscript(24)
val bar = scrollBar()
setValue(bar, bottom(bar) / 2)
val value = bar.value
emit(ChatEventDto.MessageUpdated("ses_test", message("tail1")))
drainScroll()
assertEquals(value, bar.value)
setBottom(bar)
emit(ChatEventDto.MessageUpdated("ses_test", message("tail2")))
drainScroll()
assertBottom(bar)
}
fun `test batched update samples scroll once before model changes`() {
showMessages()
fillTranscript(24)
val bar = scrollBar()
setValue(bar, bottom(bar) / 2)
val value = bar.value
emit(ChatEventDto.MessageUpdated("ses_test", message("batch")), flush = false)
emit(ChatEventDto.PartUpdated("ses_test", part("part", "batch", "text", "hello")), flush = false)
forceFlush()
drainScroll()
assertEquals(value, bar.value)
}
fun `test state changes do not force scroll when user is in middle`() {
showMessages()
fillTranscript(24)
val bar = scrollBar()
setValue(bar, bottom(bar) / 2)
val value = bar.value
emit(ChatEventDto.TurnOpen("ses_test"))
drainScroll()
assertEquals(value, bar.value)
}
fun `test scroll button appears only when transcript is away from bottom`() {
showMessages()
fillTranscript(24)
val button = jumpButton()
val bar = scrollBar()
setBottom(bar)
drainScroll()
assertFalse(button.isVisible)
setValue(bar, bottom(bar) / 2)
drainScroll()
assertTrue(button.isVisible)
setBottom(bar)
drainScroll()
assertFalse(button.isVisible)
}
fun `test scroll button scrolls transcript to bottom`() {
showMessages()
fillTranscript(24)
val button = jumpButton()
val bar = scrollBar()
setValue(bar, bottom(bar) / 2)
drainScroll()
assertTrue(button.isVisible)
button.doClick()
drainScroll()
assertBottom(bar)
assertFalse(button.isVisible)
}
fun `test scroll button remains hidden outside transcript body`() {
val button = jumpButton()
settle()
layout()
assertFalse(button.isVisible)
}
fun `test history load follows initially empty transcript`() {
rpc.history.addAll(history(24))
ui = SessionUi(project, workspace, sessions, app, scope, id = "ses_test", displayMs = 0).apply {
setSize(800, 600)
}
settle()
drainScroll()
assertBottom(scrollBar())
}
fun `test recovered state after history preserves user scroll position`() {
rpc.history.addAll(history(24))
rpc.statuses.value = mapOf("ses_test" to ai.kilocode.rpc.dto.SessionStatusDto("busy"))
ui = SessionUi(project, workspace, sessions, app, scope, id = "ses_test", displayMs = 0).apply {
setSize(800, 600)
}
settle()
drainScroll()
val bar = scrollBar()
setValue(bar, bottom(bar) / 2)
val value = bar.value
emit(ChatEventDto.TurnOpen("ses_test"))
drainScroll()
assertEquals(value, bar.value)
}
private fun layout() {
ui.doLayout()
val root = find<SessionRootPanel>(ui)
root.doLayout()
root.content.doLayout()
find<PromptPanel>(ui).parent.doLayout()
find<JBScrollPane>(ui).doLayout()
(find<JBScrollPane>(ui).viewport.view as? java.awt.Container)?.doLayout()
}
private fun settle() = runBlocking {
repeat(5) {
delay(100)
com.intellij.util.ui.UIUtil.dispatchAllInvocationEvents()
}
}
private fun settleShort(ms: Long) = runBlocking {
delay(ms)
com.intellij.util.ui.UIUtil.dispatchAllInvocationEvents()
}
private fun showConnection() {
find<ConnectionPanel>(ui).onEvent(SessionControllerEvent.ConnectionChanged.ShowConnecting)
}
private fun showMessages() {
controller().prompt("hello")
settle()
layout()
}
private fun fillTranscript(count: Int, start: Int = 0) {
repeat(count) { offset ->
val i = start + offset
val id = "msg_$i"
emit(ChatEventDto.MessageUpdated("ses_test", message(id)), flush = false)
emit(ChatEventDto.PartUpdated("ses_test", part("part_$i", id, "text", text(i))), flush = false)
}
settleShort(100)
forceFlush()
drainScroll()
}
private fun emit(event: ChatEventDto, flush: Boolean = true) {
runBlocking { rpc.events.emit(event) }
if (flush) {
settleShort(20)
forceFlush()
}
}
private fun forceFlush() {
controller().flushEvents()
com.intellij.util.ui.UIUtil.dispatchAllInvocationEvents()
}
private fun drainScroll() {
repeat(4) {
layout()
com.intellij.util.ui.UIUtil.dispatchAllInvocationEvents()
}
}
private fun scrollBar(): JScrollBar = find<JBScrollPane>(ui).verticalScrollBar
private fun jumpButton(): JButton {
return find<SessionRootPanel>(ui).overlay.components.single() as JButton
}
private fun bottom(bar: JScrollBar): Int = (bar.maximum - bar.visibleAmount).coerceAtLeast(0)
private fun setBottom(bar: JScrollBar) {
setValue(bar, bottom(bar))
}
private fun setValue(bar: JScrollBar, value: Int) {
bar.value = value.coerceIn(bar.minimum, bottom(bar))
}
private fun assertBottom(bar: JScrollBar) {
assertTrue("value=${bar.value} bottom=${bottom(bar)} max=${bar.maximum} visible=${bar.visibleAmount}", bar.value >= bottom(bar) - 1)
}
private inline fun <reified T> find(root: java.awt.Container): T {
return find(root, T::class.java) ?: error("missing ${T::class.java.simpleName}")
}
private fun <T> find(root: java.awt.Container, cls: Class<T>): T? {
if (cls.isInstance(root)) return cls.cast(root)
for (child in root.components) {
if (cls.isInstance(child)) return cls.cast(child)
if (child is java.awt.Container) {
val item = find(child, cls)
if (item != null) return item
}
}
return null
}
private fun controller(): SessionController {
val field = SessionUi::class.java.getDeclaredField("controller")
field.isAccessible = true
return field.get(ui) as SessionController
}
private fun questionStateChanged() = SessionState.AwaitingQuestion(
Question(
id = "q1",
@@ -583,35 +209,4 @@ class SessionUiLayoutTest : BasePlatformTestCase() {
meta = PermissionMeta(raw = emptyMap()),
)
)
private fun session(id: String) = SessionDto(
id = id,
projectID = "prj",
directory = "/test",
title = "Recent $id",
version = "1",
time = SessionTimeDto(created = 1.0, updated = 2.0),
)
private fun message(id: String) = MessageDto(
id = id,
sessionID = "ses_test",
role = "user",
time = MessageTimeDto(created = 0.0),
)
private fun part(id: String, mid: String, type: String, text: String? = null) = PartDto(
id = id,
sessionID = "ses_test",
messageID = mid,
type = type,
text = text,
)
private fun history(count: Int): List<MessageWithPartsDto> = List(count) { i ->
val id = "hist_$i"
MessageWithPartsDto(message(id), listOf(part("hist_part_$i", id, "text", text(i))))
}
private fun text(i: Int): String = "line $i\n".repeat(12)
}
@@ -0,0 +1,232 @@
package ai.kilocode.client.session
import ai.kilocode.client.app.KiloAppService
import ai.kilocode.client.app.KiloSessionService
import ai.kilocode.client.app.KiloWorkspaceService
import ai.kilocode.client.app.Workspace
import ai.kilocode.client.session.ui.SessionRootPanel
import ai.kilocode.client.session.ui.prompt.PromptPanel
import ai.kilocode.client.session.update.SessionController
import ai.kilocode.client.testing.FakeAppRpcApi
import ai.kilocode.client.testing.FakeSessionRpcApi
import ai.kilocode.client.testing.FakeWorkspaceRpcApi
import ai.kilocode.rpc.dto.ChatEventDto
import ai.kilocode.rpc.dto.KiloAppStateDto
import ai.kilocode.rpc.dto.KiloAppStatusDto
import ai.kilocode.rpc.dto.KiloWorkspaceStateDto
import ai.kilocode.rpc.dto.KiloWorkspaceStatusDto
import ai.kilocode.rpc.dto.MessageDto
import ai.kilocode.rpc.dto.MessageTimeDto
import ai.kilocode.rpc.dto.MessageWithPartsDto
import ai.kilocode.rpc.dto.PartDto
import ai.kilocode.rpc.dto.SessionDto
import ai.kilocode.rpc.dto.SessionTimeDto
import com.intellij.testFramework.fixtures.BasePlatformTestCase
import com.intellij.util.ui.UIUtil
import kotlinx.coroutines.CoroutineScope
import kotlinx.coroutines.SupervisorJob
import kotlinx.coroutines.cancel
import kotlinx.coroutines.delay
import kotlinx.coroutines.runBlocking
import java.awt.Container
import java.awt.event.MouseEvent
import javax.swing.JLabel
import javax.swing.JComponent
import javax.swing.JScrollBar
@Suppress("UnstableApiUsage")
abstract class SessionUiTestBase : BasePlatformTestCase() {
protected lateinit var scope: CoroutineScope
protected lateinit var sessions: KiloSessionService
protected lateinit var app: KiloAppService
protected lateinit var workspaces: KiloWorkspaceService
protected lateinit var rpc: FakeSessionRpcApi
protected lateinit var workspace: Workspace
protected lateinit var ui: SessionUi
override fun setUp() {
super.setUp()
scope = CoroutineScope(SupervisorJob())
rpc = FakeSessionRpcApi()
val appRpc = FakeAppRpcApi().also {
it.state.value = KiloAppStateDto(KiloAppStatusDto.READY)
}
val workspaceRpc = FakeWorkspaceRpcApi().also {
it.state.value = KiloWorkspaceStateDto(status = KiloWorkspaceStatusDto.READY)
}
sessions = KiloSessionService(project, scope, rpc)
app = KiloAppService(scope, appRpc)
workspaces = KiloWorkspaceService(scope, workspaceRpc)
workspace = workspaces.workspace("/test")
ui = newUi()
layout()
}
override fun tearDown() {
try {
scope.cancel()
} finally {
super.tearDown()
}
}
protected fun newUi(
id: String? = null,
displayMs: Long = 0,
loading: Boolean = id == null,
open: (SessionDto) -> Unit = {},
): SessionUi {
return SessionUi(project, workspace, sessions, app, scope, id = id, displayMs = displayMs, loading = loading, open = open).apply {
setSize(800, 600)
}
}
protected fun layout() {
ui.doLayout()
val root = find<SessionRootPanel>(ui)
root.doLayout()
root.content.doLayout()
find<PromptPanel>(ui).parent.doLayout()
scrollComponent().doLayout()
(scrollView() as? Container)?.doLayout()
}
protected fun settle() = runBlocking {
repeat(5) {
delay(100)
UIUtil.dispatchAllInvocationEvents()
}
}
protected fun settleShort(ms: Long) = runBlocking {
delay(ms)
UIUtil.dispatchAllInvocationEvents()
}
protected fun showMessages() {
controller().prompt("hello")
settle()
layout()
}
protected fun fillTranscript(count: Int, start: Int = 0) {
repeat(count) { offset ->
val i = start + offset
val id = "msg_$i"
emit(ChatEventDto.MessageUpdated("ses_test", message(id)), flush = false)
emit(ChatEventDto.PartUpdated("ses_test", part("part_$i", id, "text", text(i))), flush = false)
}
settleShort(100)
forceFlush()
drainScroll()
}
protected fun emit(event: ChatEventDto, flush: Boolean = true) {
runBlocking { rpc.events.emit(event) }
if (flush) {
settleShort(20)
forceFlush()
}
}
protected fun forceFlush() {
controller().flushEvents()
UIUtil.dispatchAllInvocationEvents()
}
protected fun forceFlushWithoutDispatch() {
controller().flushEvents()
}
protected fun drainScroll() {
repeat(4) {
layout()
UIUtil.dispatchAllInvocationEvents()
}
}
private fun scroll(): SessionScroll = ui.scroll
protected fun scrollComponent(): JComponent = scroll().component
protected fun scrollView(): JComponent? = scroll().view
protected fun scrollBar(): JScrollBar = scroll().bar
protected fun jumpButton(): JLabel = scroll().jump
protected fun click(label: JLabel) {
val event = MouseEvent(label, MouseEvent.MOUSE_CLICKED, System.currentTimeMillis(), 0, 1, 1, 1, false)
for (listener in label.mouseListeners) listener.mouseClicked(event)
}
protected fun bottom(bar: JScrollBar): Int = (bar.maximum - bar.visibleAmount).coerceAtLeast(0)
protected fun setBottom(bar: JScrollBar) {
setValue(bar, bottom(bar))
}
protected fun setValue(bar: JScrollBar, value: Int) {
bar.value = value.coerceIn(bar.minimum, bottom(bar))
}
protected fun assertBottom(bar: JScrollBar) {
assertTrue("value=${bar.value} bottom=${bottom(bar)} max=${bar.maximum} visible=${bar.visibleAmount}", bar.value >= bottom(bar) - 1)
}
protected inline fun <reified T> find(root: Container): T {
return find(root, T::class.java) ?: error("missing ${T::class.java.simpleName}")
}
protected fun <T> find(root: Container, cls: Class<T>): T? {
if (cls.isInstance(root)) return cls.cast(root)
for (child in root.components) {
if (cls.isInstance(child)) return cls.cast(child)
if (child is Container) {
val item = find(child, cls)
if (item != null) return item
}
}
return null
}
protected fun controller(): SessionController {
val field = SessionUi::class.java.getDeclaredField("controller")
field.isAccessible = true
return field.get(ui) as SessionController
}
protected fun session(id: String) = SessionDto(
id = id,
projectID = "prj",
directory = "/test",
title = "Recent $id",
version = "1",
time = SessionTimeDto(created = 1.0, updated = 2.0),
)
protected fun message(id: String) = MessageDto(
id = id,
sessionID = "ses_test",
role = "user",
time = MessageTimeDto(created = 0.0),
)
protected fun part(id: String, mid: String, type: String, text: String? = null) = PartDto(
id = id,
sessionID = "ses_test",
messageID = mid,
type = type,
text = text,
)
protected fun history(count: Int): List<MessageWithPartsDto> = List(count) { i ->
val id = "hist_$i"
MessageWithPartsDto(message(id), listOf(part("hist_part_$i", id, "text", text(i))))
}
protected fun text(i: Int): String = "line $i\n".repeat(12)
}
@@ -42,6 +42,7 @@ class FakeSessionRpcApi : KiloSessionRpcApi {
/** Message history returned by [messages]. */
val history = mutableListOf<MessageWithPartsDto>()
var historyGate: CompletableDeferred<Unit>? = null
/** Recent sessions returned by [recent]. */
val recent = mutableListOf<SessionDto>()
@@ -135,6 +136,7 @@ class FakeSessionRpcApi : KiloSessionRpcApi {
override suspend fun messages(id: String, directory: String): List<MessageWithPartsDto> {
assertNotEdt("messages")
historyGate?.await()
return history.toList()
}