Parcourir la source

fix(core): fail interrupted session steps

Dax Raad il y a 1 mois
Parent
commit
b6553d14e1

+ 6 - 7
packages/core/src/session/runner/llm.ts

@@ -287,19 +287,18 @@ export const layer = Layer.effect(
           }
           if (stream._tag === "Failure" && Cause.hasInterrupts(stream.cause)) yield* FiberSet.clear(toolFibers)
           const settled = yield* restore(awaitToolFibers(toolFibers)).pipe(Effect.exit)
+          const streamInterrupted = stream._tag === "Failure" && Cause.hasInterrupts(stream.cause)
+          const toolsInterrupted = settled._tag === "Failure" && Cause.hasInterrupts(settled.cause)
           if (settled._tag === "Failure" && isQuestionRejected(settled.cause)) {
             yield* FiberSet.clear(toolFibers)
             yield* withPublication(publisher.failUnsettledTools("Tool execution interrupted"))
+            yield* withPublication(publisher.failAssistant("Provider turn interrupted"))
             return yield* Effect.interrupt
           }
-          if (
-            (stream._tag === "Failure" && Cause.hasInterrupts(stream.cause)) ||
-            (settled._tag === "Failure" && Cause.hasInterrupts(settled.cause))
-          ) {
+          if (streamInterrupted || toolsInterrupted) {
             yield* FiberSet.clear(toolFibers)
             yield* withPublication(publisher.failUnsettledTools("Tool execution interrupted"))
-            if (publisher.hasActiveAssistant())
-              yield* withPublication(publisher.failAssistant("Provider turn interrupted"))
+            yield* withPublication(publisher.failAssistant("Provider turn interrupted"))
           }
           if (settled._tag === "Failure" && !Cause.hasInterrupts(settled.cause)) {
             const failure = Cause.squash(settled.cause)
@@ -307,7 +306,7 @@ export const layer = Layer.effect(
             yield* withPublication(publisher.failUnsettledTools(`Tool execution failed: ${message}`))
           }
           const stepSettlement = publisher.stepSettlement()
-          if (stepSettlement && !publisher.hasProviderError()) {
+          if (stepSettlement && !streamInterrupted && !toolsInterrupted && !publisher.hasProviderError()) {
             const endSnapshot = yield* snapshots.capture()
             const files =
               startSnapshot && endSnapshot

+ 22 - 0
packages/core/test/session-runner.test.ts

@@ -371,6 +371,18 @@ const messageTexts = (request: LLMRequest, role: "user" | "system") =>
 const userTexts = (request: LLMRequest) => messageTexts(request, "user")
 const systemTexts = (request: LLMRequest) => messageTexts(request, "system")
 
+const recordedEventTypes = (id: SessionV2.ID) =>
+  Effect.gen(function* () {
+    const { db } = yield* Database.Service
+    return yield* db
+      .select({ type: EventTable.type })
+      .from(EventTable)
+      .where(eq(EventTable.aggregate_id, id))
+      .orderBy(asc(EventTable.seq))
+      .all()
+      .pipe(Effect.orDie, Effect.map((rows) => rows.map((row) => row.type)))
+  })
+
 const replaySessionProjection = (id: SessionV2.ID) =>
   Effect.gen(function* () {
     const { db } = yield* Database.Service
@@ -2772,6 +2784,11 @@ describe("SessionRunnerLLM", () => {
 
       expect(Exit.isFailure(exit) && Cause.hasInterruptsOnly(exit.cause)).toBeTrue()
       expect(requests).toHaveLength(1)
+      expect(yield* session.context(sessionID)).toMatchObject([
+        { type: "user", text: "Interrupt provider" },
+        { type: "assistant", finish: "error", error: { type: "unknown", message: "Provider turn interrupted" } },
+      ])
+      expect(yield* recordedEventTypes(sessionID)).toContain("session.next.step.failed.2")
       yield* session.interrupt(sessionID)
     }),
   )
@@ -2801,6 +2818,8 @@ describe("SessionRunnerLLM", () => {
         { type: "user", text: "Interrupt tool settlement" },
         {
           type: "assistant",
+          finish: "error",
+          error: { type: "unknown", message: "Provider turn interrupted" },
           content: [
             {
               type: "tool",
@@ -2810,6 +2829,9 @@ describe("SessionRunnerLLM", () => {
           ],
         },
       ])
+      const eventTypes = yield* recordedEventTypes(sessionID)
+      expect(eventTypes).toContain("session.next.step.failed.2")
+      expect(eventTypes).not.toContain("session.next.step.ended.2")
     }),
   )
 

+ 25 - 0
packages/tui/test/cli/tui/data.test.tsx

@@ -224,6 +224,31 @@ test("tracks session status from active sessions and execution events", async ()
       },
     })
     await wait(() => data.session.status("session-live") === "idle")
+
+    emitEvent(events, {
+      id: "evt_failed_step_started",
+      type: "session.next.step.started",
+      data: {
+        sessionID: "session-failed",
+        assistantMessageID: "message-failed",
+        timestamp: 3,
+        agent: "build",
+        model: { id: "model", providerID: "provider" },
+      },
+    })
+    await wait(() => data.session.status("session-failed") === "running")
+
+    emitEvent(events, {
+      id: "evt_step_failed",
+      type: "session.next.step.failed",
+      data: {
+        sessionID: "session-failed",
+        assistantMessageID: "message-failed",
+        timestamp: 4,
+        error: { type: "unknown", message: "Provider unavailable" },
+      },
+    })
+    await wait(() => data.session.status("session-failed") === "idle")
   } finally {
     app.renderer.destroy()
   }