move image generation behind provider capability

This commit is contained in:
Anton Bogdanovich 2026-05-05 13:54:53 -07:00
parent 32d4029f2b
commit cd0b4f025e
12 changed files with 479 additions and 548 deletions

View file

@ -418,6 +418,30 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
}
}
// CreateImageGenerationProviderFromModel creates a provider for image generation
// from a provider-prefixed model string. It returns the provider plus the model
// identifier stripped of the provider prefix.
func CreateImageGenerationProviderFromModel(model string) (ImageGenerationCapable, string, error) {
providerName, modelID := ExtractProtocol(&config.ModelConfig{Model: model})
if modelID == "" {
modelID = "gpt-image-2"
}
switch providerName {
case "", "openai", "openai-codex":
provider, err := createCodexAuthProvider()
if err != nil {
return nil, "", err
}
imageProvider, ok := provider.(ImageGenerationCapable)
if !ok || !imageProvider.SupportsImageGeneration() {
return nil, "", fmt.Errorf("provider %q does not support image generation", providerName)
}
return imageProvider, modelID, nil
default:
return nil, "", fmt.Errorf("provider %q does not support image generation", providerName)
}
}
func finalizeProviderFromConfig(
provider LLMProvider,
modelID string,

View file

@ -104,6 +104,32 @@ func TestCreateProviderReturnsClaudeProviderForAnthropicOAuth(t *testing.T) {
// TODO: Test custom APIBase when createClaudeAuthProvider supports it
}
func TestCreateImageGenerationProviderFromModelUsesCodexOAuth(t *testing.T) {
originalGetCredential := getCredential
t.Cleanup(func() { getCredential = originalGetCredential })
getCredential = func(provider string) (*auth.AuthCredential, error) {
if provider != "openai" {
t.Fatalf("provider = %q, want openai", provider)
}
return &auth.AuthCredential{
AccessToken: "openai-token",
AccountID: "acct-123",
}, nil
}
provider, model, err := CreateImageGenerationProviderFromModel("openai/gpt-image-2")
if err != nil {
t.Fatalf("CreateImageGenerationProviderFromModel() error = %v", err)
}
if model != "gpt-image-2" {
t.Fatalf("model = %q, want gpt-image-2", model)
}
if provider.ImageGenerationProviderID() != "openai-codex" {
t.Fatalf("provider id = %q, want openai-codex", provider.ImageGenerationProviderID())
}
}
func TestCreateProviderReturnsCodexProviderForOpenAIOAuth(t *testing.T) {
// TODO: This test requires openai protocol to support auth_method: "oauth"
// which is not yet implemented in the new factory_provider.go

View file

@ -0,0 +1,229 @@
package oauthprovider
import (
"context"
"encoding/base64"
"encoding/json"
"fmt"
"strings"
"github.com/openai/openai-go/v3"
"github.com/openai/openai-go/v3/option"
"github.com/openai/openai-go/v3/responses"
)
const (
codexDefaultImageGenerationModel = "gpt-image-2"
maxImageGenerationResults = 4
maxImageGenerationSSEBytes = 64 * 1024 * 1024
maxImageGenerationEvents = 512
)
func (p *CodexProvider) SupportsImageGeneration() bool {
return true
}
func (p *CodexProvider) ImageGenerationProviderID() string {
return "openai-codex"
}
func (p *CodexProvider) DefaultImageGenerationModel() string {
return codexDefaultImageGenerationModel
}
func (p *CodexProvider) GenerateImage(
ctx context.Context,
req ImageGenerationRequest,
) (*ImageGenerationResponse, error) {
opts, accountID, err := p.requestOptions()
if err != nil {
return nil, err
}
if accountID == "" {
return nil, fmt.Errorf("no account id found for Codex image generation")
}
if strings.TrimSpace(req.Model) == "" {
req.Model = p.DefaultImageGenerationModel()
}
if req.Count < 1 {
req.Count = 1
}
if req.Count > maxImageGenerationResults {
req.Count = maxImageGenerationResults
}
images := make([]GeneratedImage, 0, req.Count)
for i := 0; i < req.Count; i++ {
params := buildCodexImageParams(req)
stream := p.client.Responses.NewStreaming(ctx, params, opts...)
eventImages, readErr := parseCodexImageSSE(stream, req.OutputFormat)
closeErr := stream.Close()
if readErr != nil {
return nil, readErr
}
if closeErr != nil {
return nil, closeErr
}
images = append(images, eventImages...)
}
if len(images) > maxImageGenerationResults {
images = images[:maxImageGenerationResults]
}
return &ImageGenerationResponse{Images: images}, nil
}
func (p *CodexProvider) requestOptions() ([]option.RequestOption, string, error) {
var opts []option.RequestOption
accountID := p.accountID
if p.tokenSource != nil {
tok, accID, err := p.tokenSource()
if err != nil {
return nil, "", fmt.Errorf("refreshing token: %w", err)
}
opts = append(opts, option.WithAPIKey(tok))
if accID != "" {
accountID = accID
}
}
if accountID != "" {
opts = append(opts, option.WithHeader("Chatgpt-Account-Id", accountID))
}
return opts, accountID, nil
}
func buildCodexImageParams(req ImageGenerationRequest) responses.ResponseNewParams {
tool := responses.ToolUnionParam{OfImageGeneration: &responses.ToolImageGenerationParam{
Model: req.Model,
Size: req.Size,
}}
if req.Quality != "" {
tool.OfImageGeneration.Quality = req.Quality
}
if req.OutputFormat != "" {
tool.OfImageGeneration.OutputFormat = req.OutputFormat
}
return responses.ResponseNewParams{
Model: "gpt-5.4",
Input: responses.ResponseNewParamsInputUnion{
OfString: openai.Opt(req.Prompt),
},
Instructions: openai.Opt("You are an image generation assistant."),
Tools: []responses.ToolUnionParam{tool},
ToolChoice: responses.ResponseNewParamsToolChoiceUnion{
OfHostedTool: &responses.ToolChoiceTypesParam{
Type: responses.ToolChoiceTypesTypeImageGeneration,
},
},
Store: openai.Opt(false),
}
}
type codexImageStream interface {
Next() bool
Current() responses.ResponseStreamEventUnion
Err() error
}
func parseCodexImageSSE(stream codexImageStream, outputFormat string) ([]GeneratedImage, error) {
var totalBytes int
var events int
var images []GeneratedImage
var completedImages []GeneratedImage
for stream.Next() {
evt := stream.Current()
events++
if events > maxImageGenerationEvents {
return nil, fmt.Errorf("codex image response exceeded event limit")
}
data, err := json.Marshal(evt)
if err == nil {
totalBytes += len(data)
if totalBytes > maxImageGenerationSSEBytes {
return nil, fmt.Errorf("codex image response exceeded size limit")
}
}
eventImages, eventCompletedImages, parseErr := parseCodexImageEventUnion(evt, outputFormat)
if parseErr != nil {
return nil, parseErr
}
images = append(images, eventImages...)
completedImages = append(completedImages, eventCompletedImages...)
}
if err := stream.Err(); err != nil {
return nil, err
}
if len(images) > 0 {
return images, nil
}
return completedImages, nil
}
func parseCodexImageEventUnion(
evt responses.ResponseStreamEventUnion,
outputFormat string,
) ([]GeneratedImage, []GeneratedImage, error) {
switch evt.Type {
case "response.output_item.done":
if image, ok, err := imageFromCodexItemUnion(evt.Item, outputFormat); err != nil {
return nil, nil, err
} else if ok {
return []GeneratedImage{image}, nil, nil
}
case "response.completed":
images := make([]GeneratedImage, 0)
for _, item := range evt.Response.Output {
if image, ok, err := imageFromCodexResponseItem(item, outputFormat); err != nil {
return nil, nil, err
} else if ok {
images = append(images, image)
}
}
return nil, images, nil
case "response.failed", "error":
return nil, nil, fmt.Errorf("codex image generation failed")
}
return nil, nil, nil
}
func imageFromCodexItemUnion(
item responses.ResponseOutputItemUnion,
outputFormat string,
) (GeneratedImage, bool, error) {
if item.Type != "image_generation_call" {
return GeneratedImage{}, false, nil
}
return imageFromCodexPayload(item.Result, outputFormat)
}
func imageFromCodexResponseItem(
item responses.ResponseOutputItemUnion,
outputFormat string,
) (GeneratedImage, bool, error) {
return imageFromCodexItemUnion(item, outputFormat)
}
func imageFromCodexPayload(payload string, outputFormat string) (GeneratedImage, bool, error) {
if payload == "" {
return GeneratedImage{}, false, nil
}
data, err := base64.StdEncoding.DecodeString(payload)
if err != nil {
return GeneratedImage{}, false, err
}
mime, ext := imageMimeAndExtension(outputFormat)
return GeneratedImage{Data: data, MimeType: mime, Ext: ext}, true, nil
}
func imageMimeAndExtension(outputFormat string) (string, string) {
switch strings.ToLower(strings.TrimSpace(outputFormat)) {
case "jpeg", "jpg":
return "image/jpeg", "jpg"
case "webp":
return "image/webp", "webp"
default:
return "image/png", "png"
}
}

View file

@ -0,0 +1,93 @@
package oauthprovider
import (
"encoding/base64"
"testing"
"github.com/openai/openai-go/v3/responses"
)
type mockCodexImageStream struct {
events []responses.ResponseStreamEventUnion
index int
err error
}
func (s *mockCodexImageStream) Next() bool {
if s.index >= len(s.events) {
return false
}
s.index++
return true
}
func (s *mockCodexImageStream) Current() responses.ResponseStreamEventUnion {
return s.events[s.index-1]
}
func (s *mockCodexImageStream) Err() error { return s.err }
func TestCodexProviderSupportsImageGeneration(t *testing.T) {
provider := NewCodexProvider("test-token", "acct-123")
if !provider.SupportsImageGeneration() {
t.Fatal("SupportsImageGeneration = false, want true")
}
if provider.ImageGenerationProviderID() != "openai-codex" {
t.Fatalf("provider id = %q, want openai-codex", provider.ImageGenerationProviderID())
}
if provider.DefaultImageGenerationModel() != "gpt-image-2" {
t.Fatalf("default image model = %q, want gpt-image-2", provider.DefaultImageGenerationModel())
}
}
func TestBuildCodexImageParams(t *testing.T) {
params := buildCodexImageParams(ImageGenerationRequest{
Prompt: "make a tiny icon",
Model: "gpt-image-2",
Size: "1536x1024",
Quality: "medium",
OutputFormat: "png",
})
if params.Model != "gpt-5.4" {
t.Fatalf("request model = %q, want gpt-5.4", params.Model)
}
if len(params.Tools) != 1 || params.Tools[0].OfImageGeneration == nil {
t.Fatalf("expected one image_generation tool, got %#v", params.Tools)
}
tool := params.Tools[0].OfImageGeneration
if tool.Model != "gpt-image-2" {
t.Fatalf("image model = %q, want gpt-image-2", tool.Model)
}
if tool.Size != "1536x1024" {
t.Fatalf("size = %q, want 1536x1024", tool.Size)
}
if tool.Quality != "medium" {
t.Fatalf("quality = %q, want medium", tool.Quality)
}
}
func TestParseCodexImageSSECompletedResponseFallback(t *testing.T) {
payload := base64.StdEncoding.EncodeToString([]byte("fake-png"))
stream := &mockCodexImageStream{
events: []responses.ResponseStreamEventUnion{{
Type: "response.completed",
Response: responses.Response{
Output: []responses.ResponseOutputItemUnion{{
Type: "image_generation_call",
Result: payload,
}},
},
}},
}
images, err := parseCodexImageSSE(stream, "png")
if err != nil {
t.Fatalf("parseCodexImageSSE: %v", err)
}
if len(images) != 1 {
t.Fatalf("images = %d, want 1", len(images))
}
if string(images[0].Data) != "fake-png" {
t.Fatalf("image data = %q, want fake-png", string(images[0].Data))
}
}

View file

@ -58,8 +58,6 @@ func NewCodexProviderWithTokenSource(
func (p *CodexProvider) Chat(
ctx context.Context, messages []Message, tools []ToolDefinition, model string, options map[string]any,
) (*LLMResponse, error) {
var opts []option.RequestOption
accountID := p.accountID
resolvedModel, fallbackReason := resolveCodexModel(model)
if fallbackReason != "" {
logger.WarnCF(
@ -72,18 +70,11 @@ func (p *CodexProvider) Chat(
},
)
}
if p.tokenSource != nil {
tok, accID, err := p.tokenSource()
if err != nil {
return nil, fmt.Errorf("refreshing token: %w", err)
}
opts = append(opts, option.WithAPIKey(tok))
if accID != "" {
accountID = accID
}
opts, accountID, err := p.requestOptions()
if err != nil {
return nil, err
}
if accountID != "" {
opts = append(opts, option.WithHeader("Chatgpt-Account-Id", accountID))
} else {
logger.WarnCF(
"provider.codex",
@ -114,7 +105,7 @@ func (p *CodexProvider) Chat(
}
}
}
err := stream.Err()
err = stream.Err()
if err != nil {
fields := map[string]any{
"requested_model": model,

View file

@ -7,17 +7,20 @@ import (
)
type (
ToolCall = protocoltypes.ToolCall
FunctionCall = protocoltypes.FunctionCall
LLMResponse = protocoltypes.LLMResponse
UsageInfo = protocoltypes.UsageInfo
Message = protocoltypes.Message
ToolDefinition = protocoltypes.ToolDefinition
ToolFunctionDefinition = protocoltypes.ToolFunctionDefinition
ExtraContent = protocoltypes.ExtraContent
GoogleExtra = protocoltypes.GoogleExtra
ContentBlock = protocoltypes.ContentBlock
CacheControl = protocoltypes.CacheControl
ToolCall = protocoltypes.ToolCall
FunctionCall = protocoltypes.FunctionCall
LLMResponse = protocoltypes.LLMResponse
UsageInfo = protocoltypes.UsageInfo
Message = protocoltypes.Message
ToolDefinition = protocoltypes.ToolDefinition
ToolFunctionDefinition = protocoltypes.ToolFunctionDefinition
ExtraContent = protocoltypes.ExtraContent
GoogleExtra = protocoltypes.GoogleExtra
ContentBlock = protocoltypes.ContentBlock
CacheControl = protocoltypes.CacheControl
ImageGenerationRequest = protocoltypes.ImageGenerationRequest
GeneratedImage = protocoltypes.GeneratedImage
ImageGenerationResponse = protocoltypes.ImageGenerationResponse
)
type LLMProvider interface {

View file

@ -78,6 +78,25 @@ type Attachment struct {
ContentType string `json:"content_type,omitempty"`
}
type ImageGenerationRequest struct {
Prompt string
Model string
Size string
Quality string
OutputFormat string
Count int
}
type GeneratedImage struct {
Data []byte
MimeType string
Ext string
}
type ImageGenerationResponse struct {
Images []GeneratedImage
}
type Message struct {
Role string `json:"role"`
Content string `json:"content"`

View file

@ -8,18 +8,21 @@ import (
)
type (
ToolCall = protocoltypes.ToolCall
FunctionCall = protocoltypes.FunctionCall
LLMResponse = protocoltypes.LLMResponse
UsageInfo = protocoltypes.UsageInfo
Message = protocoltypes.Message
ToolDefinition = protocoltypes.ToolDefinition
ToolFunctionDefinition = protocoltypes.ToolFunctionDefinition
ExtraContent = protocoltypes.ExtraContent
GoogleExtra = protocoltypes.GoogleExtra
ContentBlock = protocoltypes.ContentBlock
CacheControl = protocoltypes.CacheControl
Attachment = protocoltypes.Attachment
ToolCall = protocoltypes.ToolCall
FunctionCall = protocoltypes.FunctionCall
LLMResponse = protocoltypes.LLMResponse
UsageInfo = protocoltypes.UsageInfo
Message = protocoltypes.Message
ToolDefinition = protocoltypes.ToolDefinition
ToolFunctionDefinition = protocoltypes.ToolFunctionDefinition
ExtraContent = protocoltypes.ExtraContent
GoogleExtra = protocoltypes.GoogleExtra
ContentBlock = protocoltypes.ContentBlock
CacheControl = protocoltypes.CacheControl
Attachment = protocoltypes.Attachment
ImageGenerationRequest = protocoltypes.ImageGenerationRequest
GeneratedImage = protocoltypes.GeneratedImage
ImageGenerationResponse = protocoltypes.ImageGenerationResponse
)
type LLMProvider interface {
@ -68,6 +71,17 @@ type NativeSearchCapable interface {
SupportsNativeSearch() bool
}
// ImageGenerationCapable is an optional interface for providers that can
// generate raster images outside the normal chat inference loop. Core tools can
// use this to reuse provider-owned auth/client behavior while keeping channel
// media delivery in the agent runtime.
type ImageGenerationCapable interface {
SupportsImageGeneration() bool
ImageGenerationProviderID() string
DefaultImageGenerationModel() string
GenerateImage(ctx context.Context, req ImageGenerationRequest) (*ImageGenerationResponse, error)
}
// FailoverReason classifies why an LLM request failed for fallback decisions.
type FailoverReason string

View file

@ -11,6 +11,7 @@ import (
"github.com/google/uuid"
"github.com/sipeed/picoclaw/pkg/media"
"github.com/sipeed/picoclaw/pkg/providers"
)
const (
@ -23,21 +24,16 @@ const (
type ImageGenerateTool struct {
workspace string
model string
provider imageGenerationProvider
provider providers.ImageGenerationCapable
resolver ImageGenerationProviderResolver
mediaStore media.MediaStore
}
type ImageGenerateToolOption func(*ImageGenerateTool)
type imageGenerationProvider interface {
ID() string
DefaultModel() string
GenerateImages(ctx context.Context, req imageGenerationRequest) ([]generatedImage, error)
}
type ImageGenerationProviderResolver func(model string) (providers.ImageGenerationCapable, string, error)
type imageGenerationProviderFactory func() imageGenerationProvider
func WithImageGenerationProvider(provider imageGenerationProvider) ImageGenerateToolOption {
func WithImageGenerationProvider(provider providers.ImageGenerationCapable) ImageGenerateToolOption {
return func(t *ImageGenerateTool) {
if provider != nil {
t.provider = provider
@ -45,25 +41,24 @@ func WithImageGenerationProvider(provider imageGenerationProvider) ImageGenerate
}
}
func WithImageGenerationProviderResolver(resolver ImageGenerationProviderResolver) ImageGenerateToolOption {
return func(t *ImageGenerateTool) {
if resolver != nil {
t.resolver = resolver
}
}
}
func NewImageGenerateTool(
workspace string,
model string,
store media.MediaStore,
options ...ImageGenerateToolOption,
) *ImageGenerateTool {
spec := parseImageGenerationModel(model)
factory := imageGenerationProviderFactories[spec.Provider]
if factory == nil {
factory = imageGenerationProviderFactories[defaultImageGenerationProvider]
}
provider := factory()
if spec.Model == "" && provider != nil {
spec.Model = provider.DefaultModel()
}
tool := &ImageGenerateTool{
workspace: workspace,
model: spec.Model,
provider: provider,
model: model,
resolver: providers.CreateImageGenerationProviderFromModel,
mediaStore: store,
}
for _, option := range options {
@ -124,11 +119,19 @@ func (t *ImageGenerateTool) Execute(ctx context.Context, args map[string]any) *T
if t.mediaStore == nil {
return ErrorResult("media store not configured")
}
if t.provider == nil && t.resolver != nil {
provider, model, err := t.resolver(t.model)
if err != nil {
return ErrorResult(fmt.Sprintf("image generation provider not configured: %v", err)).WithError(err)
}
t.provider = provider
t.model = model
}
if t.provider == nil {
return ErrorResult("image generation provider not configured")
}
req := imageGenerationRequest{
req := providers.ImageGenerationRequest{
Prompt: prompt,
Model: t.model,
Size: readStringDefault(args, "size", defaultImageGenerationSize),
@ -137,12 +140,16 @@ func (t *ImageGenerateTool) Execute(ctx context.Context, args map[string]any) *T
Count: readImageCount(args["count"]),
}
if strings.TrimSpace(req.Model) == "" {
req.Model = t.provider.DefaultModel()
req.Model = t.provider.DefaultImageGenerationModel()
}
images, err := t.provider.GenerateImages(ctx, req)
resp, err := t.provider.GenerateImage(ctx, req)
if err != nil {
return ErrorResult(fmt.Sprintf("image generation failed: %v", err)).WithError(err)
}
if resp == nil {
return ErrorResult("image generation returned no response")
}
images := resp.Images
if len(images) == 0 {
return ErrorResult("image generation returned no images")
}
@ -168,7 +175,7 @@ func (t *ImageGenerateTool) Execute(ctx context.Context, args map[string]any) *T
paths = append(paths, path)
}
message := fmt.Sprintf("Generated %d image(s) with %s via %s.", len(refs), req.Model, t.provider.ID())
message := fmt.Sprintf("Generated %d image(s) with %s via %s.", len(refs), req.Model, t.provider.ImageGenerationProviderID())
result := MediaResult(message, refs).WithResponseHandled()
result.ArtifactTags = make([]string, 0, len(paths))
for _, path := range paths {
@ -177,22 +184,7 @@ func (t *ImageGenerateTool) Execute(ctx context.Context, args map[string]any) *T
return result
}
type imageGenerationRequest struct {
Prompt string
Model string
Size string
Quality string
OutputFormat string
Count int
}
type generatedImage struct {
Data []byte
MimeType string
Ext string
}
func writeGeneratedImage(image generatedImage, index int) (string, error) {
func writeGeneratedImage(image providers.GeneratedImage, index int) (string, error) {
dir, err := os.MkdirTemp("", "picoclaw-image-generate-*")
if err != nil {
return "", err
@ -248,14 +240,3 @@ func readImageCount(raw any) int {
}
return count
}
func imageMimeAndExtension(outputFormat string) (string, string) {
switch strings.ToLower(strings.TrimSpace(outputFormat)) {
case "jpeg", "jpg":
return "image/jpeg", "jpg"
case "webp":
return "image/webp", "webp"
default:
return "image/png", "png"
}
}

View file

@ -1,262 +0,0 @@
package tools
import (
"bufio"
"bytes"
"context"
"encoding/base64"
"encoding/json"
"fmt"
"io"
"net/http"
"strings"
"time"
oauthprovider "github.com/sipeed/picoclaw/pkg/providers/oauth"
)
const (
defaultImageGenerationBaseURL = "https://chatgpt.com/backend-api/codex"
defaultImageGenerationTimeout = 180 * time.Second
maxImageGenerationSSEBytes = 64 * 1024 * 1024
maxImageGenerationEvents = 512
)
type openAICodexImageGenerationProvider struct {
baseURL string
timeout time.Duration
httpClient *http.Client
tokenSource func() (accessToken, accountID string, err error)
}
func newOpenAICodexImageGenerationProvider() *openAICodexImageGenerationProvider {
return &openAICodexImageGenerationProvider{
baseURL: defaultImageGenerationBaseURL,
timeout: defaultImageGenerationTimeout,
httpClient: http.DefaultClient,
tokenSource: oauthprovider.CreateCodexTokenSource(),
}
}
func WithImageGenerateBaseURL(baseURL string) ImageGenerateToolOption {
return func(t *ImageGenerateTool) {
if provider, ok := t.provider.(*openAICodexImageGenerationProvider); ok {
provider.baseURL = baseURL
}
}
}
func WithImageGenerateHTTPClient(client *http.Client) ImageGenerateToolOption {
return func(t *ImageGenerateTool) {
if provider, ok := t.provider.(*openAICodexImageGenerationProvider); ok {
provider.httpClient = client
}
}
}
func WithImageGenerateTokenSource(source func() (string, string, error)) ImageGenerateToolOption {
return func(t *ImageGenerateTool) {
if provider, ok := t.provider.(*openAICodexImageGenerationProvider); ok {
provider.tokenSource = source
}
}
}
func (p *openAICodexImageGenerationProvider) ID() string { return "openai-codex" }
func (p *openAICodexImageGenerationProvider) DefaultModel() string {
return defaultImageGenerationModel
}
func (p *openAICodexImageGenerationProvider) GenerateImages(
ctx context.Context,
req imageGenerationRequest,
) ([]generatedImage, error) {
accessToken, accountID, err := p.tokenSource()
if err != nil {
return nil, fmt.Errorf("OpenAI/Codex OAuth not configured: %w", err)
}
timeout := p.timeout
if timeout <= 0 {
timeout = defaultImageGenerationTimeout
}
callCtx, cancel := context.WithTimeout(ctx, timeout)
defer cancel()
images := make([]generatedImage, 0, req.Count)
for i := 0; i < req.Count; i++ {
body, err := json.Marshal(buildCodexImageRequest(req))
if err != nil {
return nil, err
}
httpReq, err := http.NewRequestWithContext(
callCtx,
http.MethodPost,
strings.TrimRight(p.baseURL, "/")+"/responses",
bytes.NewReader(body),
)
if err != nil {
return nil, err
}
httpReq.Header.Set("Authorization", "Bearer "+accessToken)
httpReq.Header.Set("Content-Type", "application/json")
httpReq.Header.Set("Accept", "text/event-stream")
httpReq.Header.Set("originator", "codex_cli_rs")
httpReq.Header.Set("OpenAI-Beta", "responses=experimental")
if strings.TrimSpace(accountID) != "" {
httpReq.Header.Set("Chatgpt-Account-Id", accountID)
}
resp, err := p.httpClient.Do(httpReq)
if err != nil {
return nil, err
}
eventImages, readErr := parseCodexImageSSE(resp.Body, req.OutputFormat)
closeErr := resp.Body.Close()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return nil, fmt.Errorf("codex image request failed: HTTP %d", resp.StatusCode)
}
if readErr != nil {
return nil, readErr
}
if closeErr != nil {
return nil, closeErr
}
images = append(images, eventImages...)
}
if len(images) > maxImageGenerationResults {
images = images[:maxImageGenerationResults]
}
return images, nil
}
func buildCodexImageRequest(req imageGenerationRequest) map[string]any {
tool := map[string]any{
"type": "image_generation",
"model": req.Model,
"size": req.Size,
}
if req.Quality != "" {
tool["quality"] = req.Quality
}
if req.OutputFormat != "" {
tool["output_format"] = req.OutputFormat
}
return map[string]any{
"model": "gpt-5.4",
"input": []map[string]any{{
"role": "user",
"content": []map[string]any{{
"type": "input_text",
"text": req.Prompt,
}},
}},
"instructions": "You are an image generation assistant.",
"tools": []map[string]any{tool},
"tool_choice": map[string]any{"type": "image_generation"},
"stream": true,
"store": false,
}
}
func parseCodexImageSSE(r io.Reader, outputFormat string) ([]generatedImage, error) {
reader := bufio.NewReader(r)
var totalBytes int
var events int
var images []generatedImage
var completedImages []generatedImage
for {
line, err := reader.ReadString('\n')
if len(line) > 0 {
totalBytes += len(line)
if totalBytes > maxImageGenerationSSEBytes {
return nil, fmt.Errorf("codex image response exceeded size limit")
}
if strings.HasPrefix(line, "data: ") {
events++
if events > maxImageGenerationEvents {
return nil, fmt.Errorf("codex image response exceeded event limit")
}
eventImages, eventCompletedImages, parseErr := parseCodexImageEvent(
strings.TrimSpace(strings.TrimPrefix(line, "data: ")),
outputFormat,
)
if parseErr != nil {
return nil, parseErr
}
images = append(images, eventImages...)
completedImages = append(completedImages, eventCompletedImages...)
}
}
if err == io.EOF {
break
}
if err != nil {
return nil, err
}
}
if len(images) > 0 {
return images, nil
}
return completedImages, nil
}
func parseCodexImageEvent(data string, outputFormat string) ([]generatedImage, []generatedImage, error) {
if data == "" || data == "[DONE]" {
return nil, nil, nil
}
var event map[string]any
if err := json.Unmarshal([]byte(data), &event); err != nil {
return nil, nil, nil
}
eventType, _ := event["type"].(string)
if eventType == "response.failed" || eventType == "error" {
return nil, nil, fmt.Errorf("codex image generation failed")
}
var itemImages []generatedImage
if eventType == "response.output_item.done" {
if item, _ := event["item"].(map[string]any); item != nil {
if image, ok, err := imageFromCodexItem(item, outputFormat); err != nil {
return nil, nil, err
} else if ok {
itemImages = append(itemImages, image)
}
}
}
var completedImages []generatedImage
if eventType == "response.completed" {
if response, _ := event["response"].(map[string]any); response != nil {
if output, _ := response["output"].([]any); output != nil {
for _, raw := range output {
item, _ := raw.(map[string]any)
if image, ok, err := imageFromCodexItem(item, outputFormat); err != nil {
return nil, nil, err
} else if ok {
completedImages = append(completedImages, image)
}
}
}
}
}
return itemImages, completedImages, nil
}
func imageFromCodexItem(item map[string]any, outputFormat string) (generatedImage, bool, error) {
if item == nil || item["type"] != "image_generation_call" {
return generatedImage{}, false, nil
}
payload, _ := item["result"].(string)
if payload == "" {
return generatedImage{}, false, nil
}
data, err := base64.StdEncoding.DecodeString(payload)
if err != nil {
return generatedImage{}, false, err
}
mime, ext := imageMimeAndExtension(outputFormat)
return generatedImage{Data: data, MimeType: mime, Ext: ext}, true, nil
}

View file

@ -1,45 +0,0 @@
package tools
import "strings"
const (
defaultImageGenerationProvider = "openai-codex"
defaultImageGenerationModel = "gpt-image-2"
)
var imageGenerationProviderFactories = map[string]imageGenerationProviderFactory{
defaultImageGenerationProvider: func() imageGenerationProvider {
return newOpenAICodexImageGenerationProvider()
},
}
type imageGenerationModelSpec struct {
Provider string
Model string
}
func parseImageGenerationModel(model string) imageGenerationModelSpec {
model = strings.TrimSpace(model)
if model == "" {
return imageGenerationModelSpec{
Provider: defaultImageGenerationProvider,
Model: defaultImageGenerationModel,
}
}
provider, modelName, ok := strings.Cut(model, "/")
if !ok || strings.TrimSpace(provider) == "" || strings.TrimSpace(modelName) == "" {
return imageGenerationModelSpec{
Provider: defaultImageGenerationProvider,
Model: model,
}
}
provider = strings.TrimSpace(provider)
modelName = strings.TrimSpace(modelName)
if provider == "openai" {
provider = defaultImageGenerationProvider
}
return imageGenerationModelSpec{
Provider: provider,
Model: modelName,
}
}

View file

@ -2,115 +2,34 @@ package tools
import (
"context"
"encoding/base64"
"encoding/json"
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/sipeed/picoclaw/pkg/media"
"github.com/sipeed/picoclaw/pkg/providers"
)
type fakeImageGenerationProvider struct {
id string
defaultModel string
request imageGenerationRequest
request providers.ImageGenerationRequest
}
func (p *fakeImageGenerationProvider) ID() string { return p.id }
func (p *fakeImageGenerationProvider) SupportsImageGeneration() bool { return true }
func (p *fakeImageGenerationProvider) DefaultModel() string { return p.defaultModel }
func (p *fakeImageGenerationProvider) ImageGenerationProviderID() string { return p.id }
func (p *fakeImageGenerationProvider) GenerateImages(
func (p *fakeImageGenerationProvider) DefaultImageGenerationModel() string { return p.defaultModel }
func (p *fakeImageGenerationProvider) GenerateImage(
_ context.Context,
req imageGenerationRequest,
) ([]generatedImage, error) {
req providers.ImageGenerationRequest,
) (*providers.ImageGenerationResponse, error) {
p.request = req
return []generatedImage{{
return &providers.ImageGenerationResponse{Images: []providers.GeneratedImage{{
Data: []byte("fake-image"),
MimeType: "image/png",
Ext: "png",
}}, nil
}
func TestImageGenerateToolCodexOAuthRequestAndMediaResult(t *testing.T) {
var captured map[string]any
var gotAuth string
var gotAccount string
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
gotAuth = r.Header.Get("Authorization")
gotAccount = r.Header.Get("Chatgpt-Account-Id")
if r.URL.Path != "/responses" {
t.Fatalf("path = %q, want /responses", r.URL.Path)
}
if err := json.NewDecoder(r.Body).Decode(&captured); err != nil {
t.Fatalf("decode request: %v", err)
}
payload := base64.StdEncoding.EncodeToString([]byte("fake-png"))
w.Header().Set("Content-Type", "text/event-stream")
_, _ = w.Write([]byte(`data: {"type":"response.output_item.done","item":{"type":"image_generation_call","result":"` + payload + `"}}` + "\n\n"))
_, _ = w.Write([]byte("data: [DONE]\n\n"))
}))
defer server.Close()
store := media.NewFileMediaStore()
tool := NewImageGenerateTool(
t.TempDir(),
"openai/gpt-image-2",
store,
WithImageGenerateBaseURL(server.URL),
WithImageGenerateTokenSource(func() (string, string, error) {
return "test-token", "acct-123", nil
}),
)
result := tool.Execute(
WithToolContext(t.Context(), "telegram", "chat-1"),
map[string]any{
"prompt": "make a tiny icon",
"size": "1536x1024",
"quality": "medium",
"output_format": "png",
},
)
if result.IsError {
t.Fatalf("Execute returned error: %s", result.ContentForLLM())
}
if !result.ResponseHandled {
t.Fatal("ResponseHandled = false, want true")
}
if len(result.Media) != 1 {
t.Fatalf("media refs = %d, want 1", len(result.Media))
}
path, err := store.Resolve(result.Media[0])
if err != nil {
t.Fatalf("resolve media: %v", err)
}
if gotAuth != "Bearer test-token" {
t.Fatalf("Authorization = %q, want Bearer test-token", gotAuth)
}
if gotAccount != "acct-123" {
t.Fatalf("Chatgpt-Account-Id = %q, want acct-123", gotAccount)
}
if captured["model"] != "gpt-5.4" {
t.Fatalf("request model = %v, want gpt-5.4", captured["model"])
}
toolsRaw := captured["tools"].([]any)
imageTool := toolsRaw[0].(map[string]any)
if imageTool["model"] != "gpt-image-2" {
t.Fatalf("image model = %v, want gpt-image-2", imageTool["model"])
}
if imageTool["size"] != "1536x1024" {
t.Fatalf("size = %v, want 1536x1024", imageTool["size"])
}
if imageTool["quality"] != "medium" {
t.Fatalf("quality = %v, want medium", imageTool["quality"])
}
if path == "" {
t.Fatal("generated media path is empty")
}
}}}, nil
}
func TestImageGenerateToolCanUseInjectedProvider(t *testing.T) {
@ -121,7 +40,7 @@ func TestImageGenerateToolCanUseInjectedProvider(t *testing.T) {
}
tool := NewImageGenerateTool(
t.TempDir(),
"test-provider/custom-image-model",
"custom-image-model",
store,
WithImageGenerationProvider(provider),
)
@ -140,64 +59,3 @@ func TestImageGenerateToolCanUseInjectedProvider(t *testing.T) {
t.Fatalf("media refs = %d, want 1", len(result.Media))
}
}
func TestParseImageGenerationModel(t *testing.T) {
tests := []struct {
name string
model string
wantProvider string
wantModel string
}{
{
name: "empty uses default provider and model",
model: "",
wantProvider: "openai-codex",
wantModel: "gpt-image-2",
},
{
name: "bare model uses default provider",
model: "custom-image-model",
wantProvider: "openai-codex",
wantModel: "custom-image-model",
},
{
name: "openai alias routes to codex oauth provider",
model: "openai/gpt-image-2",
wantProvider: "openai-codex",
wantModel: "gpt-image-2",
},
{
name: "future provider prefix is preserved",
model: "gemini/imagen-4",
wantProvider: "gemini",
wantModel: "imagen-4",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := parseImageGenerationModel(tt.model)
if got.Provider != tt.wantProvider {
t.Fatalf("provider = %q, want %q", got.Provider, tt.wantProvider)
}
if got.Model != tt.wantModel {
t.Fatalf("model = %q, want %q", got.Model, tt.wantModel)
}
})
}
}
func TestParseCodexImageSSECompletedResponseFallback(t *testing.T) {
payload := base64.StdEncoding.EncodeToString([]byte("fake-png"))
body := `data: {"type":"response.completed","response":{"output":[{"type":"image_generation_call","result":"` + payload + `"}]}}` + "\n\n"
images, err := parseCodexImageSSE(strings.NewReader(body), "png")
if err != nil {
t.Fatalf("parseCodexImageSSE: %v", err)
}
if len(images) != 1 {
t.Fatalf("images = %d, want 1", len(images))
}
if string(images[0].Data) != "fake-png" {
t.Fatalf("image data = %q, want fake-png", string(images[0].Data))
}
}