Skip to content
Open
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
119 changes: 119 additions & 0 deletions core/internal/llmtests/responses_image_generation.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,119 @@
package llmtests

import (
"context"
"os"
"strings"
"testing"
"time"

bifrost "github.com/maximhq/bifrost/core"
"github.com/maximhq/bifrost/core/schemas"
)

// RunResponsesImageGenerationToolTest exercises the hosted image_generation
// tool through Responses, rather than the standalone /images endpoint. It is
// opt-in because it consumes a real image-generation request. Set
// BIFROST_RESPONSES_IMAGE_GENERATION_MODEL to override the default Luna model.
func RunResponsesImageGenerationToolTest(t *testing.T, client *bifrost.Bifrost, ctx context.Context, testConfig ComprehensiveTestConfig) {
if os.Getenv("BIFROST_RUN_RESPONSES_IMAGE_GENERATION_TESTS") != "true" {
return
}
model := os.Getenv("BIFROST_RESPONSES_IMAGE_GENERATION_MODEL")
if model == "" {
model = "gpt-5.6-luna"
}

t.Run("ResponsesImageGenerationTool", func(t *testing.T) {
request := func() *schemas.BifrostResponsesRequest {
return &schemas.BifrostResponsesRequest{
Provider: testConfig.Provider,
Model: model,
Input: []schemas.ResponsesMessage{{
Role: schemas.Ptr(schemas.ResponsesInputMessageRoleUser),
Content: &schemas.ResponsesMessageContent{ContentStr: schemas.Ptr("Generate an image of a cute orange cat sitting in a coffee shop.")},
}},
Params: &schemas.ResponsesParameters{
Tools: []schemas.ResponsesTool{{
Type: schemas.ResponsesToolTypeImageGeneration,
ResponsesToolImageGeneration: &schemas.ResponsesToolImageGeneration{
Action: schemas.Ptr(schemas.ResponsesImageGenerationActionGenerate),
},
}},
},
}
}

t.Run("non_streaming", func(t *testing.T) {
response, err := client.ResponsesRequest(schemas.NewBifrostContext(ctx, schemas.NoDeadline), request())
if err != nil {
if strings.Contains(GetErrorMessage(err), "failed to peek at type field") {
t.Fatalf("image_generation response schema regression: %s", GetErrorMessage(err))
}
t.Fatalf("Responses image_generation request failed for %s/%s: %s", testConfig.Provider, model, GetErrorMessage(err))
}
if response == nil {
t.Fatal("Responses image_generation returned a nil response")
}
if !hasImageGenerationCall(response.Output) {
t.Fatalf("Responses image_generation response contained no image_generation_call")
}
})

t.Run("streaming", func(t *testing.T) {
streamCtx, cancel := context.WithTimeout(ctx, 180*time.Second)
defer cancel()
stream, err := client.ResponsesStreamRequest(schemas.NewBifrostContext(streamCtx, schemas.NoDeadline), request())
if err != nil {
t.Fatalf("Responses image_generation stream setup failed for %s/%s: %s", testConfig.Provider, model, GetErrorMessage(err))
}

var sawImageCall, sawCompleted bool
for {
select {
case chunk, ok := <-stream:
if !ok {
if !sawImageCall {
t.Fatal("Responses image_generation stream contained no image_generation_call")
}
if !sawCompleted {
t.Fatal("Responses image_generation stream ended without response.completed")
}
return
}
if chunk == nil {
continue
}
if chunk.BifrostError != nil {
message := GetErrorMessage(chunk.BifrostError)
if strings.Contains(message, "failed to peek at type field") {
t.Fatalf("image_generation stream schema regression: %s", message)
}
t.Fatalf("Responses image_generation stream failed: %s", message)
}
if chunk.BifrostResponsesStreamResponse == nil {
continue
}
response := chunk.BifrostResponsesStreamResponse
if response.Item != nil && response.Item.Type != nil && *response.Item.Type == schemas.ResponsesMessageTypeImageGenerationCall {
sawImageCall = true
}
if response.Type == schemas.ResponsesStreamResponseTypeCompleted {
sawCompleted = true
}
case <-streamCtx.Done():
t.Fatalf("timed out waiting for Responses image_generation stream: %v", streamCtx.Err())
}
}
})
})
}

