make image generation provider pluggable

This commit is contained in:
Anton Bogdanovich 2026-05-05 13:31:53 -07:00
parent 31a771b1f7
commit 8844289268
2 changed files with 149 additions and 19 deletions

View file

@ -21,6 +21,7 @@ import (
) )
const ( const (
defaultImageGenerationProvider = "openai-codex"
defaultImageGenerationModel = "gpt-image-2" defaultImageGenerationModel = "gpt-image-2"
defaultImageGenerationBaseURL = "https://chatgpt.com/backend-api/codex" defaultImageGenerationBaseURL = "https://chatgpt.com/backend-api/codex"
defaultImageGenerationSize = "1024x1024" defaultImageGenerationSize = "1024x1024"
@ -30,7 +31,7 @@ const (
maxImageGenerationEvents = 512 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. // generated files through the MediaStore outbound media pipeline.
type ImageGenerateTool struct { type ImageGenerateTool struct {
workspace string workspace string
@ -47,6 +48,14 @@ type imageGenerationProvider interface {
GenerateImages(ctx context.Context, req imageGenerationRequest) ([]generatedImage, error) 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 { func WithImageGenerateBaseURL(baseURL string) ImageGenerateToolOption {
return func(t *ImageGenerateTool) { return func(t *ImageGenerateTool) {
if provider, ok := t.provider.(*openAICodexImageGenerationProvider); ok { if provider, ok := t.provider.(*openAICodexImageGenerationProvider); ok {
@ -85,13 +94,18 @@ func NewImageGenerateTool(
store media.MediaStore, store media.MediaStore,
options ...ImageGenerateToolOption, options ...ImageGenerateToolOption,
) *ImageGenerateTool { ) *ImageGenerateTool {
if strings.TrimSpace(model) == "" { spec := parseImageGenerationModel(model)
model = defaultImageGenerationModel 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{ tool := &ImageGenerateTool{
workspace: workspace, workspace: workspace,
model: stripProviderPrefix(model), model: spec.Model,
provider: provider, provider: provider,
mediaStore: store, mediaStore: store,
} }
@ -110,7 +124,7 @@ func (t *ImageGenerateTool) Name() string { return "image_generate" }
func (t *ImageGenerateTool) Description() string { func (t *ImageGenerateTool) Description() string {
return `Generate an image from a prompt and send it to the current chat. 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 { func (t *ImageGenerateTool) Parameters() map[string]any {
@ -493,15 +507,35 @@ func readImageCount(raw any) int {
return count return count
} }
func stripProviderPrefix(model string) string { type imageGenerationModelSpec struct {
Provider string
Model string
}
func parseImageGenerationModel(model string) imageGenerationModelSpec {
model = strings.TrimSpace(model) model = strings.TrimSpace(model)
if strings.HasPrefix(model, "openai/") { if model == "" {
return strings.TrimPrefix(model, "openai/") return imageGenerationModelSpec{
Provider: defaultImageGenerationProvider,
Model: defaultImageGenerationModel,
} }
if strings.HasPrefix(model, "openai-codex/") {
return strings.TrimPrefix(model, "openai-codex/")
} }
return model 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,
}
} }
func imageMimeAndExtension(outputFormat string) (string, string) { func imageMimeAndExtension(outputFormat string) (string, string) {

View file

@ -1,6 +1,7 @@
package tools package tools
import ( import (
"context"
"encoding/base64" "encoding/base64"
"encoding/json" "encoding/json"
"net/http" "net/http"
@ -11,6 +12,28 @@ import (
"github.com/sipeed/picoclaw/pkg/media" "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) { func TestImageGenerateToolCodexOAuthRequestAndMediaResult(t *testing.T) {
var captured map[string]any var captured map[string]any
var gotAuth string 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) { func TestParseCodexImageSSECompletedResponseFallback(t *testing.T) {
payload := base64.StdEncoding.EncodeToString([]byte("fake-png")) payload := base64.StdEncoding.EncodeToString([]byte("fake-png"))
body := `data: {"type":"response.completed","response":{"output":[{"type":"image_generation_call","result":"` + payload + `"}]}}` + "\n\n" body := `data: {"type":"response.completed","response":{"output":[{"type":"image_generation_call","result":"` + payload + `"}]}}` + "\n\n"