diff --git a/pkg/providers/factory_provider.go b/pkg/providers/factory_provider.go index a59e2de25..a61e34711 100644 --- a/pkg/providers/factory_provider.go +++ b/pkg/providers/factory_provider.go @@ -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, diff --git a/pkg/providers/factory_test.go b/pkg/providers/factory_test.go index b99f5baf9..d3c33b61b 100644 --- a/pkg/providers/factory_test.go +++ b/pkg/providers/factory_test.go @@ -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 diff --git a/pkg/providers/oauth/codex_image_generation.go b/pkg/providers/oauth/codex_image_generation.go new file mode 100644 index 000000000..43d07e5a2 --- /dev/null +++ b/pkg/providers/oauth/codex_image_generation.go @@ -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" + } +} diff --git a/pkg/providers/oauth/codex_image_generation_test.go b/pkg/providers/oauth/codex_image_generation_test.go new file mode 100644 index 000000000..a1323b83f --- /dev/null +++ b/pkg/providers/oauth/codex_image_generation_test.go @@ -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)) + } +} diff --git a/pkg/providers/oauth/codex_provider.go b/pkg/providers/oauth/codex_provider.go index 0b125997b..1fd3ca944 100644 --- a/pkg/providers/oauth/codex_provider.go +++ b/pkg/providers/oauth/codex_provider.go @@ -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, diff --git a/pkg/providers/oauth/types.go b/pkg/providers/oauth/types.go index 02ea4a21c..8828fc52d 100644 --- a/pkg/providers/oauth/types.go +++ b/pkg/providers/oauth/types.go @@ -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 { diff --git a/pkg/providers/protocoltypes/types.go b/pkg/providers/protocoltypes/types.go index bab4433e7..5e2f18ce6 100644 --- a/pkg/providers/protocoltypes/types.go +++ b/pkg/providers/protocoltypes/types.go @@ -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"` diff --git a/pkg/providers/types.go b/pkg/providers/types.go index 23406bc45..1a4f81b71 100644 --- a/pkg/providers/types.go +++ b/pkg/providers/types.go @@ -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 diff --git a/pkg/tools/image_generate.go b/pkg/tools/image_generate.go index 4967d3336..a2a59eab2 100644 --- a/pkg/tools/image_generate.go +++ b/pkg/tools/image_generate.go @@ -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" - } -} diff --git a/pkg/tools/image_generate_openai_codex.go b/pkg/tools/image_generate_openai_codex.go deleted file mode 100644 index b6a1cbe5d..000000000 --- a/pkg/tools/image_generate_openai_codex.go +++ /dev/null @@ -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 -} diff --git a/pkg/tools/image_generate_registry.go b/pkg/tools/image_generate_registry.go deleted file mode 100644 index 85ddfb8b2..000000000 --- a/pkg/tools/image_generate_registry.go +++ /dev/null @@ -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, - } -} diff --git a/pkg/tools/image_generate_test.go b/pkg/tools/image_generate_test.go index 52d8ac425..c9613568c 100644 --- a/pkg/tools/image_generate_test.go +++ b/pkg/tools/image_generate_test.go @@ -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)) - } -}