func hasImageGenerationCall(output []schemas.ResponsesMessage) bool {
for _, item := range output {
if item.Type != nil && *item.Type == schemas.ResponsesMessageTypeImageGenerationCall {
return true
}
}
return false
}
1 change: 1 addition & 0 deletions core/internal/llmtests/tests.go
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@ func RunAllComprehensiveTests(t *testing.T, client *bifrost.Bifrost, ctx context
RunSimpleChatTest,
RunChatCompletionStreamTest,
RunResponsesStreamTest,
RunResponsesImageGenerationToolTest,
RunMultiTurnConversationTest,
RunToolCallsTest,
RunToolCallsWithEmptyPropertiesTest,
Expand Down
137 changes: 137 additions & 0 deletions core/providers/openai/imagegeneration_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,137 @@
package openai

import (
"encoding/json"
"fmt"
"io"
"net/http"
"net/http/httptest"
"testing"

"github.com/stretchr/testify/require"

"github.com/maximhq/bifrost/core/schemas"
)

func imageGenerationResponsesRequest() *schemas.BifrostResponsesRequest {
return &schemas.BifrostResponsesRequest{
Provider: schemas.OpenAI,
Model: "gpt-5.6-luna",
Input: []schemas.ResponsesMessage{{
Role: schemas.Ptr(schemas.ResponsesInputMessageRoleUser),
Content: &schemas.ResponsesMessageContent{ContentStr: schemas.Ptr("Generate a small orange cat in a coffee shop.")},
}},
Params: &schemas.ResponsesParameters{
Tools: []schemas.ResponsesTool{{
Type: schemas.ResponsesToolTypeImageGeneration,
ResponsesToolImageGeneration: &schemas.ResponsesToolImageGeneration{
Action: schemas.Ptr(schemas.ResponsesImageGenerationActionGenerate),
},
}},
},
}
}

func TestResponsesImageGenerationUnaryProviderIntegration(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
body, err := io.ReadAll(r.Body)
if err != nil {
t.Errorf("failed to read request body: %v", err)
w.WriteHeader(http.StatusInternalServerError)
return
}
var request struct {
Tools []struct {
Type string `json:"type"`
Action string `json:"action"`
} `json:"tools"`
}
if err := json.Unmarshal(body, &request); err != nil {
t.Errorf("failed to decode request body: %v", err)
} else if len(request.Tools) != 1 || request.Tools[0].Type != "image_generation" || request.Tools[0].Action != "generate" {
t.Errorf("unexpected image_generation tool payload: %s", body)
}
w.Header().Set("Content-Type", "application/json")
_, _ = fmt.Fprint(w, `{"id":"resp_unary","object":"response","status":"completed","output":[{"type":"image_generation_call","id":"ig_unary","status":"completed","action":"generate","result":"aGVsbG8="}]}`)
}))
defer server.Close()

provider := NewOpenAIProvider(&schemas.ProviderConfig{NetworkConfig: schemas.NetworkConfig{BaseURL: server.URL}}, testNoopLogger{})
response, bifrostErr := provider.Responses(newStreamTestContext(), testKey(), imageGenerationResponsesRequest())
require.Nil(t, bifrostErr)
require.NotNil(t, response)
require.Len(t, response.Output, 1)
require.NotNil(t, response.Output[0].ResponsesImageGenerationCall)
require.Equal(t, "aGVsbG8=", response.Output[0].ResponsesImageGenerationCall.Result)
require.NotNil(t, response.Output[0].Action)
require.NotNil(t, response.Output[0].Action.ResponsesImageGenerationToolCallAction)
require.Equal(t, schemas.ResponsesImageGenerationActionGenerate, response.Output[0].Action.ResponsesImageGenerationToolCallAction.Type)
encodedResponse, err := json.Marshal(response)
require.NoError(t, err)
require.Contains(t, string(encodedResponse), `"action":"generate"`)
}

