Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
27 changes: 21 additions & 6 deletions providers/openai/responses_language_model.go
Original file line number Diff line number Diff line change
Expand Up @@ -581,7 +581,13 @@ func toResponsesPromptWithValidation(prompt fantasy.Prompt, systemMessageMode st
})
continue
}
input = append(input, responses.ResponseInputItemParamOfMessage(textPart.Text, responses.EasyInputMessageRoleAssistant))
message := responses.ResponseInputItemParamOfMessage(textPart.Text, responses.EasyInputMessageRoleAssistant)
// Resend the phase whether or not items are stored: models
// that label messages degrade when follow-ups drop it.
if metadata, ok := textPart.ProviderOptions[Name].(*ResponsesTextMetadata); ok && metadata != nil {
message.OfMessage.Phase = responses.EasyInputMessagePhase(metadata.Phase)
}
input = append(input, message)
lastEmittedReasoningReference = false

case fantasy.ContentTypeToolCall:
Expand Down Expand Up @@ -1055,7 +1061,8 @@ func (o responsesLanguageModel) Generate(ctx context.Context, call fantasy.Call)
for _, contentPart := range outputItem.Content {
if contentPart.Type == "output_text" {
content = append(content, fantasy.TextContent{
Text: contentPart.Text,
Text: contentPart.Text,
ProviderMetadata: responsesTextMetadata(outputItem.ID, outputItem.Phase),
})

for _, annotation := range contentPart.Annotations {
Expand Down Expand Up @@ -1281,8 +1288,9 @@ func (o responsesLanguageModel) Stream(ctx context.Context, call fantasy.Call) (

case "message":
if !yield(fantasy.StreamPart{
Type: fantasy.StreamPartTypeTextStart,
ID: added.Item.ID,
Type: fantasy.StreamPartTypeTextStart,
ID: added.Item.ID,
ProviderMetadata: responsesTextMetadata(added.Item.ID, added.Item.Phase),
}) {
return
}
Expand Down Expand Up @@ -1368,8 +1376,9 @@ func (o responsesLanguageModel) Stream(ctx context.Context, call fantasy.Call) (
}
case "message":
if !yield(fantasy.StreamPart{
Type: fantasy.StreamPartTypeTextEnd,
ID: done.Item.ID,
Type: fantasy.StreamPartTypeTextEnd,
ID: done.Item.ID,
ProviderMetadata: responsesTextMetadata(done.Item.ID, done.Item.Phase),
}) {
return
}
Expand Down Expand Up @@ -1659,6 +1668,12 @@ func webSearchCallToMetadata(itemID string, action responses.ResponseOutputItemU
return meta
}

func responsesTextMetadata(itemID string, phase responses.ResponseOutputMessagePhase) fantasy.ProviderMetadata {
return fantasy.ProviderMetadata{
Name: &ResponsesTextMetadata{ItemID: itemID, Phase: string(phase)},
}
}

// GetReasoningMetadata extracts reasoning metadata from provider options for responses models.
func GetReasoningMetadata(providerOptions fantasy.ProviderOptions) *ResponsesReasoningMetadata {
if openaiResponsesOptions, ok := providerOptions[Name]; ok {
Expand Down
36 changes: 36 additions & 0 deletions providers/openai/responses_options.go
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@ import (
const (
TypeResponsesProviderMetadata = Name + ".responses.metadata"
TypeResponsesProviderOptions = Name + ".responses.options"
TypeResponsesTextMetadata = Name + ".responses.text_metadata"
TypeResponsesReasoningMetadata = Name + ".responses.reasoning_metadata"
TypeWebSearchCallMetadata = Name + ".responses.web_search_call_metadata"
)
Expand All @@ -34,6 +35,13 @@ func init() {
}
return &v, nil
})
fantasy.RegisterProviderType(TypeResponsesTextMetadata, func(data []byte) (fantasy.ProviderOptionsData, error) {
var v ResponsesTextMetadata
if err := json.Unmarshal(data, &v); err != nil {
return nil, err
}
return &v, nil
})
fantasy.RegisterProviderType(TypeResponsesReasoningMetadata, func(data []byte) (fantasy.ProviderOptionsData, error) {
var v ResponsesReasoningMetadata
if err := json.Unmarshal(data, &v); err != nil {
Expand Down Expand Up @@ -92,6 +100,34 @@ func (m *ResponsesProviderMetadata) UnmarshalJSON(data []byte) error {
return nil
}

// ResponsesTextMetadata identifies the output message a text part came from.
// Phase labels the message as intermediate commentary or the final answer
// on models that send it, and is replayed on follow-up requests.
type ResponsesTextMetadata struct {
ItemID string `json:"item_id"`
Phase string `json:"phase,omitempty"`
}

// Options implements the ProviderOptions interface.
func (*ResponsesTextMetadata) Options() {}

// MarshalJSON implements custom JSON marshaling with type info for ResponsesTextMetadata.
func (m ResponsesTextMetadata) MarshalJSON() ([]byte, error) {
type plain ResponsesTextMetadata
return fantasy.MarshalProviderType(TypeResponsesTextMetadata, plain(m))
}

// UnmarshalJSON implements custom JSON unmarshaling with type info for ResponsesTextMetadata.
func (m *ResponsesTextMetadata) UnmarshalJSON(data []byte) error {
type plain ResponsesTextMetadata
var p plain
if err := fantasy.UnmarshalProviderType(data, &p); err != nil {
return err
}
*m = ResponsesTextMetadata(p)
return nil
}

// ResponsesReasoningMetadata represents reasoning metadata for OpenAI Responses API.
type ResponsesReasoningMetadata struct {
ItemID string `json:"item_id"`
Expand Down
146 changes: 146 additions & 0 deletions providers/openai/responses_text_phase_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,146 @@
package openai

import (
"context"
"encoding/json"
"fmt"
"testing"

"charm.land/fantasy"
"github.com/stretchr/testify/require"
)

func TestResponsesStream_TextPhaseMetadata(t *testing.T) {
t.Parallel()

for _, phase := range []string{"commentary", "final_answer", ""} {
t.Run(fmt.Sprintf("phase %q", phase), func(t *testing.T) {
t.Parallel()

phaseField := ""
if phase != "" {
phaseField = fmt.Sprintf(`,"phase":%q`, phase)
}
sms := newStreamingMockServer()
defer sms.close()
sms.chunks = []string{
responsesSSEEvent("response.created", `{"type":"response.created","response":{"id":"resp_01","status":"in_progress","output":[]}}`),
responsesSSEEvent("response.output_item.added", `{"type":"response.output_item.added","output_index":0,"item":{"id":"msg_01","type":"message","role":"assistant","status":"in_progress"`+phaseField+`,"content":[]}}`),
responsesSSEEvent("response.output_text.delta", `{"type":"response.output_text.delta","output_index":0,"content_index":0,"item_id":"msg_01","delta":"Reading the file."}`),
responsesSSEEvent("response.output_item.done", `{"type":"response.output_item.done","output_index":0,"item":{"id":"msg_01","type":"message","role":"assistant","status":"completed"`+phaseField+`,"content":[{"type":"output_text","text":"Reading the file.","annotations":[]}]}}`),
responsesSSEEvent("response.completed", `{"type":"response.completed","response":{"id":"resp_01","status":"completed","output":[],"usage":{"input_tokens":1,"output_tokens":1,"total_tokens":2}}}`),
}

model := newResponsesProvider(t, sms.server.URL)
stream, err := model.Stream(context.Background(), fantasy.Call{Prompt: testPrompt})
require.NoError(t, err)
parts, err := collectStreamParts(stream)
require.NoError(t, err)

want := &ResponsesTextMetadata{ItemID: "msg_01", Phase: phase}
var bounds []fantasy.StreamPartType
for _, part := range parts {
if part.Type != fantasy.StreamPartTypeTextStart && part.Type != fantasy.StreamPartTypeTextEnd {
continue
}
bounds = append(bounds, part.Type)
require.Equal(t, want, part.ProviderMetadata[Name], part.Type)
}
require.Equal(t, []fantasy.StreamPartType{fantasy.StreamPartTypeTextStart, fantasy.StreamPartTypeTextEnd}, bounds)
})
}
}

func TestResponsesGenerate_TextPhaseMetadata(t *testing.T) {
t.Parallel()

message := func(id, phase, text string) map[string]any {
return map[string]any{
"type": "message",
"id": id,
"role": "assistant",
"status": "completed",
"phase": phase,
"content": []any{map[string]any{"type": "output_text", "text": text, "annotations": []any{}}},
}
}
server := newMockServer()
defer server.close()
server.response = map[string]any{
"id": "resp_01",
"object": "response",
"model": "gpt-4.1",
"status": "completed",
"output": []any{
message("msg_commentary", "commentary", "Reading the file."),
message("msg_answer", "final_answer", "It is flaky."),
},
"usage": map[string]any{"input_tokens": 1, "output_tokens": 1, "total_tokens": 2},
}

model := newResponsesProvider(t, server.server.URL)
resp, err := model.Generate(context.Background(), fantasy.Call{Prompt: testPrompt})
require.NoError(t, err)

var metadata []fantasy.ProviderOptionsData
for _, content := range resp.Content {
if text, ok := content.(fantasy.TextContent); ok {
metadata = append(metadata, text.ProviderMetadata[Name])
}
}
require.Equal(t, []fantasy.ProviderOptionsData{
&ResponsesTextMetadata{ItemID: "msg_commentary", Phase: "commentary"},
&ResponsesTextMetadata{ItemID: "msg_answer", Phase: "final_answer"},
}, metadata)
}

func TestResponsesToPrompt_ReplaysTextPhase(t *testing.T) {
t.Parallel()

// Callers store the metadata as JSON and restore it through the registry.
stored, err := json.Marshal(&ResponsesTextMetadata{ItemID: "msg_commentary", Phase: "commentary"})
require.NoError(t, err)
restored, err := fantasy.UnmarshalProviderOptions(map[string]json.RawMessage{Name: stored})
require.NoError(t, err)

prompt := fantasy.Prompt{
{
Role: fantasy.MessageRoleUser,
Content: []fantasy.MessagePart{fantasy.TextPart{Text: "Why is TestFoo flaky?"}},
},
{
Role: fantasy.MessageRoleAssistant,
Content: []fantasy.MessagePart{
fantasy.TextPart{Text: "Reading the file.", ProviderOptions: restored},
fantasy.TextPart{Text: "It is flaky.", ProviderOptions: fantasy.ProviderOptions{
Name: &ResponsesTextMetadata{ItemID: "msg_answer", Phase: "final_answer"},
}},
fantasy.TextPart{Text: "Unlabeled."},
fantasy.TextPart{Text: "No phase.", ProviderOptions: fantasy.ProviderOptions{
Name: &ResponsesTextMetadata{ItemID: "msg_no_phase"},
}},
},
},
}

for _, store := range []bool{true, false} {
t.Run(fmt.Sprintf("store %t", store), func(t *testing.T) {
t.Parallel()

input, warnings, err := toResponsesPrompt(prompt, "system", store)
require.NoError(t, err)
require.Empty(t, warnings)
require.Len(t, input, 5)

var phases []any
for _, item := range input[1:] {
raw, err := json.Marshal(item)
require.NoError(t, err)
var fields map[string]any
require.NoError(t, json.Unmarshal(raw, &fields))
phases = append(phases, fields["phase"])
}
require.Equal(t, []any{"commentary", "final_answer", nil, nil}, phases)
})
}
}
Loading