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,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) {

View file

@ -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"