fix(jetbrains): harden revert operation state

This commit is contained in:
kirillk
2026-07-10 11:38:08 -04:00
parent c1b206b161
commit 3a26cb844b
3 changed files with 109 additions and 48 deletions
@@ -104,6 +104,7 @@ class SessionController(
private data class OrganizationTarget(val org: String?)
private data class Followup(val dir: String, val time: Long)
private data class Pref(val agent: String?, val model: String?, val variants: List<String>, val variant: String?, val reset: Boolean)
private data class RevertOp(val key: Long)
private data class Dispatch(
val kind: String,
val source: String,
@@ -151,6 +152,8 @@ class SessionController(
private var eventJob: Job? = null
private var drainJob: Job? = null
private var revertJob: Job? = null
private var revertOp: RevertOp? = null
private var revertSeq = 0L
private var revertWatchdog: UiTimer? = null
private var creating: CompletableDeferred<String?>? = null
private val childJobs: MutableMap<String, Job> = mutableMapOf()
@@ -276,6 +279,7 @@ class SessionController(
private fun dispatch(data: Dispatch, send: suspend (String) -> Unit) {
assertEdt()
if (revertOp != null) return
val props = data.props + if (data.kind == "command") slashProps() else emptyMap()
capture("Conversation Send Clicked", sessionProps(sid ?: ref?.key) + mapOf(
"source" to data.source,
@@ -348,7 +352,7 @@ class SessionController(
fun abort() {
assertEdt()
LOG.debug { "${ChatLogSummary.sid(sid ?: ref?.key ?: "pending")} kind=abort" }
if (model.state is SessionState.Reverting) {
if (revertOp != null) {
cancelRevert()
return
}
@@ -388,7 +392,7 @@ class SessionController(
fun compact() {
assertEdt()
val id = sid ?: return
if (model.state.isBusy()) return
if (revertOp != null || model.state.isBusy()) return
if (model.isEmpty()) return
val parsed = model.model?.let(::parseModel) ?: return
val sel = ModelSelectionDto(parsed.first, parsed.second)
@@ -421,18 +425,17 @@ class SessionController(
return
}
val state = model.state
if (state is SessionState.Reverting) return
if (revertOp != null) return
val busy = state.isBusy()
LOG.info(
"${ChatLogSummary.sid(id)} kind=revert clicked=true message=$message " +
"part=${part ?: "none"} busy=$busy",
)
model.setState(SessionState.Reverting(
val op = beginReverting(
KiloBundle.message("session.status.rollingback"),
SessionState.Reverting.Kind.ROLLBACK,
message,
))
startRevertWatchdog()
) ?: return
revertJob = cs.launch {
try {
if (busy) {
@@ -444,13 +447,13 @@ class SessionController(
capture("Session Rollback", sessionProps(id))
synchronizeFromDisk(id, "revert")
LOG.info("${ChatLogSummary.sid(id)} kind=revert ok=true")
edt { clearReverting() }
edt { clearReverting(op) }
} catch (e: CancellationException) {
throw e
} catch (e: Exception) {
capture("Session Error", sessionProps(id) + mapOf("context" to "revert", "errorClass" to e::class.java.name))
LOG.warn("${ChatLogSummary.sid(id)} kind=revert dir=${ChatLogSummary.dir(directory)} failed message=${e.message}", e)
edt { failReverting(e) }
edt { failReverting(op, e) }
}
}
}
@@ -458,24 +461,23 @@ class SessionController(
fun unrevert() {
assertEdt()
val id = sid ?: return
if (model.state is SessionState.Reverting) return
model.setState(SessionState.Reverting(
if (revertOp != null) return
val op = beginReverting(
KiloBundle.message("session.status.redoing"),
SessionState.Reverting.Kind.REDO,
))
startRevertWatchdog()
) ?: return
revertJob = cs.launch {
try {
sessions.unrevert(id, directory)
capture("Session Unrevert", sessionProps(id))
synchronizeFromDisk(id, "unrevert")
edt { clearReverting() }
edt { clearReverting(op) }
} catch (e: CancellationException) {
throw e
} catch (e: Exception) {
capture("Session Error", sessionProps(id) + mapOf("context" to "unrevert", "errorClass" to e::class.java.name))
LOG.warn("${ChatLogSummary.sid(id)} kind=unrevert dir=${ChatLogSummary.dir(directory)} failed message=${e.message}", e)
edt { failReverting(e) }
edt { failReverting(op, e) }
}
}
}
@@ -503,22 +505,28 @@ class SessionController(
private fun redoTo(message: String) {
assertEdt()
val id = sid ?: return
if (model.state is SessionState.Reverting) return
model.setState(SessionState.Reverting(
val state = model.state
if (revertOp != null) return
val busy = state.isBusy()
val op = beginReverting(
KiloBundle.message("session.status.redoing"),
SessionState.Reverting.Kind.REDO,
message,
))
startRevertWatchdog()
) ?: return
revertJob = cs.launch {
try {
if (busy) {
LOG.info("${ChatLogSummary.sid(id)} kind=redo abort=true reason=busy")
sessions.abort(id, directory)
LOG.info("${ChatLogSummary.sid(id)} kind=redo abort=true ok=true")
}
sessions.revert(id, directory, message, null)
synchronizeFromDisk(id, "redo")
edt { clearReverting() }
edt { clearReverting(op) }
} catch (e: CancellationException) {
throw e
} catch (e: Exception) {
edt { failReverting(e) }
edt { failReverting(op, e) }
}
}
}
@@ -531,13 +539,13 @@ class SessionController(
fun cancelRevert() {
assertEdt()
if (model.state !is SessionState.Reverting) return
LOG.info("${ChatLogSummary.sid(sid ?: "?")} kind=revert cancelled=true")
stopRevertWatchdog()
revertJob?.cancel()
revertJob = null
sid?.let { capture("Session Revert Cancelled", sessionProps(it)) }
model.setState(SessionState.Idle)
if (revertOp == null) return
LOG.info("${ChatLogSummary.sid(sid ?: "?")} kind=revert cancelRequested=true")
sid?.let { capture("Session Revert Cancel Requested", sessionProps(it)) }
val state = model.state
if (state is SessionState.Reverting) {
model.setState(state.copy(text = KiloBundle.message("session.status.operation.finishing")))
}
}
private fun synchronizeFromDisk(id: String, kind: String) {
@@ -1458,6 +1466,7 @@ class SessionController(
}
private fun status(dto: SessionStatusDto) {
if (revertOp != null) return
val state = when (dto.type) {
"idle" -> {
val current = model.state
@@ -1484,11 +1493,21 @@ class SessionController(
model.setState(state)
}
private fun startRevertWatchdog() {
private fun beginReverting(text: String, kind: SessionState.Reverting.Kind, message: String? = null): RevertOp? {
assertEdt()
if (revertOp != null) return null
val op = RevertOp(++revertSeq)
revertOp = op
model.setState(SessionState.Reverting(text, kind, message))
startRevertWatchdog(op)
return op
}
private fun startRevertWatchdog(op: RevertOp) {
assertEdt()
stopRevertWatchdog()
val ms = revertTimeoutMs.coerceIn(1, Int.MAX_VALUE.toLong()).toInt()
revertWatchdog = timers.timer(ms, repeats = false) { onRevertTimeout() }.also { it.start() }
revertWatchdog = timers.timer(ms, repeats = false) { onRevertTimeout(op) }.also { it.start() }
}
private fun stopRevertWatchdog() {
@@ -1496,32 +1515,38 @@ class SessionController(
revertWatchdog = null
}
private fun onRevertTimeout() {
private fun onRevertTimeout(op: RevertOp) {
assertEdt()
if (model.state !is SessionState.Reverting) return
if (revertOp?.key != op.key) return
LOG.warn("${ChatLogSummary.sid(sid ?: "?")} kind=revert timeout=true after=${revertTimeoutMs}ms")
revertJob?.cancel()
revertJob = null
stopRevertWatchdog()
sid?.let { capture("Session Revert Timeout", sessionProps(it)) }
model.setState(SessionState.Error(KiloBundle.message("session.error.revert.timeout")))
val state = model.state
if (state is SessionState.Reverting) {
model.setState(state.copy(text = KiloBundle.message("session.error.revert.timeout")))
}
}
private fun clearReverting() {
private fun clearReverting(op: RevertOp) {
assertEdt()
if (revertOp?.key != op.key) return
stopRevertWatchdog()
revertJob = null
revertOp = null
if (model.state is SessionState.Reverting) model.setState(SessionState.Idle)
}
private fun failReverting(e: Exception) {
private fun failReverting(op: RevertOp, e: Exception) {
assertEdt()
if (revertOp?.key != op.key) return
stopRevertWatchdog()
revertJob = null
revertOp = null
model.setState(SessionState.Error(e.message ?: KiloBundle.message("session.error.unknown")))
}
private fun idle() {
if (revertOp != null) return
// Treat session.idle as an explicit signal to return to Idle.
// Only apply if we're not in a more specific non-terminal state.
val current = model.state
@@ -2126,6 +2151,7 @@ class SessionController(
revertWatchdog = null
revertJob?.cancel()
revertJob = null
revertOp = null
val callbacks = enhancements.values.toList()
enhancements.clear()
cs.cancel()
@@ -43,7 +43,8 @@ revert.banner.filesNotRestored=Snapshots are off - only the conversation was rev
revert.message.rollback=Rollback to this message
session.status.rollingback=Rolling back\u2026
session.status.redoing=Redoing\u2026
session.error.revert.timeout=Rollback timed out. Please try again.
session.status.operation.finishing=Waiting for the operation to finish\u2026
session.error.revert.timeout=Operation timed out. Waiting for it to finish before continuing.
session.permission.title=Permission required
session.permission.title.subagent=Permission required (subagent)
@@ -66,6 +66,19 @@ class TurnLifecycleTest : SessionControllerTestBase() {
assertTrue(appRpc.telemetry.any { it.event == "Session Redo" })
}
fun `test redo aborts busy session before partial redo`() {
val (m, _, _) = prompted()
seedRevertMessages()
emit(ChatEventDto.SessionUpdated("ses_test", session("ses_test").copy(revert = SessionRevertDto("u1"))))
emit(ChatEventDto.TurnOpen("ses_test"))
edt { m.redo() }
flush()
assertEquals(listOf("ses_test" to "/test"), rpc.aborts)
assertEquals(listOf(FakeSessionRpcApi.RevertCall("ses_test", "/test", "u2", null)), rpc.reverts)
}
fun `test redo at final user message unreverts`() {
val (m, _, _) = prompted()
seedRevertMessages()
@@ -573,6 +586,9 @@ class TurnLifecycleTest : SessionControllerTestBase() {
emit(ChatEventDto.SessionStatusChanged("ses_test", SessionStatusDto("idle")))
assertTrue("status idle must not clear reverting", m.model.state is SessionState.Reverting)
emit(ChatEventDto.SessionStatusChanged("ses_test", SessionStatusDto("retry", message = "retrying", attempt = 1, next = 0L)))
assertTrue("status retry must not clear reverting", m.model.state is SessionState.Reverting)
gate.complete(Unit)
flush()
}
@@ -676,7 +692,7 @@ class TurnLifecycleTest : SessionControllerTestBase() {
}
fun `test cancelRevert cancels in-flight rollback and returns to idle`() {
fun `test cancelRevert waits for in-flight rollback before returning to idle`() {
val (m, _, _) = prompted()
val gate = CompletableDeferred<Unit>()
rpc.revertGate = gate
@@ -688,16 +704,19 @@ class TurnLifecycleTest : SessionControllerTestBase() {
edt { m.cancelRevert() }
settle()
assertTrue(m.model.state is SessionState.Idle)
assertTrue("rpc must not have recorded a revert", rpc.reverts.isEmpty())
assertTrue(appRpc.telemetry.any { it.event == "Session Revert Cancelled" })
val state = m.model.state
assertTrue("expected Reverting, was $state", state is SessionState.Reverting)
assertEquals(KiloBundle.message("session.status.operation.finishing"), (state as SessionState.Reverting).text)
assertTrue("rpc must still be waiting", rpc.reverts.isEmpty())
assertTrue(appRpc.telemetry.any { it.event == "Session Revert Cancel Requested" })
gate.complete(Unit)
flush()
assertEquals(listOf(FakeSessionRpcApi.RevertCall("ses_test", "/test", "msg1", null)), rpc.reverts)
assertTrue(m.model.state is SessionState.Idle)
}
fun `test cancelRevert cancels in-flight redoAll`() {
fun `test cancelRevert waits for in-flight redoAll`() {
val (m, _, _) = prompted()
val gate = CompletableDeferred<Unit>()
rpc.unrevertGate = gate
@@ -709,8 +728,13 @@ class TurnLifecycleTest : SessionControllerTestBase() {
edt { m.cancelRevert() }
settle()
assertTrue(m.model.state is SessionState.Idle)
assertTrue(m.model.state is SessionState.Reverting)
assertTrue(rpc.unreverts.isEmpty())
gate.complete(Unit)
flush()
assertEquals(listOf("ses_test" to "/test"), rpc.unreverts)
assertTrue(m.model.state is SessionState.Idle)
}
fun `test cancelRevert is ignored when not reverting`() {
@@ -720,7 +744,7 @@ class TurnLifecycleTest : SessionControllerTestBase() {
settle()
assertTrue(m.model.state is SessionState.Idle)
assertFalse(appRpc.telemetry.any { it.event == "Session Revert Cancelled" })
assertFalse(appRpc.telemetry.any { it.event == "Session Revert Cancel Requested" })
}
fun `test abort while reverting cancels rollback instead of aborting turn`() {
@@ -735,14 +759,17 @@ class TurnLifecycleTest : SessionControllerTestBase() {
edt { m.abort() }
settle()
assertTrue(m.model.state is SessionState.Idle)
assertTrue(m.model.state is SessionState.Reverting)
assertTrue(rpc.aborts.isEmpty())
assertTrue(rpc.reverts.isEmpty())
gate.complete(Unit)
flush()
assertEquals(listOf(FakeSessionRpcApi.RevertCall("ses_test", "/test", "msg1", null)), rpc.reverts)
assertTrue(m.model.state is SessionState.Idle)
}
fun `test rollback watchdog times out to error`() {
fun `test rollback watchdog times out without opening a second rollback`() {
val (m, _, _) = prompted()
val gate = CompletableDeferred<Unit>()
rpc.revertGate = gate
@@ -754,11 +781,18 @@ class TurnLifecycleTest : SessionControllerTestBase() {
pause(SessionController.REVERT_TIMEOUT_MS + 1)
val state = m.model.state
assertTrue("expected Error, was $state", state is SessionState.Error)
assertEquals(KiloBundle.message("session.error.revert.timeout"), (state as SessionState.Error).message)
assertTrue("expected Reverting, was $state", state is SessionState.Reverting)
assertEquals(KiloBundle.message("session.error.revert.timeout"), (state as SessionState.Reverting).text)
assertTrue(appRpc.telemetry.any { it.event == "Session Revert Timeout" })
edt { m.revert("msg2") }
settle()
assertTrue(rpc.reverts.isEmpty())
gate.complete(Unit)
flush()
assertEquals(listOf(FakeSessionRpcApi.RevertCall("ses_test", "/test", "msg1", null)), rpc.reverts)
assertTrue(m.model.state is SessionState.Idle)
}
fun `test successful rollback stops watchdog`() {