diff --git a/pkg/tools/image_generate.go b/pkg/tools/image_generate.go index 4ed1a33de..a74cc6827 100644 --- a/pkg/tools/image_generate.go +++ b/pkg/tools/image_generate.go @@ -21,16 +21,17 @@ import ( ) const ( - defaultImageGenerationModel = "gpt-image-2" - defaultImageGenerationBaseURL = "https://chatgpt.com/backend-api/codex" - defaultImageGenerationSize = "1024x1024" - defaultImageGenerationTimeout = 180 * time.Second - maxImageGenerationResults = 4 - maxImageGenerationSSEBytes = 64 * 1024 * 1024 - maxImageGenerationEvents = 512 + defaultImageGenerationProvider = "openai-codex" + defaultImageGenerationModel = "gpt-image-2" + defaultImageGenerationBaseURL = "https://chatgpt.com/backend-api/codex" + defaultImageGenerationSize = "1024x1024" + defaultImageGenerationTimeout = 180 * time.Second + maxImageGenerationResults = 4 + maxImageGenerationSSEBytes = 64 * 1024 * 1024 + maxImageGenerationEvents = 512 ) -// ImageGenerateTool generates images through OpenAI/Codex OAuth and returns +// ImageGenerateTool generates images through a provider adapter and returns // generated files through the MediaStore outbound media pipeline. type ImageGenerateTool struct { workspace string @@ -47,6 +48,14 @@ type imageGenerationProvider interface { GenerateImages(ctx context.Context, req imageGenerationRequest) ([]generatedImage, error) } +type imageGenerationProviderFactory func() imageGenerationProvider + +var imageGenerationProviderFactories = map[string]imageGenerationProviderFactory{ + defaultImageGenerationProvider: func() imageGenerationProvider { + return newOpenAICodexImageGenerationProvider() + }, +} + func WithImageGenerateBaseURL(baseURL string) ImageGenerateToolOption { return func(t *ImageGenerateTool) { if provider, ok := t.provider.(*openAICodexImageGenerationProvider); ok { @@ -85,13 +94,18 @@ func NewImageGenerateTool( store media.MediaStore, options ...ImageGenerateToolOption, ) *ImageGenerateTool { - if strings.TrimSpace(model) == "" { - model = defaultImageGenerationModel + spec := parseImageGenerationModel(model) + factory := imageGenerationProviderFactories[spec.Provider] + if factory == nil { + factory = imageGenerationProviderFactories[defaultImageGenerationProvider] + } + provider := factory() + if spec.Model == "" && provider != nil { + spec.Model = provider.DefaultModel() } - provider := newOpenAICodexImageGenerationProvider() tool := &ImageGenerateTool{ workspace: workspace, - model: stripProviderPrefix(model), + model: spec.Model, provider: provider, mediaStore: store, } @@ -110,7 +124,7 @@ func (t *ImageGenerateTool) Name() string { return "image_generate" } func (t *ImageGenerateTool) Description() string { return `Generate an image from a prompt and send it to the current chat. -Uses OpenAI/Codex OAuth with gpt-image-2 by default. Use this when the user asks to create an image, infographic, diagram, poster, visual summary, or other generated raster artwork.` +Use this when the user asks to create an image, infographic, diagram, poster, visual summary, or other generated raster artwork. The active image backend is selected from the configured image model provider prefix.` } func (t *ImageGenerateTool) Parameters() map[string]any { @@ -493,15 +507,35 @@ func readImageCount(raw any) int { return count } -func stripProviderPrefix(model string) string { +type imageGenerationModelSpec struct { + Provider string + Model string +} + +func parseImageGenerationModel(model string) imageGenerationModelSpec { model = strings.TrimSpace(model) - if strings.HasPrefix(model, "openai/") { - return strings.TrimPrefix(model, "openai/") + if model == "" { + return imageGenerationModelSpec{ + Provider: defaultImageGenerationProvider, + Model: defaultImageGenerationModel, + } } - if strings.HasPrefix(model, "openai-codex/") { - return strings.TrimPrefix(model, "openai-codex/") + 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, } - return model } func imageMimeAndExtension(outputFormat string) (string, string) { diff --git a/pkg/tools/image_generate_test.go b/pkg/tools/image_generate_test.go index 9d7978239..52d8ac425 100644 --- a/pkg/tools/image_generate_test.go +++ b/pkg/tools/image_generate_test.go @@ -1,6 +1,7 @@ package tools import ( + "context" "encoding/base64" "encoding/json" "net/http" @@ -11,6 +12,28 @@ import ( "github.com/sipeed/picoclaw/pkg/media" ) +type fakeImageGenerationProvider struct { + id string + defaultModel string + request imageGenerationRequest +} + +func (p *fakeImageGenerationProvider) ID() string { return p.id } + +func (p *fakeImageGenerationProvider) DefaultModel() string { return p.defaultModel } + +func (p *fakeImageGenerationProvider) GenerateImages( + _ context.Context, + req imageGenerationRequest, +) ([]generatedImage, error) { + p.request = req + return []generatedImage{{ + Data: []byte("fake-image"), + MimeType: "image/png", + Ext: "png", + }}, nil +} + func TestImageGenerateToolCodexOAuthRequestAndMediaResult(t *testing.T) { var captured map[string]any var gotAuth string @@ -90,6 +113,79 @@ func TestImageGenerateToolCodexOAuthRequestAndMediaResult(t *testing.T) { } } +func TestImageGenerateToolCanUseInjectedProvider(t *testing.T) { + store := media.NewFileMediaStore() + provider := &fakeImageGenerationProvider{ + id: "test-provider", + defaultModel: "test-default-image-model", + } + tool := NewImageGenerateTool( + t.TempDir(), + "test-provider/custom-image-model", + store, + WithImageGenerationProvider(provider), + ) + + result := tool.Execute( + WithToolContext(t.Context(), "telegram", "chat-1"), + map[string]any{"prompt": "make a tiny icon"}, + ) + if result.IsError { + t.Fatalf("Execute returned error: %s", result.ContentForLLM()) + } + if provider.request.Model != "custom-image-model" { + t.Fatalf("model = %q, want custom-image-model", provider.request.Model) + } + if len(result.Media) != 1 { + 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"