func TestResponsesImageGenerationStreamingProviderIntegration(t *testing.T) {
streamBody := "event: response.image_generation_call.in_progress\n" +
`data: {"type":"response.image_generation_call.in_progress","sequence_number":1,"item":{"type":"image_generation_call","id":"ig_stream","status":"in_progress","action":"generate"}}` + "\n\n" +
"event: response.image_generation_call.generating\n" +
`data: {"type":"response.image_generation_call.generating","sequence_number":2,"item":{"type":"image_generation_call","id":"ig_stream","status":"generating","action":"generate"}}` + "\n\n" +
"event: response.image_generation_call.partial_image\n" +
`data: {"type":"response.image_generation_call.partial_image","sequence_number":3,"item_id":"ig_stream","partial_image_b64":"aGk=","partial_image_index":0}` + "\n\n" +
// OpenAI's live stream returns the final base64 image in output_item.done,
// rather than emitting a response.image_generation_call.completed event.
"event: response.output_item.done\n" +
`data: {"type":"response.output_item.done","sequence_number":4,"item":{"type":"image_generation_call","id":"ig_stream","status":"generating","action":"generate","result":"aGVsbG8="}}` + "\n\n" +
"event: response.completed\n" +
`data: {"type":"response.completed","sequence_number":5,"response":{"id":"resp_stream","object":"response","status":"completed","output":[{"type":"image_generation_call","id":"ig_stream","status":"generating","action":"generate","result":"aGVsbG8="}]}}` + "\n\n"

server := completeSSEServer(t, streamBody)
defer server.Close()
provider := newStreamTestProvider(server.URL)
stream, bifrostErr := provider.ResponsesStream(newStreamTestContext(), passthroughPostHook, nil, testKey(), imageGenerationResponsesRequest())
require.Nil(t, bifrostErr)

wantedEvents := map[schemas.ResponsesStreamResponseType]bool{
schemas.ResponsesStreamResponseTypeImageGenerationCallInProgress: false,
schemas.ResponsesStreamResponseTypeImageGenerationCallGenerating: false,
schemas.ResponsesStreamResponseTypeImageGenerationCallPartialImage: false,
schemas.ResponsesStreamResponseTypeOutputItemDone: false,
}
var completed *schemas.BifrostResponsesStreamResponse
for _, chunk := range collectChunks(t, stream) {
require.Nil(t, chunk.BifrostError)
if chunk.BifrostResponsesStreamResponse == nil {
continue
}
response := chunk.BifrostResponsesStreamResponse
if _, ok := wantedEvents[response.Type]; ok {
wantedEvents[response.Type] = true
}
if response.Type == schemas.ResponsesStreamResponseTypeImageGenerationCallInProgress ||
response.Type == schemas.ResponsesStreamResponseTypeImageGenerationCallGenerating ||
response.Type == schemas.ResponsesStreamResponseTypeOutputItemDone {
require.NotNil(t, response.Item)
require.NotNil(t, response.Item.Action)
require.NotNil(t, response.Item.Action.ResponsesImageGenerationToolCallAction)
encodedEvent, err := json.Marshal(response)
require.NoError(t, err)
require.Contains(t, string(encodedEvent), `"action":"generate"`)
if response.Type == schemas.ResponsesStreamResponseTypeOutputItemDone {
require.NotNil(t, response.Item.ResponsesImageGenerationCall)
require.Equal(t, "aGVsbG8=", response.Item.ResponsesImageGenerationCall.Result)
}
}
if response.Type == schemas.ResponsesStreamResponseTypeCompleted {
completed = response
}
}

for eventType, seen := range wantedEvents {
require.True(t, seen, "missing image-generation stream event %s", eventType)
}
require.NotNil(t, completed)
require.NotNil(t, completed.Response)
require.Len(t, completed.Response.Output, 1)
require.NotNil(t, completed.Response.Output[0].ResponsesImageGenerationCall)
require.Equal(t, "aGVsbG8=", completed.Response.Output[0].ResponsesImageGenerationCall.Result)
}
13 changes: 9 additions & 4 deletions core/providers/openai/responses_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -1205,16 +1205,21 @@ func TestResponsesToolMessageActionStruct_EdgeCases(t *testing.T) {
}
})

t.Run("unknown action type - unmarshal to computer tool (default)", func(t *testing.T) {
t.Run("unknown action type - preserve without computer coercion", func(t *testing.T) {
jsonData := `{"type":"unknown_action"}`
var action schemas.ResponsesToolMessageActionStruct
if err := json.Unmarshal([]byte(jsonData), &action); err != nil {
t.Fatalf("failed to unmarshal: %v", err)
}

// Default behavior is to unmarshal to computer tool
if action.ResponsesComputerToolCallAction == nil {
t.Error("expected ResponsesComputerToolCallAction to be populated for unknown type")
if action.ResponsesComputerToolCallAction != nil {
t.Error("unknown provider action must not be coerced to a computer action")
}
if len(action.Raw) == 0 {
t.Fatal("expected unknown action to be preserved in the raw action field")
}
if string(action.Raw) != jsonData {
t.Errorf("raw action mismatch: got %s", action.Raw)
}
})

Expand Down
Loading