fix(core): persist assistant text state
This commit is contained in:
parent
fc8bbb153f
commit
ce1203ce83
8 changed files with 144 additions and 36 deletions
|
|
@ -98,8 +98,6 @@ export type SessionMessageShell = {
|
|||
output?: { output: string; cursor: number; size: number; truncated: boolean }
|
||||
}
|
||||
|
||||
export type SessionMessageAssistantText = { type: "text"; text: string }
|
||||
|
||||
export type SessionMessageProviderState = { [x: string]: JsonValue }
|
||||
|
||||
export type SessionMessageToolStateStreaming = { status: "streaming"; input: string }
|
||||
|
|
@ -151,8 +149,6 @@ export type ShellInfo = {
|
|||
time: { started: number; completed?: number }
|
||||
}
|
||||
|
||||
export type SessionMessageProviderState3 = { [x: string]: any }
|
||||
|
||||
export type SessionMessageProviderState4 = { [x: string]: any }
|
||||
|
||||
export type SessionMessageProviderState5 = { [x: string]: any }
|
||||
|
|
@ -161,6 +157,12 @@ export type SessionMessageProviderState6 = { [x: string]: any }
|
|||
|
||||
export type SessionMessageProviderState7 = { [x: string]: any }
|
||||
|
||||
export type SessionMessageProviderState8 = { [x: string]: any }
|
||||
|
||||
export type SessionMessageProviderState9 = { [x: string]: any }
|
||||
|
||||
export type SessionMessageProviderState10 = { [x: string]: any }
|
||||
|
||||
export type EventLogSynced = { type: "log.synced"; aggregateID: string; seq?: number }
|
||||
|
||||
export type ModelReasoningField = "reasoning" | "reasoning_content" | "reasoning_text" | (string & {})
|
||||
|
|
@ -717,26 +719,6 @@ export type SessionStepEnded = {
|
|||
}
|
||||
}
|
||||
|
||||
export type SessionTextStarted = {
|
||||
id: string
|
||||
created: number
|
||||
metadata?: { [x: string]: any }
|
||||
type: "session.text.started"
|
||||
durable: { aggregateID: string; seq: number; version: 1 }
|
||||
location?: LocationRef
|
||||
data: { sessionID: string; assistantMessageID: string; ordinal: number }
|
||||
}
|
||||
|
||||
export type SessionTextEnded = {
|
||||
id: string
|
||||
created: number
|
||||
metadata?: { [x: string]: any }
|
||||
type: "session.text.ended"
|
||||
durable: { aggregateID: string; seq: number; version: 1 }
|
||||
location?: LocationRef
|
||||
data: { sessionID: string; assistantMessageID: string; ordinal: number; text: string }
|
||||
}
|
||||
|
||||
export type SessionToolInputStarted = {
|
||||
id: string
|
||||
created: number
|
||||
|
|
@ -1234,6 +1216,8 @@ export type SessionPendingSynthetic = {
|
|||
delivery: "steer" | "queue"
|
||||
}
|
||||
|
||||
export type SessionMessageAssistantText = { type: "text"; text: string; state?: SessionMessageProviderState }
|
||||
|
||||
export type SessionMessageAssistantReasoning = {
|
||||
type: "reasoning"
|
||||
text: string
|
||||
|
|
@ -1344,6 +1328,32 @@ export type ShellCreated = {
|
|||
data: { info: ShellInfo }
|
||||
}
|
||||
|
||||
export type SessionTextStarted = {
|
||||
id: string
|
||||
created: number
|
||||
metadata?: { [x: string]: any }
|
||||
type: "session.text.started"
|
||||
durable: { aggregateID: string; seq: number; version: 1 }
|
||||
location?: LocationRef
|
||||
data: { sessionID: string; assistantMessageID: string; ordinal: number; state?: SessionMessageProviderState4 }
|
||||
}
|
||||
|
||||
export type SessionTextEnded = {
|
||||
id: string
|
||||
created: number
|
||||
metadata?: { [x: string]: any }
|
||||
type: "session.text.ended"
|
||||
durable: { aggregateID: string; seq: number; version: 1 }
|
||||
location?: LocationRef
|
||||
data: {
|
||||
sessionID: string
|
||||
assistantMessageID: string
|
||||
ordinal: number
|
||||
text: string
|
||||
state?: SessionMessageProviderState5
|
||||
}
|
||||
}
|
||||
|
||||
export type SessionReasoningStarted = {
|
||||
id: string
|
||||
created: number
|
||||
|
|
@ -1351,7 +1361,7 @@ export type SessionReasoningStarted = {
|
|||
type: "session.reasoning.started"
|
||||
durable: { aggregateID: string; seq: number; version: 1 }
|
||||
location?: LocationRef
|
||||
data: { sessionID: string; assistantMessageID: string; ordinal: number; state?: SessionMessageProviderState3 }
|
||||
data: { sessionID: string; assistantMessageID: string; ordinal: number; state?: SessionMessageProviderState6 }
|
||||
}
|
||||
|
||||
export type SessionReasoningEnded = {
|
||||
|
|
@ -1366,7 +1376,7 @@ export type SessionReasoningEnded = {
|
|||
assistantMessageID: string
|
||||
ordinal: number
|
||||
text: string
|
||||
state?: SessionMessageProviderState4
|
||||
state?: SessionMessageProviderState7
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -1383,7 +1393,7 @@ export type SessionToolCalled = {
|
|||
callID: string
|
||||
input: { [x: string]: any }
|
||||
executed: boolean
|
||||
state?: SessionMessageProviderState5
|
||||
state?: SessionMessageProviderState8
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -1848,7 +1858,7 @@ export type SessionToolSuccess = {
|
|||
content: Array<LLMToolContent>
|
||||
result?: any
|
||||
executed: boolean
|
||||
resultState?: SessionMessageProviderState6
|
||||
resultState?: SessionMessageProviderState9
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -1868,7 +1878,7 @@ export type SessionToolFailed = {
|
|||
metadata?: { [x: string]: any }
|
||||
result?: any
|
||||
executed: boolean
|
||||
resultState?: SessionMessageProviderState7
|
||||
resultState?: SessionMessageProviderState10
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -307,7 +307,9 @@ export function update(adapter: Adapter, event: SessionEvent.Event) {
|
|||
},
|
||||
"session.text.started": (event) => {
|
||||
return updateOwnedAssistant(event.data.assistantMessageID, (draft) => {
|
||||
draft.content.push(castDraft(SessionMessage.AssistantText.make({ type: "text", text: "" })))
|
||||
draft.content.push(
|
||||
castDraft(SessionMessage.AssistantText.make({ type: "text", text: "", state: event.data.state })),
|
||||
)
|
||||
})
|
||||
},
|
||||
"session.text.delta": (event) => {
|
||||
|
|
@ -319,7 +321,10 @@ export function update(adapter: Adapter, event: SessionEvent.Event) {
|
|||
"session.text.ended": (event) => {
|
||||
return updateOwnedAssistant(event.data.assistantMessageID, (draft) => {
|
||||
const match = latestText(draft)
|
||||
if (match) match.text = event.data.text
|
||||
if (match) {
|
||||
match.text = event.data.text
|
||||
if (event.data.state !== undefined) match.state = event.data.state
|
||||
}
|
||||
})
|
||||
},
|
||||
"session.tool.input.started": (event) => {
|
||||
|
|
|
|||
|
|
@ -142,13 +142,14 @@ export const createLLMEventPublisher = (events: Pick<EventV2.Interface, "publish
|
|||
|
||||
const text = fragments(
|
||||
"text",
|
||||
(_textID, value, ordinal) =>
|
||||
(_textID, value, ordinal, state) =>
|
||||
Effect.gen(function* () {
|
||||
yield* events.publish(SessionEvent.Text.Ended, {
|
||||
sessionID: input.sessionID,
|
||||
assistantMessageID: yield* currentAssistantMessageID(),
|
||||
ordinal,
|
||||
text: value,
|
||||
state,
|
||||
})
|
||||
}),
|
||||
true,
|
||||
|
|
@ -317,15 +318,16 @@ export const createLLMEventPublisher = (events: Pick<EventV2.Interface, "publish
|
|||
return
|
||||
case "text-start":
|
||||
retryEvidence = true
|
||||
const startedTextOrdinal = yield* text.start(event.id)
|
||||
const startedTextOrdinal = yield* text.start(event.id, providerState(event.providerMetadata))
|
||||
yield* events.publish(SessionEvent.Text.Started, {
|
||||
sessionID: input.sessionID,
|
||||
assistantMessageID: yield* startAssistant(),
|
||||
ordinal: startedTextOrdinal,
|
||||
state: providerState(event.providerMetadata),
|
||||
})
|
||||
return
|
||||
case "text-delta":
|
||||
const deltaTextOrdinal = yield* text.append(event.id, event.text)
|
||||
const deltaTextOrdinal = yield* text.append(event.id, event.text, providerState(event.providerMetadata))
|
||||
yield* events.publish(SessionEvent.Text.Delta, {
|
||||
sessionID: input.sessionID,
|
||||
assistantMessageID: yield* currentAssistantMessageID(),
|
||||
|
|
@ -334,7 +336,7 @@ export const createLLMEventPublisher = (events: Pick<EventV2.Interface, "publish
|
|||
})
|
||||
return
|
||||
case "text-end":
|
||||
yield* text.end(event.id)
|
||||
yield* text.end(event.id, providerState(event.providerMetadata))
|
||||
return
|
||||
case "reasoning-start":
|
||||
retryEvidence = true
|
||||
|
|
|
|||
|
|
@ -123,7 +123,14 @@ const assistant = (message: SessionMessage.Assistant, model: ModelV2.Ref, provid
|
|||
String(message.model.providerID) === String(model.providerID) && String(message.model.id) === String(model.id)
|
||||
const reuseProviderMetadata = sameModel && message.error === undefined
|
||||
const content = message.content.flatMap((item): ContentPart[] => {
|
||||
if (item.type === "text") return [{ type: "text", text: item.text }]
|
||||
if (item.type === "text")
|
||||
return [
|
||||
{
|
||||
type: "text",
|
||||
text: item.text,
|
||||
providerMetadata: reuseProviderMetadata ? providerMetadata(providerMetadataKey, item.state) : undefined,
|
||||
},
|
||||
]
|
||||
if (item.type === "reasoning")
|
||||
return reuseProviderMetadata
|
||||
? [
|
||||
|
|
|
|||
|
|
@ -15,6 +15,36 @@ const model = ModelV2.Ref.make({ id: ModelV2.ID.make("model"), providerID: Provi
|
|||
const build = AgentV2.defaultID
|
||||
|
||||
describe("toLLMMessages", () => {
|
||||
test("restores same-model provider state on assistant text", () => {
|
||||
const messages = toLLMMessages(
|
||||
[
|
||||
SessionMessage.Assistant.make({
|
||||
id: id("assistant-text-state"),
|
||||
type: "assistant",
|
||||
agent: build,
|
||||
model,
|
||||
content: [
|
||||
SessionMessage.AssistantText.make({
|
||||
type: "text",
|
||||
text: "Checking.",
|
||||
state: { itemId: "msg_commentary", phase: "commentary" },
|
||||
}),
|
||||
],
|
||||
time: { created, completed: created },
|
||||
}),
|
||||
],
|
||||
model,
|
||||
)
|
||||
|
||||
expect(messages[0]?.content).toEqual([
|
||||
{
|
||||
type: "text",
|
||||
text: "Checking.",
|
||||
providerMetadata: { provider: { itemId: "msg_commentary", phase: "commentary" } },
|
||||
},
|
||||
])
|
||||
})
|
||||
|
||||
test("omits empty assistant turns", () => {
|
||||
const assistant = (value: string, content: SessionMessage.Assistant["content"]) =>
|
||||
SessionMessage.Assistant.make({
|
||||
|
|
@ -661,6 +691,11 @@ Recent work
|
|||
agent: build,
|
||||
model: { id: ModelV2.ID.make("old-model"), providerID: ProviderV2.ID.make("provider") },
|
||||
content: [
|
||||
SessionMessage.AssistantText.make({
|
||||
type: "text",
|
||||
text: "Checking",
|
||||
state: { itemId: "msg_old", phase: "commentary" },
|
||||
}),
|
||||
SessionMessage.AssistantReasoning.make({
|
||||
type: "reasoning",
|
||||
text: "Visible thought",
|
||||
|
|
@ -705,6 +740,7 @@ Recent work
|
|||
)
|
||||
|
||||
expect(messages[0]?.content).toEqual([
|
||||
{ type: "text", text: "Checking", providerMetadata: undefined },
|
||||
{ type: "text", text: "Visible thought" },
|
||||
{
|
||||
type: "tool-call",
|
||||
|
|
|
|||
|
|
@ -2434,6 +2434,51 @@ describe("SessionRunnerLLM", () => {
|
|||
}),
|
||||
)
|
||||
|
||||
it.effect("preserves assistant text provider state across durable continuation", () =>
|
||||
Effect.gen(function* () {
|
||||
const session = yield* setup
|
||||
yield* admit(session, "Check this")
|
||||
|
||||
const providerMetadata = { openai: { itemId: "msg_commentary", phase: "commentary" } }
|
||||
responses = [
|
||||
[
|
||||
LLMEvent.stepStart({ index: 0 }),
|
||||
LLMEvent.textStart({ id: "msg_commentary", providerMetadata }),
|
||||
LLMEvent.textDelta({ id: "msg_commentary", text: "Checking.", providerMetadata }),
|
||||
LLMEvent.textEnd({ id: "msg_commentary", providerMetadata }),
|
||||
LLMEvent.toolCall({ id: "call-echo", name: "echo", input: { text: "hello" } }),
|
||||
LLMEvent.stepFinish({ index: 0, reason: "tool-calls" }),
|
||||
LLMEvent.finish({ reason: "tool-calls" }),
|
||||
],
|
||||
reply.text("Done", "text-final"),
|
||||
]
|
||||
|
||||
yield* session.resume(sessionID)
|
||||
|
||||
expect(requests).toHaveLength(2)
|
||||
expect(requests[1]?.messages[1]?.content[0]).toEqual({
|
||||
type: "text",
|
||||
text: "Checking.",
|
||||
providerMetadata,
|
||||
})
|
||||
expect(yield* session.context(sessionID)).toMatchObject([
|
||||
{ type: "user", text: "Check this" },
|
||||
{
|
||||
type: "assistant",
|
||||
content: [
|
||||
{
|
||||
type: "text",
|
||||
text: "Checking.",
|
||||
state: { itemId: "msg_commentary", phase: "commentary" },
|
||||
},
|
||||
{ type: "tool", id: "call-echo" },
|
||||
],
|
||||
},
|
||||
{ type: "assistant", content: [{ type: "text", text: "Done" }] },
|
||||
])
|
||||
}),
|
||||
)
|
||||
|
||||
it.effect("reloads a model switch before a tool-driven continuation step", () =>
|
||||
Effect.gen(function* () {
|
||||
const session = yield* setup
|
||||
|
|
|
|||
|
|
@ -290,6 +290,7 @@ export namespace Text {
|
|||
...Base,
|
||||
assistantMessageID: SessionMessage.ID,
|
||||
ordinal: NonNegativeInt,
|
||||
state: SessionMessage.ProviderState.pipe(optional),
|
||||
},
|
||||
})
|
||||
export type Started = typeof Started.Type
|
||||
|
|
@ -314,6 +315,7 @@ export namespace Text {
|
|||
assistantMessageID: SessionMessage.ID,
|
||||
ordinal: NonNegativeInt,
|
||||
text: Schema.String,
|
||||
state: SessionMessage.ProviderState.pipe(optional),
|
||||
},
|
||||
})
|
||||
export type Ended = typeof Ended.Type
|
||||
|
|
|
|||
|
|
@ -158,6 +158,7 @@ export interface AssistantText extends Schema.Schema.Type<typeof AssistantText>
|
|||
export const AssistantText = Schema.Struct({
|
||||
type: Schema.tag("text"),
|
||||
text: Schema.String,
|
||||
state: ProviderState.pipe(optional),
|
||||
}).annotate({ identifier: "Session.Message.Assistant.Text" })
|
||||
|
||||
export interface AssistantReasoning extends Schema.Schema.Type<typeof AssistantReasoning> {}
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue