fix(ai): preserve OpenAI message phases
This commit is contained in:
parent
a817fe5e6c
commit
fc8bbb153f
3 changed files with 131 additions and 14 deletions
|
|
@ -62,6 +62,9 @@ const OpenAIResponsesOutputText = Schema.Struct({
|
||||||
text: Schema.String,
|
text: Schema.String,
|
||||||
})
|
})
|
||||||
|
|
||||||
|
const OpenAIResponsesMessagePhase = Schema.Literals(["commentary", "final_answer"])
|
||||||
|
type OpenAIResponsesMessagePhase = Schema.Schema.Type<typeof OpenAIResponsesMessagePhase>
|
||||||
|
|
||||||
const OpenAIResponsesReasoningSummaryText = Schema.Struct({
|
const OpenAIResponsesReasoningSummaryText = Schema.Struct({
|
||||||
type: Schema.tag("summary_text"),
|
type: Schema.tag("summary_text"),
|
||||||
text: Schema.String,
|
text: Schema.String,
|
||||||
|
|
@ -96,7 +99,11 @@ const OpenAIResponsesFunctionCallOutput = Schema.Union([
|
||||||
const OpenAIResponsesInputItem = Schema.Union([
|
const OpenAIResponsesInputItem = Schema.Union([
|
||||||
Schema.Struct({ role: Schema.tag("system"), content: Schema.String }),
|
Schema.Struct({ role: Schema.tag("system"), content: Schema.String }),
|
||||||
Schema.Struct({ role: Schema.tag("user"), content: Schema.Array(OpenAIResponsesInputContent) }),
|
Schema.Struct({ role: Schema.tag("user"), content: Schema.Array(OpenAIResponsesInputContent) }),
|
||||||
Schema.Struct({ role: Schema.tag("assistant"), content: Schema.Array(OpenAIResponsesOutputText) }),
|
Schema.Struct({
|
||||||
|
role: Schema.tag("assistant"),
|
||||||
|
content: Schema.Array(OpenAIResponsesOutputText),
|
||||||
|
phase: optionalNull(OpenAIResponsesMessagePhase),
|
||||||
|
}),
|
||||||
OpenAIResponsesReasoningItem,
|
OpenAIResponsesReasoningItem,
|
||||||
OpenAIResponsesItemReference,
|
OpenAIResponsesItemReference,
|
||||||
Schema.Struct({
|
Schema.Struct({
|
||||||
|
|
@ -216,6 +223,7 @@ const OpenAIResponsesStreamItem = Schema.Struct({
|
||||||
// call's typed input portion and round-trip the full result payload without
|
// call's typed input portion and round-trip the full result payload without
|
||||||
// hand-rolling a per-tool schema.
|
// hand-rolling a per-tool schema.
|
||||||
status: Schema.optional(Schema.String),
|
status: Schema.optional(Schema.String),
|
||||||
|
phase: optionalNull(OpenAIResponsesMessagePhase),
|
||||||
action: Schema.optional(Schema.Unknown),
|
action: Schema.optional(Schema.Unknown),
|
||||||
queries: Schema.optional(Schema.Unknown),
|
queries: Schema.optional(Schema.Unknown),
|
||||||
results: Schema.optional(Schema.Unknown),
|
results: Schema.optional(Schema.Unknown),
|
||||||
|
|
@ -271,6 +279,7 @@ interface ParserState {
|
||||||
readonly tools: ToolStream.State<string>
|
readonly tools: ToolStream.State<string>
|
||||||
readonly hasFunctionCall: boolean
|
readonly hasFunctionCall: boolean
|
||||||
readonly lifecycle: Lifecycle.State
|
readonly lifecycle: Lifecycle.State
|
||||||
|
readonly textItems: Readonly<Record<string, ProviderMetadata>>
|
||||||
readonly reasoningItems: Readonly<Record<string, ReasoningStreamItem>>
|
readonly reasoningItems: Readonly<Record<string, ReasoningStreamItem>>
|
||||||
readonly store: boolean | undefined
|
readonly store: boolean | undefined
|
||||||
}
|
}
|
||||||
|
|
@ -395,10 +404,7 @@ const lowerToolResultContentItem = Effect.fn("OpenAIResponses.lowerToolResultCon
|
||||||
provider: string,
|
provider: string,
|
||||||
) {
|
) {
|
||||||
if (item.type === "text") return { type: "input_text" as const, text: item.text }
|
if (item.type === "text") return { type: "input_text" as const, text: item.text }
|
||||||
return yield* lowerMedia(
|
return yield* lowerMedia({ type: "media", mediaType: item.mime, data: item.uri, filename: item.name }, provider)
|
||||||
{ type: "media", mediaType: item.mime, data: item.uri, filename: item.name },
|
|
||||||
provider,
|
|
||||||
)
|
|
||||||
})
|
})
|
||||||
|
|
||||||
const lowerToolResultOutput = Effect.fn("OpenAIResponses.lowerToolResultOutput")(function* (
|
const lowerToolResultOutput = Effect.fn("OpenAIResponses.lowerToolResultOutput")(function* (
|
||||||
|
|
@ -442,16 +448,29 @@ const lowerMessages = Effect.fn("OpenAIResponses.lowerMessages")(function* (requ
|
||||||
|
|
||||||
if (message.role === "assistant") {
|
if (message.role === "assistant") {
|
||||||
const content: TextPart[] = []
|
const content: TextPart[] = []
|
||||||
|
let phase: OpenAIResponsesMessagePhase | undefined
|
||||||
const reasoningItems: Record<string, OpenAIResponsesReasoningReplay> = {}
|
const reasoningItems: Record<string, OpenAIResponsesReasoningReplay> = {}
|
||||||
const reasoningReferences = new Set<string>()
|
const reasoningReferences = new Set<string>()
|
||||||
const hostedToolReferences = new Set<string>()
|
const hostedToolReferences = new Set<string>()
|
||||||
const flushText = () => {
|
const flushText = () => {
|
||||||
if (content.length === 0) return
|
if (content.length === 0) return
|
||||||
input.push({ role: "assistant", content: content.map((part) => ({ type: "output_text", text: part.text })) })
|
input.push({
|
||||||
|
role: "assistant",
|
||||||
|
content: content.map((part) => ({ type: "output_text", text: part.text })),
|
||||||
|
...(phase === undefined ? {} : { phase }),
|
||||||
|
})
|
||||||
content.splice(0, content.length)
|
content.splice(0, content.length)
|
||||||
|
phase = undefined
|
||||||
}
|
}
|
||||||
for (const part of message.content) {
|
for (const part of message.content) {
|
||||||
if (part.type === "text") {
|
if (part.type === "text") {
|
||||||
|
const openai = part.providerMetadata?.openai
|
||||||
|
const nextPhase =
|
||||||
|
ProviderShared.isRecord(openai) && (openai.phase === "commentary" || openai.phase === "final_answer")
|
||||||
|
? openai.phase
|
||||||
|
: undefined
|
||||||
|
if (content.length > 0 && phase !== nextPhase) flushText()
|
||||||
|
phase = nextPhase
|
||||||
content.push(part)
|
content.push(part)
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
@ -709,15 +728,28 @@ const TERMINAL_TYPES = new Set(["response.completed", "response.incomplete", "re
|
||||||
const onOutputTextDelta = (state: ParserState, event: OpenAIResponsesEvent): StepResult => {
|
const onOutputTextDelta = (state: ParserState, event: OpenAIResponsesEvent): StepResult => {
|
||||||
if (!event.delta) return [state, NO_EVENTS]
|
if (!event.delta) return [state, NO_EVENTS]
|
||||||
const events: LLMEvent[] = []
|
const events: LLMEvent[] = []
|
||||||
|
const itemID = event.item_id ?? "text-0"
|
||||||
return [
|
return [
|
||||||
{ ...state, lifecycle: Lifecycle.textDelta(state.lifecycle, events, event.item_id ?? "text-0", event.delta) },
|
{
|
||||||
|
...state,
|
||||||
|
lifecycle: Lifecycle.textDelta(state.lifecycle, events, itemID, event.delta, state.textItems[itemID]),
|
||||||
|
},
|
||||||
events,
|
events,
|
||||||
]
|
]
|
||||||
}
|
}
|
||||||
|
|
||||||
const onOutputTextDone = (state: ParserState, event: OpenAIResponsesEvent): StepResult => {
|
const onOutputTextDone = (state: ParserState, event: OpenAIResponsesEvent): StepResult => {
|
||||||
const events: LLMEvent[] = []
|
const events: LLMEvent[] = []
|
||||||
return [{ ...state, lifecycle: Lifecycle.textEnd(state.lifecycle, events, event.item_id ?? "text-0") }, events]
|
const itemID = event.item_id ?? "text-0"
|
||||||
|
const { [itemID]: _completed, ...textItems } = state.textItems
|
||||||
|
return [
|
||||||
|
{
|
||||||
|
...state,
|
||||||
|
lifecycle: Lifecycle.textEnd(state.lifecycle, events, itemID, state.textItems[itemID]),
|
||||||
|
textItems,
|
||||||
|
},
|
||||||
|
events,
|
||||||
|
]
|
||||||
}
|
}
|
||||||
|
|
||||||
const onReasoningDelta = (state: ParserState, event: OpenAIResponsesEvent): StepResult => {
|
const onReasoningDelta = (state: ParserState, event: OpenAIResponsesEvent): StepResult => {
|
||||||
|
|
@ -754,6 +786,21 @@ const reasoningMetadata = (item: OpenAIResponsesStreamItem & { id: string }) =>
|
||||||
// best-effort, not guaranteed.
|
// best-effort, not guaranteed.
|
||||||
const onOutputItemAdded = (state: ParserState, event: OpenAIResponsesEvent): StepResult => {
|
const onOutputItemAdded = (state: ParserState, event: OpenAIResponsesEvent): StepResult => {
|
||||||
const item = event.item
|
const item = event.item
|
||||||
|
if (item?.type === "message" && item.id) {
|
||||||
|
return [
|
||||||
|
{
|
||||||
|
...state,
|
||||||
|
textItems: {
|
||||||
|
...state.textItems,
|
||||||
|
[item.id]: openaiMetadata({
|
||||||
|
itemId: item.id,
|
||||||
|
...(item.phase === undefined || item.phase === null ? {} : { phase: item.phase }),
|
||||||
|
}),
|
||||||
|
},
|
||||||
|
},
|
||||||
|
NO_EVENTS,
|
||||||
|
]
|
||||||
|
}
|
||||||
if (item && isReasoningItem(item)) {
|
if (item && isReasoningItem(item)) {
|
||||||
const events: LLMEvent[] = []
|
const events: LLMEvent[] = []
|
||||||
return [
|
return [
|
||||||
|
|
@ -1063,6 +1110,7 @@ export const protocol = Protocol.make({
|
||||||
hasFunctionCall: false,
|
hasFunctionCall: false,
|
||||||
tools: ToolStream.empty<string>(),
|
tools: ToolStream.empty<string>(),
|
||||||
lifecycle: Lifecycle.initial(),
|
lifecycle: Lifecycle.initial(),
|
||||||
|
textItems: {},
|
||||||
reasoningItems: {},
|
reasoningItems: {},
|
||||||
store: OpenAIOptions.store(request),
|
store: OpenAIOptions.store(request),
|
||||||
}),
|
}),
|
||||||
|
|
|
||||||
|
|
@ -14,16 +14,25 @@ export const stepStart = (state: State, events: LLMEvent[]): State => {
|
||||||
return { ...state, stepStarted: true }
|
return { ...state, stepStarted: true }
|
||||||
}
|
}
|
||||||
|
|
||||||
export const textDelta = (state: State, events: LLMEvent[], id: string, text: string): State => {
|
export const textStart = (state: State, events: LLMEvent[], id: string, providerMetadata?: ProviderMetadata): State => {
|
||||||
|
if (state.text.has(id)) return state
|
||||||
const stepped = stepStart(state, events)
|
const stepped = stepStart(state, events)
|
||||||
if (stepped.text.has(id)) {
|
events.push(LLMEvent.textStart({ id, providerMetadata }))
|
||||||
events.push(LLMEvent.textDelta({ id, text }))
|
|
||||||
return stepped
|
|
||||||
}
|
|
||||||
events.push(LLMEvent.textStart({ id }), LLMEvent.textDelta({ id, text }))
|
|
||||||
return { ...stepped, text: new Set([...stepped.text, id]) }
|
return { ...stepped, text: new Set([...stepped.text, id]) }
|
||||||
}
|
}
|
||||||
|
|
||||||
|
export const textDelta = (
|
||||||
|
state: State,
|
||||||
|
events: LLMEvent[],
|
||||||
|
id: string,
|
||||||
|
text: string,
|
||||||
|
providerMetadata?: ProviderMetadata,
|
||||||
|
): State => {
|
||||||
|
const started = textStart(state, events, id, providerMetadata)
|
||||||
|
events.push(LLMEvent.textDelta({ id, text, providerMetadata }))
|
||||||
|
return started
|
||||||
|
}
|
||||||
|
|
||||||
export const reasoningStart = (
|
export const reasoningStart = (
|
||||||
state: State,
|
state: State,
|
||||||
events: LLMEvent[],
|
events: LLMEvent[],
|
||||||
|
|
|
||||||
|
|
@ -899,6 +899,66 @@ describe("OpenAI Responses route", () => {
|
||||||
}),
|
}),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
it.effect("preserves output message phases in follow-up requests", () =>
|
||||||
|
Effect.gen(function* () {
|
||||||
|
const response = yield* LLMClient.generate(request).pipe(
|
||||||
|
Effect.provide(
|
||||||
|
fixedResponse(
|
||||||
|
sseEvents(
|
||||||
|
{
|
||||||
|
type: "response.output_item.added",
|
||||||
|
item: { type: "message", id: "msg_commentary", phase: "commentary" },
|
||||||
|
},
|
||||||
|
{ type: "response.output_text.delta", item_id: "msg_commentary", delta: "Checking." },
|
||||||
|
{ type: "response.output_text.done", item_id: "msg_commentary" },
|
||||||
|
{
|
||||||
|
type: "response.output_item.added",
|
||||||
|
item: { type: "message", id: "msg_final", phase: "final_answer" },
|
||||||
|
},
|
||||||
|
{ type: "response.output_text.delta", item_id: "msg_final", delta: "Done." },
|
||||||
|
{ type: "response.output_text.done", item_id: "msg_final" },
|
||||||
|
{ type: "response.completed", response: { id: "resp_1" } },
|
||||||
|
),
|
||||||
|
),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
expect(response.message.content).toEqual([
|
||||||
|
{
|
||||||
|
type: "text",
|
||||||
|
text: "Checking.",
|
||||||
|
providerMetadata: { openai: { itemId: "msg_commentary", phase: "commentary" } },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
type: "text",
|
||||||
|
text: "Done.",
|
||||||
|
providerMetadata: { openai: { itemId: "msg_final", phase: "final_answer" } },
|
||||||
|
},
|
||||||
|
])
|
||||||
|
|
||||||
|
const followUp = yield* LLMClient.prepare<OpenAIResponses.OpenAIResponsesBody>(
|
||||||
|
LLM.request({
|
||||||
|
model,
|
||||||
|
messages: [Message.user("Start"), response.message, Message.user("Continue")],
|
||||||
|
}),
|
||||||
|
)
|
||||||
|
expect(followUp.body.input).toEqual([
|
||||||
|
{ role: "user", content: [{ type: "input_text", text: "Start" }] },
|
||||||
|
{
|
||||||
|
role: "assistant",
|
||||||
|
phase: "commentary",
|
||||||
|
content: [{ type: "output_text", text: "Checking." }],
|
||||||
|
},
|
||||||
|
{
|
||||||
|
role: "assistant",
|
||||||
|
phase: "final_answer",
|
||||||
|
content: [{ type: "output_text", text: "Done." }],
|
||||||
|
},
|
||||||
|
{ role: "user", content: [{ type: "input_text", text: "Continue" }] },
|
||||||
|
])
|
||||||
|
}),
|
||||||
|
)
|
||||||
|
|
||||||
it.effect("parses reasoning summary stream fixtures", () =>
|
it.effect("parses reasoning summary stream fixtures", () =>
|
||||||
Effect.gen(function* () {
|
Effect.gen(function* () {
|
||||||
const body = sseEvents(
|
const body = sseEvents(
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue