1
0

session-tool-progress.test.ts 6.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161
  1. import { describe, expect } from "bun:test"
  2. import { asc, eq } from "drizzle-orm"
  3. import { DateTime, Effect, Layer, Schema } from "effect"
  4. import { Database } from "@opencode-ai/core/database/database"
  5. import { EventV2 } from "@opencode-ai/core/event"
  6. import { EventTable } from "@opencode-ai/core/event/sql"
  7. import { ModelV2 } from "@opencode-ai/core/model"
  8. import { Project } from "@opencode-ai/core/project"
  9. import { ProjectTable } from "@opencode-ai/core/project/sql"
  10. import { ProviderV2 } from "@opencode-ai/core/provider"
  11. import { AbsolutePath } from "@opencode-ai/core/schema"
  12. import { SessionV2 } from "@opencode-ai/core/session"
  13. import { SessionEvent } from "@opencode-ai/core/session/event"
  14. import { SessionMessage } from "@opencode-ai/core/session/message"
  15. import { SessionProjector } from "@opencode-ai/core/session/projector"
  16. import { SessionTable, SessionMessageTable } from "@opencode-ai/core/session/sql"
  17. import { ToolOutput } from "@opencode-ai/core/tool-output"
  18. import { testEffect } from "./lib/effect"
  19. const database = Database.layerFromPath(":memory:")
  20. const events = EventV2.layer.pipe(Layer.provide(database))
  21. const projector = SessionProjector.layer.pipe(Layer.provide(events), Layer.provide(database))
  22. const it = testEffect(Layer.mergeAll(database, events, projector))
  23. const timestamp = DateTime.makeUnsafe(1)
  24. const model = { id: ModelV2.ID.make("model"), providerID: ProviderV2.ID.make("provider") }
  25. const content = (text: string) => [ToolOutput.text({ type: "text", text })]
  26. describe("Tool.Progress", () => {
  27. it.effect("projects durable progress and keeps final settlements durable", () =>
  28. Effect.gen(function* () {
  29. const { db } = yield* Database.Service
  30. const service = yield* EventV2.Service
  31. const sessionID = SessionV2.ID.make("ses_tool_progress_projector")
  32. yield* db
  33. .insert(ProjectTable)
  34. .values({ id: Project.ID.global, worktree: AbsolutePath.make("/project"), sandboxes: [] })
  35. .onConflictDoNothing()
  36. .run()
  37. .pipe(Effect.orDie)
  38. yield* db
  39. .insert(SessionTable)
  40. .values({
  41. id: sessionID,
  42. project_id: Project.ID.global,
  43. slug: "progress",
  44. directory: "/project",
  45. title: "progress",
  46. version: "test",
  47. })
  48. .run()
  49. .pipe(Effect.orDie)
  50. const assistantMessageID = SessionMessage.ID.create()
  51. yield* service.publish(SessionEvent.Step.Started, {
  52. sessionID,
  53. assistantMessageID,
  54. timestamp,
  55. agent: "build",
  56. model,
  57. })
  58. const readAssistant = Effect.gen(function* () {
  59. const row = yield* db
  60. .select()
  61. .from(SessionMessageTable)
  62. .where(eq(SessionMessageTable.id, assistantMessageID))
  63. .get()
  64. .pipe(Effect.orDie)
  65. if (!row) return yield* Effect.die("Missing projected assistant")
  66. return Schema.decodeUnknownSync(SessionMessage.Assistant)({ ...row.data, id: row.id, type: row.type })
  67. })
  68. const start = (callID: string) =>
  69. Effect.gen(function* () {
  70. yield* service.publish(SessionEvent.Tool.Input.Started, {
  71. sessionID,
  72. timestamp,
  73. assistantMessageID,
  74. callID,
  75. name: "bash",
  76. })
  77. yield* service.publish(SessionEvent.Tool.Called, {
  78. sessionID,
  79. timestamp,
  80. assistantMessageID,
  81. callID,
  82. tool: "bash",
  83. input: { command: "pwd" },
  84. provider: { executed: false },
  85. })
  86. })
  87. yield* start("call-success")
  88. expect((yield* readAssistant).content[0]).toMatchObject({
  89. state: { status: "running", structured: {}, content: [] },
  90. })
  91. yield* service.publish(SessionEvent.Tool.Progress, {
  92. sessionID,
  93. timestamp,
  94. assistantMessageID,
  95. callID: "call-success",
  96. structured: { phase: "checkpoint" },
  97. content: content("saved"),
  98. })
  99. expect((yield* readAssistant).content[0]).toMatchObject({
  100. state: { status: "running", structured: { phase: "checkpoint" }, content: content("saved") },
  101. })
  102. const success = yield* service.publish(SessionEvent.Tool.Success, {
  103. sessionID,
  104. timestamp,
  105. assistantMessageID,
  106. callID: "call-success",
  107. structured: { phase: "done" },
  108. content: content("complete"),
  109. provider: { executed: false },
  110. })
  111. expect((yield* readAssistant).content[0]).toMatchObject({
  112. state: { status: "completed", structured: { phase: "done" }, content: content("complete") },
  113. })
  114. yield* start("call-failed")
  115. yield* service.publish(SessionEvent.Tool.Progress, {
  116. sessionID,
  117. timestamp,
  118. assistantMessageID,
  119. callID: "call-failed",
  120. structured: { phase: "checkpoint" },
  121. content: content("before failure"),
  122. })
  123. const failed = yield* service.publish(SessionEvent.Tool.Failed, {
  124. sessionID,
  125. timestamp,
  126. assistantMessageID,
  127. callID: "call-failed",
  128. error: { type: "unknown", message: "boom" },
  129. provider: { executed: false },
  130. })
  131. expect((yield* readAssistant).content[1]).toMatchObject({
  132. state: {
  133. status: "error",
  134. structured: { phase: "checkpoint" },
  135. content: content("before failure"),
  136. error: { type: "unknown", message: "boom" },
  137. },
  138. })
  139. expect(Schema.is(SessionEvent.Durable)(success)).toBe(true)
  140. expect(Schema.is(SessionEvent.Durable)(failed)).toBe(true)
  141. const rows = yield* db
  142. .select({ type: EventTable.type })
  143. .from(EventTable)
  144. .where(eq(EventTable.aggregate_id, sessionID))
  145. .orderBy(asc(EventTable.seq))
  146. .all()
  147. .pipe(Effect.orDie)
  148. expect(rows.map((row) => row.type)).toContain(EventV2.versionedType(SessionEvent.Tool.Progress.type, 1))
  149. expect(rows.map((row) => row.type)).toContain(EventV2.versionedType(SessionEvent.Tool.Success.type, 1))
  150. expect(rows.map((row) => row.type)).toContain(EventV2.versionedType(SessionEvent.Tool.Failed.type, 1))
  151. }),
  152. )
  153. })