make image generation provider pluggable
This commit is contained in:
parent
31a771b1f7
commit
8844289268
2 changed files with 149 additions and 19 deletions
|
|
@ -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) {
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue