mirror of
https://github.com/tiennm99/goclaw.git
synced 2026-10-03 07:12:50 +00:00
feat(tools): image reference processing and native provider support (#1251)
* feat(tools): implement image reference processing and native provider support - Support OpenAI image edits via both Multipart form-data and JSON payloads - Automatically append reference image descriptions to prompt under [Reference Image Roles] - Support downloading image URLs for Gemini native image generation - Deduplicate reference images to optimize API request size - Add unit tests for Codex, DashScope, MiniMax, BytePlus and local/remote path resolution * fix(tools): SSRF-guard reference-image URL downloads in create_image downloadImageBytes fetched caller-supplied ref_images[].url with a plain http.Client and unbounded io.ReadAll — no SSRF validation, redirect policy, or size cap, letting the gateway dial loopback/private/metadata hosts or read arbitrarily large responses. - Validate the URL via security.Validate and pin the resolved IP, then download through security.NewSafeClient (pinned dial, no redirects). - Cap the response with a bounded read (refImageMaxBytes, 20 MB). - Reject non-HTTP(S) reference URLs up front (file://, data:, gopher://) so provider-forwarded URLs stay HTTP(S)-only; document the trust boundary between gateway-side fetch and provider-forwarded URLs. - Add regression tests: blocked loopback/private/metadata, unfollowed redirect, oversized response, and non-http(s) scheme rejection.
This commit is contained in:
1 parent
389640ae51
commit
a5a853f461
9 files changed
+1465
-34
No files matched your search
@@ -57,6 +57,39 @@ func (p *CodexProvider) GenerateImage(ctx context.Context, req NativeImageReques
|
||||
// so stream is always true. Final assembly happens in parseNativeImageSSE which scans the
|
||||
// event stream for response.output_item.done (image item) or response.completed output walk.
|
||||
func (p *CodexProvider) buildNativeImageRequestBody(model string, req NativeImageRequest) map[string]any {
|
||||
tool := map[string]any{
|
||||
"type": "image_generation",
|
||||
"action": "generate",
|
||||
"model": req.ImageModel,
|
||||
"output_format": req.OutputFormat,
|
||||
"size": SizeFromAspect(req.AspectRatio),
|
||||
}
|
||||
|
||||
contentParts := []map[string]any{}
|
||||
|
||||
for _, img := range req.RefImages {
|
||||
if img.Base64 != "" {
|
||||
refMime := img.MimeType
|
||||
if refMime == "" {
|
||||
refMime = "image/png"
|
||||
}
|
||||
contentParts = append(contentParts, map[string]any{
|
||||
"type": "input_image",
|
||||
"image_url": fmt.Sprintf("data:%s;base64,%s", refMime, img.Base64),
|
||||
})
|
||||
} else if img.URL != "" {
|
||||
contentParts = append(contentParts, map[string]any{
|
||||
"type": "input_image",
|
||||
"image_url": img.URL,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
contentParts = append(contentParts, map[string]any{
|
||||
"type": "input_text",
|
||||
"text": req.Prompt,
|
||||
})
|
||||
|
||||
return map[string]any{
|
||||
"model": model,
|
||||
"stream": true,
|
||||
@@ -64,21 +97,11 @@ func (p *CodexProvider) buildNativeImageRequestBody(model string, req NativeImag
|
||||
"instructions": "Generate an image matching the user's description using the image_generation tool. Return only the image; do not describe it in text.",
|
||||
"input": []any{
|
||||
map[string]any{
|
||||
"role": "user",
|
||||
"content": []map[string]any{
|
||||
{"type": "input_text", "text": req.Prompt},
|
||||
},
|
||||
},
|
||||
},
|
||||
"tools": []map[string]any{
|
||||
{
|
||||
"type": "image_generation",
|
||||
"action": "generate",
|
||||
"model": req.ImageModel,
|
||||
"output_format": req.OutputFormat,
|
||||
"size": SizeFromAspect(req.AspectRatio),
|
||||
"role": "user",
|
||||
"content": contentParts,
|
||||
},
|
||||
},
|
||||
"tools": []map[string]any{tool},
|
||||
"tool_choice": map[string]any{
|
||||
"type": "image_generation",
|
||||
},
|
||||
|
||||
@@ -344,3 +344,140 @@ func TestSizeFromAspect(t *testing.T) {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestCodexGenerateImage_WithReferenceImage verifies that passing a single RefImage
|
||||
// embeds the image in the input content and does not populate input_reference in tools[0].
|
||||
func TestCodexGenerateImage_WithReferenceImage(t *testing.T) {
|
||||
var captured []byte
|
||||
server := mockImageServer(t, &captured)
|
||||
defer server.Close()
|
||||
|
||||
p := NewCodexProvider("codex-test", &staticTokenSource{token: "tok"}, server.URL, "gpt-image-2")
|
||||
p.retryConfig.Attempts = 1
|
||||
|
||||
req := NativeImageRequest{
|
||||
Model: "gpt-image-2",
|
||||
Prompt: "A red circle",
|
||||
RefImages: []RefImage{
|
||||
{
|
||||
URL: "https://example.com/ref.png",
|
||||
},
|
||||
},
|
||||
AspectRatio: "1:1",
|
||||
OutputFormat: "png",
|
||||
}
|
||||
|
||||
_, err := p.GenerateImage(context.Background(), req)
|
||||
if err != nil {
|
||||
t.Fatalf("GenerateImage returned error: %v", err)
|
||||
}
|
||||
|
||||
var body map[string]any
|
||||
if err := json.Unmarshal(captured, &body); err != nil {
|
||||
t.Fatalf("unmarshal captured body: %v", err)
|
||||
}
|
||||
|
||||
// 1. Verify input contains input_image and input_text in user content
|
||||
inputs, ok := body["input"].([]any)
|
||||
if !ok || len(inputs) != 1 {
|
||||
t.Fatalf("body[input]: expected []any length 1, got %T len %d", body["input"], len(inputs))
|
||||
}
|
||||
userMsg, ok := inputs[0].(map[string]any)
|
||||
if !ok {
|
||||
t.Fatalf("inputs[0] is not a map")
|
||||
}
|
||||
contents, ok := userMsg["content"].([]any)
|
||||
if !ok || len(contents) != 2 {
|
||||
t.Fatalf("content: expected []any length 2, got %T len %d", userMsg["content"], len(contents))
|
||||
}
|
||||
|
||||
imgPart, ok := contents[0].(map[string]any)
|
||||
if !ok {
|
||||
t.Fatalf("contents[0] is not a map")
|
||||
}
|
||||
if imgPart["type"] != "input_image" || imgPart["image_url"] != "https://example.com/ref.png" {
|
||||
t.Errorf("expected input_image with url, got: %v", imgPart)
|
||||
}
|
||||
|
||||
textPart, ok := contents[1].(map[string]any)
|
||||
if !ok {
|
||||
t.Fatalf("contents[1] is not a map")
|
||||
}
|
||||
if textPart["type"] != "input_text" || textPart["text"] != "A red circle" {
|
||||
t.Errorf("expected input_text with prompt, got: %v", textPart)
|
||||
}
|
||||
|
||||
// 2. Verify tools[0] does not contain input_reference
|
||||
tools, ok := body["tools"].([]any)
|
||||
if !ok || len(tools) != 1 {
|
||||
t.Fatalf("tools shape invalid")
|
||||
}
|
||||
tool, ok := tools[0].(map[string]any)
|
||||
if !ok {
|
||||
t.Fatalf("tools[0] not map")
|
||||
}
|
||||
if _, has := tool["input_reference"]; has {
|
||||
t.Error("image_generation tool must not contain 'input_reference' field")
|
||||
}
|
||||
}
|
||||
|
||||
// TestCodexGenerateImage_WithMultipleReferenceImages verifies that passing multiple RefImages
|
||||
// embeds all of them in the input content in the correct order.
|
||||
func TestCodexGenerateImage_WithMultipleReferenceImages(t *testing.T) {
|
||||
var captured []byte
|
||||
server := mockImageServer(t, &captured)
|
||||
defer server.Close()
|
||||
|
||||
p := NewCodexProvider("codex-test", &staticTokenSource{token: "tok"}, server.URL, "gpt-image-2")
|
||||
p.retryConfig.Attempts = 1
|
||||
|
||||
req := NativeImageRequest{
|
||||
Model: "gpt-image-2",
|
||||
Prompt: "A red circle",
|
||||
RefImages: []RefImage{
|
||||
{
|
||||
URL: "https://example.com/ref1.png",
|
||||
},
|
||||
{
|
||||
Base64: "b64data",
|
||||
MimeType: "image/jpeg",
|
||||
},
|
||||
},
|
||||
AspectRatio: "1:1",
|
||||
OutputFormat: "png",
|
||||
}
|
||||
|
||||
_, err := p.GenerateImage(context.Background(), req)
|
||||
if err != nil {
|
||||
t.Fatalf("GenerateImage returned error: %v", err)
|
||||
}
|
||||
|
||||
var body map[string]any
|
||||
if err := json.Unmarshal(captured, &body); err != nil {
|
||||
t.Fatalf("unmarshal captured body: %v", err)
|
||||
}
|
||||
|
||||
inputs, _ := body["input"].([]any)
|
||||
userMsg, _ := inputs[0].(map[string]any)
|
||||
contents, _ := userMsg["content"].([]any)
|
||||
|
||||
// Expected: 2 images + 1 text prompt = 3 content parts
|
||||
if len(contents) != 3 {
|
||||
t.Fatalf("content: expected []any length 3, got len %d", len(contents))
|
||||
}
|
||||
|
||||
imgPart1, _ := contents[0].(map[string]any)
|
||||
if imgPart1["type"] != "input_image" || imgPart1["image_url"] != "https://example.com/ref1.png" {
|
||||
t.Errorf("expected first input_image with url, got: %v", imgPart1)
|
||||
}
|
||||
|
||||
imgPart2, _ := contents[1].(map[string]any)
|
||||
if imgPart2["type"] != "input_image" || imgPart2["image_url"] != "data:image/jpeg;base64,b64data" {
|
||||
t.Errorf("expected second input_image with base64 data url, got: %v", imgPart2)
|
||||
}
|
||||
|
||||
textPart, _ := contents[2].(map[string]any)
|
||||
if textPart["type"] != "input_text" || textPart["text"] != "A red circle" {
|
||||
t.Errorf("expected input_text with prompt, got: %v", textPart)
|
||||
}
|
||||
}
|
||||
@@ -61,6 +61,18 @@ type NativeImageRequest struct {
|
||||
|
||||
// OutputFormat is the desired image format: "png" (default), "jpg", "webp".
|
||||
OutputFormat string
|
||||
|
||||
// RefImages contains the list of reference images.
|
||||
RefImages []RefImage
|
||||
}
|
||||
|
||||
// RefImage represents a single reference image for image-to-image or styling tasks.
|
||||
type RefImage struct {
|
||||
Data []byte
|
||||
Base64 string
|
||||
MimeType string
|
||||
URL string
|
||||
Strength float64
|
||||
}
|
||||
|
||||
// NativeImageResult holds the result of a native image generation call.
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"mime/multipart"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
@@ -16,8 +17,17 @@ import (
|
||||
|
||||
"github.com/nextlevelbuilder/goclaw/internal/bus"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/providers"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/security"
|
||||
)
|
||||
|
||||
// Reference-image downloads are gateway-side fetches of caller-supplied URLs, so
|
||||
// they must go through the SSRF guard with a bounded read.
|
||||
const refImageDownloadTimeout = 30 * time.Second
|
||||
|
||||
// refImageMaxBytes caps a single reference-image download. Declared as a var so
|
||||
// tests can shrink it to exercise the overflow path cheaply.
|
||||
var refImageMaxBytes int64 = 20 * 1024 * 1024 // 20 MB
|
||||
|
||||
// credentialProvider is a narrow interface for providers that expose API credentials.
|
||||
type credentialProvider interface {
|
||||
APIKey() string
|
||||
@@ -45,6 +55,145 @@ type CreateImageTool struct {
|
||||
|
||||
func (t *CreateImageTool) SetVaultInterceptor(v *VaultInterceptor) { t.vaultIntc = v }
|
||||
|
||||
type referenceImage struct {
|
||||
Data []byte
|
||||
Base64 string
|
||||
URL string
|
||||
MimeType string
|
||||
Strength float64
|
||||
Description string
|
||||
}
|
||||
|
||||
func (t *CreateImageTool) resolveReferenceImages(ctx context.Context, args map[string]any) ([]*referenceImage, error) {
|
||||
var results []*referenceImage
|
||||
|
||||
seenPaths := make(map[string]bool)
|
||||
seenURLs := make(map[string]bool)
|
||||
seenIDs := make(map[string]bool)
|
||||
|
||||
resolveSingle := func(path, url, id string, strength float64, description string) (*referenceImage, error) {
|
||||
if path != "" {
|
||||
if seenPaths[path] {
|
||||
return nil, nil
|
||||
}
|
||||
seenPaths[path] = true
|
||||
|
||||
ext := strings.ToLower(filepath.Ext(path))
|
||||
mimeTypes := map[string]string{
|
||||
".jpg": "image/jpeg", ".jpeg": "image/jpeg",
|
||||
".png": "image/png", ".gif": "image/gif",
|
||||
".webp": "image/webp", ".bmp": "image/bmp",
|
||||
}
|
||||
mime, ok := mimeTypes[ext]
|
||||
if !ok {
|
||||
mime = "image/png"
|
||||
}
|
||||
workspace := ToolWorkspaceFromCtx(ctx)
|
||||
resolved, err := resolvePathWithAllowed(path, workspace, effectiveRestrict(ctx, true), allowedWithTeamWorkspace(ctx, nil))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid reference image path: %w", err)
|
||||
}
|
||||
if err := checkDeniedPath(resolved, workspace, nil); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
data, err := os.ReadFile(resolved)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read reference image file: %w", err)
|
||||
}
|
||||
return &referenceImage{
|
||||
Data: data,
|
||||
Base64: base64.StdEncoding.EncodeToString(data),
|
||||
MimeType: mime,
|
||||
Strength: strength,
|
||||
Description: description,
|
||||
}, nil
|
||||
}
|
||||
|
||||
if url != "" {
|
||||
// Trust boundary: a reference URL is either forwarded to the image
|
||||
// provider (provider fetches it) or fetched gateway-side via
|
||||
// downloadImageBytes (SSRF-guarded there). Either way it must be a
|
||||
// plain HTTP(S) URL — reject file://, gopher://, data:, etc. up front.
|
||||
if !isHTTPURL(url) {
|
||||
return nil, fmt.Errorf("reference image url must be http(s): %q", url)
|
||||
}
|
||||
if seenURLs[url] {
|
||||
return nil, nil
|
||||
}
|
||||
seenURLs[url] = true
|
||||
|
||||
return &referenceImage{
|
||||
URL: url,
|
||||
Strength: strength,
|
||||
Description: description,
|
||||
}, nil
|
||||
}
|
||||
|
||||
if id != "" {
|
||||
if seenIDs[id] {
|
||||
return nil, nil
|
||||
}
|
||||
seenIDs[id] = true
|
||||
|
||||
images := MediaImagesFromCtx(ctx)
|
||||
if len(images) == 0 {
|
||||
return nil, fmt.Errorf("no images available in conversation context")
|
||||
}
|
||||
var img providers.ImageContent
|
||||
if id == "latest" {
|
||||
img = images[len(images)-1]
|
||||
} else {
|
||||
var idx int
|
||||
if _, err := fmt.Sscanf(id, "%d", &idx); err == nil && idx >= 0 && idx < len(images) {
|
||||
img = images[idx]
|
||||
} else {
|
||||
img = images[len(images)-1]
|
||||
}
|
||||
}
|
||||
dataBytes, _ := base64.StdEncoding.DecodeString(img.Data)
|
||||
return &referenceImage{
|
||||
Data: dataBytes,
|
||||
Base64: img.Data,
|
||||
MimeType: img.MimeType,
|
||||
URL: img.URL,
|
||||
Strength: strength,
|
||||
Description: description,
|
||||
}, nil
|
||||
}
|
||||
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
// 1. Resolve complex ref_images array
|
||||
if refImagesRaw, ok := args["ref_images"]; ok {
|
||||
if refImagesList, ok := refImagesRaw.([]any); ok {
|
||||
for _, itemRaw := range refImagesList {
|
||||
if item, ok := itemRaw.(map[string]any); ok {
|
||||
path, _ := item["path"].(string)
|
||||
url, _ := item["url"].(string)
|
||||
id, _ := item["id"].(string)
|
||||
description, _ := item["description"].(string)
|
||||
strength := 0.6
|
||||
if strRaw, has := item["strength"]; has {
|
||||
if s, ok := strRaw.(float64); ok {
|
||||
strength = s
|
||||
}
|
||||
}
|
||||
refImg, err := resolveSingle(path, url, id, strength, description)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if refImg != nil {
|
||||
results = append(results, refImg)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return results, nil
|
||||
}
|
||||
|
||||
func NewCreateImageTool(registry *providers.Registry) *CreateImageTool {
|
||||
return &CreateImageTool{registry: registry}
|
||||
}
|
||||
@@ -71,6 +220,20 @@ func (t *CreateImageTool) Parameters() map[string]any {
|
||||
"type": "string",
|
||||
"description": "Short descriptive filename (no extension). Example: 'sunset-beach', 'company-logo'.",
|
||||
},
|
||||
"ref_images": map[string]any{
|
||||
"type": "array",
|
||||
"description": "Optional array of reference images with custom properties.",
|
||||
"items": map[string]any{
|
||||
"type": "object",
|
||||
"properties": map[string]any{
|
||||
"path": map[string]any{"type": "string", "description": "Workspace file path to a reference image."},
|
||||
"url": map[string]any{"type": "string", "description": "HTTP/HTTPS URL of a reference image."},
|
||||
"id": map[string]any{"type": "string", "description": "Media ID of a reference image from the chat."},
|
||||
"strength": map[string]any{"type": "number", "description": "Reference strength (0.0 to 1.0) specific to this image."},
|
||||
"description": map[string]any{"type": "string", "description": "Description of the role or content of this reference image (e.g. 'Lâm', 'Quân')."},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
"required": []string{"prompt"},
|
||||
}
|
||||
@@ -87,16 +250,24 @@ func (t *CreateImageTool) Execute(ctx context.Context, args map[string]any) *Res
|
||||
}
|
||||
filenameHint, _ := args["filename_hint"].(string)
|
||||
|
||||
refImgs, err := t.resolveReferenceImages(ctx, args)
|
||||
if err != nil {
|
||||
return ErrorResult(fmt.Sprintf("Failed to resolve reference images: %v", err))
|
||||
}
|
||||
|
||||
chain := ResolveMediaProviderChain(ctx, "create_image", "", "",
|
||||
imageGenProviderPriority, imageGenModelDefaults, t.registry)
|
||||
|
||||
// Inject prompt and aspect_ratio into each chain entry's params
|
||||
// Inject prompt, aspect_ratio, and ref_images into each chain entry's params
|
||||
for i := range chain {
|
||||
if chain[i].Params == nil {
|
||||
chain[i].Params = make(map[string]any)
|
||||
}
|
||||
chain[i].Params["prompt"] = prompt
|
||||
chain[i].Params["aspect_ratio"] = aspectRatio
|
||||
if len(refImgs) > 0 {
|
||||
chain[i].Params["ref_images"] = refImgs
|
||||
}
|
||||
}
|
||||
|
||||
chainResult, err := ExecuteWithChain(ctx, chain, t.registry, t.callProvider)
|
||||
@@ -165,6 +336,11 @@ func embedPromptIntoPNG(data []byte, prompt string) []byte {
|
||||
// the native path is used and cp may be nil. The credentialProvider path is only reached
|
||||
// for API-key-backed providers.
|
||||
func (t *CreateImageTool) callProvider(ctx context.Context, cp credentialProvider, providerName, model string, params map[string]any) ([]byte, *providers.Usage, error) {
|
||||
var refImgs []*referenceImage
|
||||
if rawImgs, ok := params["ref_images"]; ok {
|
||||
refImgs, _ = rawImgs.([]*referenceImage)
|
||||
}
|
||||
|
||||
// Native path: provider implements the image_generation tool natively (e.g. Codex/OAuth).
|
||||
// The raw provider object is injected into params["_native_provider"] by ExecuteWithChain.
|
||||
// Must check before the cp==nil guard — these providers intentionally have no APIKey/APIBase.
|
||||
@@ -173,13 +349,28 @@ func (t *CreateImageTool) callProvider(ctx context.Context, cp credentialProvide
|
||||
prompt := GetParamString(params, "prompt", "")
|
||||
aspectRatio := GetParamString(params, "aspect_ratio", "1:1")
|
||||
imageModel := GetParamString(params, "image_model", "")
|
||||
result, err := np.GenerateImage(ctx, providers.NativeImageRequest{
|
||||
|
||||
req := providers.NativeImageRequest{
|
||||
Model: model,
|
||||
ImageModel: imageModel,
|
||||
Prompt: prompt,
|
||||
AspectRatio: aspectRatio,
|
||||
OutputFormat: "png",
|
||||
})
|
||||
}
|
||||
if len(refImgs) > 0 {
|
||||
req.RefImages = make([]providers.RefImage, len(refImgs))
|
||||
for idx, r := range refImgs {
|
||||
req.RefImages[idx] = providers.RefImage{
|
||||
Data: r.Data,
|
||||
Base64: r.Base64,
|
||||
MimeType: r.MimeType,
|
||||
URL: r.URL,
|
||||
Strength: r.Strength,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
result, err := np.GenerateImage(ctx, req)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("native image generation: %w", err)
|
||||
}
|
||||
@@ -196,7 +387,30 @@ func (t *CreateImageTool) callProvider(ctx context.Context, cp credentialProvide
|
||||
slog.Info("create_image: calling image generation API",
|
||||
"provider", providerName, "model", model, "aspect_ratio", aspectRatio)
|
||||
|
||||
switch GetParamString(params, "_provider_type", providerTypeFromName(providerName)) {
|
||||
ptype := GetParamString(params, "_provider_type", providerTypeFromName(providerName))
|
||||
|
||||
// OpenAI image-to-image (Edits)
|
||||
if len(refImgs) > 0 && (ptype == "openai" || providerName == "openai" || ptype == "openai_compat") {
|
||||
isEditModel := model == "gpt-image-2" ||
|
||||
model == "gpt-image-1.5" ||
|
||||
model == "gpt-image-1" ||
|
||||
model == "gpt-image-1-mini" ||
|
||||
model == "chatgpt-image-latest" ||
|
||||
model == "dall-e-2" ||
|
||||
ptype == "openai_compat"
|
||||
if isEditModel {
|
||||
if model == "dall-e-2" {
|
||||
if len(refImgs) > 1 {
|
||||
slog.Warn("openai dall-e-2 only supports 1 reference image, using the first one", "count", len(refImgs))
|
||||
}
|
||||
return t.callOpenAIImageEditMultipart(ctx, cp.APIKey(), cp.APIBase(), model, prompt, refImgs[:1])
|
||||
}
|
||||
return t.callOpenAIImageEditJSON(ctx, cp.APIKey(), cp.APIBase(), model, prompt, refImgs)
|
||||
}
|
||||
slog.Warn("create_image: model does not support reference images, ignoring reference", "model", model, "provider", providerName)
|
||||
}
|
||||
|
||||
switch ptype {
|
||||
case "gemini":
|
||||
return t.callGeminiNativeImageGen(ctx, cp.APIKey(), cp.APIBase(), model, prompt, params)
|
||||
case "openrouter":
|
||||
@@ -215,11 +429,54 @@ func (t *CreateImageTool) callProvider(ctx context.Context, cp credentialProvide
|
||||
// callImageGenAPI calls the OpenAI-compatible chat completions endpoint with image modalities.
|
||||
// Works with OpenRouter (modalities: ["image","text"]).
|
||||
func (t *CreateImageTool) callImageGenAPI(ctx context.Context, apiKey, apiBase, model, prompt, aspectRatio string, params map[string]any) ([]byte, *providers.Usage, error) {
|
||||
body := map[string]any{
|
||||
"model": model,
|
||||
"messages": []map[string]any{
|
||||
var messages []map[string]any
|
||||
if rawImgs, ok := params["ref_images"]; ok {
|
||||
if refImgs, ok := rawImgs.([]*referenceImage); ok && len(refImgs) > 0 {
|
||||
contentParts := []map[string]any{
|
||||
{"type": "text", "text": prompt},
|
||||
}
|
||||
for _, refImg := range refImgs {
|
||||
// Trust boundary: this JSON path forwards the reference URL to the
|
||||
// image provider downstream (the provider fetches it), so the gateway
|
||||
// does NOT dial it here — only HTTP(S) URLs reach this point. The
|
||||
// gateway-side fetch path (OpenAI multipart) goes through
|
||||
// downloadImageBytes, which is SSRF-guarded. Keep these distinct: if a
|
||||
// provider URL ever becomes a gateway-side fetch, route it through
|
||||
// downloadImageBytes.
|
||||
var refURL string
|
||||
if refImg.URL != "" {
|
||||
refURL = refImg.URL
|
||||
} else {
|
||||
refMime := refImg.MimeType
|
||||
if refMime == "" {
|
||||
refMime = "image/png"
|
||||
}
|
||||
refURL = fmt.Sprintf("data:%s;base64,%s", refMime, refImg.Base64)
|
||||
}
|
||||
contentParts = append(contentParts, map[string]any{
|
||||
"type": "image_url",
|
||||
"image_url": map[string]any{
|
||||
"url": refURL,
|
||||
},
|
||||
})
|
||||
}
|
||||
messages = []map[string]any{
|
||||
{
|
||||
"role": "user",
|
||||
"content": contentParts,
|
||||
},
|
||||
}
|
||||
}
|
||||
}
|
||||
if len(messages) == 0 {
|
||||
messages = []map[string]any{
|
||||
{"role": "user", "content": prompt},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
body := map[string]any{
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
"modalities": []string{"image", "text"},
|
||||
}
|
||||
if aspectRatio != "" && aspectRatio != "1:1" {
|
||||
@@ -316,6 +573,286 @@ func (t *CreateImageTool) callStandardImageGenAPI(ctx context.Context, apiKey, a
|
||||
return imageBytes, nil, nil
|
||||
}
|
||||
|
||||
// isHTTPURL reports whether s is a plain http(s) URL. Reference URLs must be
|
||||
// HTTP(S) only — this blocks file://, data:, gopher://, etc. before a URL is
|
||||
// either forwarded to a provider or fetched gateway-side.
|
||||
func isHTTPURL(s string) bool {
|
||||
l := strings.ToLower(s)
|
||||
return strings.HasPrefix(l, "http://") || strings.HasPrefix(l, "https://")
|
||||
}
|
||||
|
||||
// downloadImageBytes downloads a caller-supplied reference image URL server-side
|
||||
// and returns its raw bytes and content type. The URL is attacker-controlled
|
||||
// (agent/user-provided ref_images[].url), so it is validated against the SSRF
|
||||
// guard and the resolved IP is pinned for the dial; the shared SafeClient also
|
||||
// refuses redirects. The response body is read with a hard size cap.
|
||||
func (t *CreateImageTool) downloadImageBytes(ctx context.Context, rawURL string) ([]byte, string, error) {
|
||||
// SSRF guard: rejects loopback/private/link-local (incl. cloud metadata
|
||||
// 169.254.169.254) and returns the resolved IP to pin for the dial.
|
||||
_, pinnedIP, err := security.Validate(rawURL)
|
||||
if err != nil {
|
||||
return nil, "", fmt.Errorf("invalid reference image URL: %w", err)
|
||||
}
|
||||
reqCtx := security.WithPinnedIP(ctx, pinnedIP)
|
||||
req, err := http.NewRequestWithContext(reqCtx, "GET", rawURL, nil)
|
||||
if err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
// SafeClient dials only the pinned IP and never follows redirects (a 3xx is
|
||||
// returned as-is and rejected by the status check below).
|
||||
client := security.NewSafeClient(refImageDownloadTimeout)
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, "", fmt.Errorf("HTTP error %d", resp.StatusCode)
|
||||
}
|
||||
// Bounded read: cap the download to avoid memory exhaustion from a hostile
|
||||
// or oversized response. Read one extra byte to detect overflow.
|
||||
data, err := io.ReadAll(io.LimitReader(resp.Body, refImageMaxBytes+1))
|
||||
if err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
if int64(len(data)) > refImageMaxBytes {
|
||||
return nil, "", fmt.Errorf("reference image exceeds maximum size of %d bytes", refImageMaxBytes)
|
||||
}
|
||||
return data, resp.Header.Get("Content-Type"), nil
|
||||
}
|
||||
|
||||
// callOpenAIImageEditMultipart calls the OpenAI /v1/images/edits API using multipart/form-data.
|
||||
func (t *CreateImageTool) callOpenAIImageEditMultipart(ctx context.Context, apiKey, apiBase, model, prompt string, refImgs []*referenceImage) ([]byte, *providers.Usage, error) {
|
||||
body := &bytes.Buffer{}
|
||||
writer := multipart.NewWriter(body)
|
||||
hasImages := false
|
||||
|
||||
for idx, refImg := range refImgs {
|
||||
var imageData []byte
|
||||
var err error
|
||||
|
||||
if len(refImg.Data) > 0 {
|
||||
imageData = refImg.Data
|
||||
} else if refImg.URL != "" {
|
||||
slog.Info("openai multipart: downloading reference image from URL", "url", refImg.URL)
|
||||
imageData, _, err = t.downloadImageBytes(ctx, refImg.URL)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("failed to download reference image %d: %w", idx, err)
|
||||
}
|
||||
}
|
||||
|
||||
if len(imageData) == 0 {
|
||||
continue
|
||||
}
|
||||
hasImages = true
|
||||
|
||||
fieldName := "image"
|
||||
if len(refImgs) > 1 {
|
||||
fieldName = "images"
|
||||
}
|
||||
|
||||
filename := fmt.Sprintf("image_%d.png", idx)
|
||||
if refImg.MimeType != "" {
|
||||
parts := strings.Split(refImg.MimeType, "/")
|
||||
if len(parts) == 2 {
|
||||
ext := parts[1]
|
||||
if ext == "jpeg" {
|
||||
ext = "jpg"
|
||||
}
|
||||
filename = fmt.Sprintf("image_%d.%s", idx, ext)
|
||||
}
|
||||
}
|
||||
|
||||
part, err := writer.CreateFormFile(fieldName, filename)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("create form file image %d: %w", idx, err)
|
||||
}
|
||||
if _, err := part.Write(imageData); err != nil {
|
||||
return nil, nil, fmt.Errorf("write image %d to form: %w", idx, err)
|
||||
}
|
||||
}
|
||||
|
||||
if !hasImages {
|
||||
return nil, nil, fmt.Errorf("no reference image data available")
|
||||
}
|
||||
|
||||
// Build reference image descriptions if available
|
||||
var descParts []string
|
||||
for idx, refImg := range refImgs {
|
||||
if refImg.Description != "" {
|
||||
descParts = append(descParts, fmt.Sprintf("- image_%d.png: %s", idx+1, refImg.Description))
|
||||
}
|
||||
}
|
||||
finalPrompt := prompt
|
||||
if len(descParts) > 0 {
|
||||
finalPrompt = fmt.Sprintf("%s\n\n[Reference Image Roles]\n%s", prompt, strings.Join(descParts, "\n"))
|
||||
slog.Info("openai multipart: appended image descriptions to prompt", "desc_count", len(descParts))
|
||||
}
|
||||
|
||||
// Add other fields
|
||||
if err := writer.WriteField("prompt", finalPrompt); err != nil {
|
||||
return nil, nil, fmt.Errorf("write field prompt: %w", err)
|
||||
}
|
||||
if err := writer.WriteField("model", model); err != nil {
|
||||
return nil, nil, fmt.Errorf("write field model: %w", err)
|
||||
}
|
||||
if err := writer.WriteField("response_format", "b64_json"); err != nil {
|
||||
return nil, nil, fmt.Errorf("write field response_format: %w", err)
|
||||
}
|
||||
|
||||
if err := writer.Close(); err != nil {
|
||||
return nil, nil, fmt.Errorf("close multipart writer: %w", err)
|
||||
}
|
||||
|
||||
url := strings.TrimRight(apiBase, "/") + "/images/edits"
|
||||
req, err := http.NewRequestWithContext(ctx, "POST", url, body)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("create request: %w", err)
|
||||
}
|
||||
req.Header.Set("Content-Type", writer.FormDataContentType())
|
||||
req.Header.Set("Authorization", "Bearer "+apiKey)
|
||||
|
||||
client := &http.Client{}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("http request: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
respBody, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("read response: %w", err)
|
||||
}
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, nil, fmt.Errorf("API error %d: %s", resp.StatusCode, truncateBytes(respBody, 500))
|
||||
}
|
||||
|
||||
var imgResp struct {
|
||||
Data []struct {
|
||||
B64JSON string `json:"b64_json"`
|
||||
} `json:"data"`
|
||||
}
|
||||
if err := json.Unmarshal(respBody, &imgResp); err != nil {
|
||||
return nil, nil, fmt.Errorf("parse response: %w", err)
|
||||
}
|
||||
if len(imgResp.Data) == 0 || imgResp.Data[0].B64JSON == "" {
|
||||
return nil, nil, fmt.Errorf("no image data in response")
|
||||
}
|
||||
|
||||
imageBytes, err := base64.StdEncoding.DecodeString(imgResp.Data[0].B64JSON)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("decode base64: %w", err)
|
||||
}
|
||||
|
||||
return imageBytes, nil, nil
|
||||
}
|
||||
|
||||
// callOpenAIImageEditJSON calls the OpenAI /v1/images/edits API using a JSON payload.
|
||||
func (t *CreateImageTool) callOpenAIImageEditJSON(ctx context.Context, apiKey, apiBase, model, prompt string, refImgs []*referenceImage) ([]byte, *providers.Usage, error) {
|
||||
type ImageRef struct {
|
||||
ImageURL string `json:"image_url,omitempty"`
|
||||
FileID string `json:"file_id,omitempty"`
|
||||
}
|
||||
|
||||
var images []ImageRef
|
||||
|
||||
for _, refImg := range refImgs {
|
||||
var refURL string
|
||||
if refImg.URL != "" {
|
||||
refURL = refImg.URL
|
||||
} else {
|
||||
// Convert local image data to Base64 Data URL
|
||||
var imageData []byte
|
||||
if len(refImg.Data) > 0 {
|
||||
imageData = refImg.Data
|
||||
} else {
|
||||
// No data available
|
||||
continue
|
||||
}
|
||||
|
||||
mime := refImg.MimeType
|
||||
if mime == "" {
|
||||
mime = "image/png"
|
||||
}
|
||||
refURL = fmt.Sprintf("data:%s;base64,%s", mime, base64.StdEncoding.EncodeToString(imageData))
|
||||
}
|
||||
|
||||
images = append(images, ImageRef{ImageURL: refURL})
|
||||
}
|
||||
|
||||
if len(images) == 0 {
|
||||
return nil, nil, fmt.Errorf("no reference images available")
|
||||
}
|
||||
|
||||
// Build reference image descriptions if available
|
||||
var descParts []string
|
||||
for idx, refImg := range refImgs {
|
||||
if refImg.Description != "" {
|
||||
descParts = append(descParts, fmt.Sprintf("- image_%d.png: %s", idx+1, refImg.Description))
|
||||
}
|
||||
}
|
||||
finalPrompt := prompt
|
||||
if len(descParts) > 0 {
|
||||
finalPrompt = fmt.Sprintf("%s\n\n[Reference Image Roles]\n%s", prompt, strings.Join(descParts, "\n"))
|
||||
slog.Info("openai json edits: appended image descriptions to prompt", "desc_count", len(descParts))
|
||||
}
|
||||
|
||||
body := map[string]any{
|
||||
"model": model,
|
||||
"prompt": finalPrompt,
|
||||
"images": images,
|
||||
"response_format": "b64_json",
|
||||
}
|
||||
|
||||
jsonBody, err := json.Marshal(body)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("marshal request: %w", err)
|
||||
}
|
||||
|
||||
url := strings.TrimRight(apiBase, "/") + "/images/edits"
|
||||
req, err := http.NewRequestWithContext(ctx, "POST", url, bytes.NewReader(jsonBody))
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("create request: %w", err)
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", "Bearer "+apiKey)
|
||||
|
||||
client := &http.Client{}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("http request: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
respBody, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("read response: %w", err)
|
||||
}
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, nil, fmt.Errorf("API error %d: %s", resp.StatusCode, truncateBytes(respBody, 500))
|
||||
}
|
||||
|
||||
var imgResp struct {
|
||||
Data []struct {
|
||||
B64JSON string `json:"b64_json"`
|
||||
} `json:"data"`
|
||||
}
|
||||
if err := json.Unmarshal(respBody, &imgResp); err != nil {
|
||||
return nil, nil, fmt.Errorf("parse response: %w", err)
|
||||
}
|
||||
if len(imgResp.Data) == 0 || imgResp.Data[0].B64JSON == "" {
|
||||
return nil, nil, fmt.Errorf("no image data in response")
|
||||
}
|
||||
|
||||
imageBytes, err := base64.StdEncoding.DecodeString(imgResp.Data[0].B64JSON)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("decode base64: %w", err)
|
||||
}
|
||||
|
||||
return imageBytes, nil, nil
|
||||
}
|
||||
|
||||
// callGeminiNativeImageGen uses the native Gemini generateContent API with responseModalities.
|
||||
// Gemini image models require this endpoint — they don't support OpenAI-compat endpoints.
|
||||
func (t *CreateImageTool) callGeminiNativeImageGen(ctx context.Context, apiKey, apiBase, model, prompt string, params map[string]any) ([]byte, *providers.Usage, error) {
|
||||
@@ -325,9 +862,44 @@ func (t *CreateImageTool) callGeminiNativeImageGen(ctx context.Context, apiKey,
|
||||
|
||||
url := fmt.Sprintf("%s/models/%s:generateContent?key=%s", nativeBase, model, apiKey)
|
||||
|
||||
parts := []map[string]any{
|
||||
{"text": prompt},
|
||||
}
|
||||
if rawImgs, ok := params["ref_images"]; ok {
|
||||
if refImgs, ok := rawImgs.([]*referenceImage); ok {
|
||||
for _, refImg := range refImgs {
|
||||
var dataB64 string
|
||||
var mime string
|
||||
if refImg.Base64 != "" {
|
||||
dataB64 = refImg.Base64
|
||||
mime = refImg.MimeType
|
||||
} else if refImg.URL != "" {
|
||||
dataBytes, contentType, err := t.downloadImageBytes(ctx, refImg.URL)
|
||||
if err == nil {
|
||||
dataB64 = base64.StdEncoding.EncodeToString(dataBytes)
|
||||
mime = contentType
|
||||
} else {
|
||||
slog.Warn("gemini native image gen: failed to download reference image from URL", "url", refImg.URL, "error", err)
|
||||
}
|
||||
}
|
||||
if dataB64 != "" {
|
||||
if mime == "" {
|
||||
mime = "image/png"
|
||||
}
|
||||
parts = append(parts, map[string]any{
|
||||
"inlineData": map[string]any{
|
||||
"mimeType": mime,
|
||||
"data": dataB64,
|
||||
},
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
body := map[string]any{
|
||||
"contents": []map[string]any{
|
||||
{"parts": []map[string]any{{"text": prompt}}},
|
||||
{"parts": parts},
|
||||
},
|
||||
"generationConfig": map[string]any{
|
||||
"responseModalities": []string{"TEXT", "IMAGE"},
|
||||
|
||||
@@ -47,6 +47,26 @@ func callBytePlusImageGen(ctx context.Context, apiKey, apiBase, model, prompt st
|
||||
"response_format": "url",
|
||||
}
|
||||
|
||||
if rawImgs, ok := params["ref_images"]; ok {
|
||||
if refImgs, ok := rawImgs.([]*referenceImage); ok && len(refImgs) > 0 {
|
||||
if len(refImgs) > 1 {
|
||||
slog.Warn("byteplus image gen: provider only supports 1 reference image, using the first one", "count", len(refImgs))
|
||||
}
|
||||
refImg := refImgs[0]
|
||||
var refURL string
|
||||
if refImg.URL != "" {
|
||||
refURL = refImg.URL
|
||||
} else {
|
||||
refMime := refImg.MimeType
|
||||
if refMime == "" {
|
||||
refMime = "image/png"
|
||||
}
|
||||
refURL = fmt.Sprintf("data:%s;base64,%s", refMime, refImg.Base64)
|
||||
}
|
||||
body["image"] = refURL
|
||||
}
|
||||
}
|
||||
|
||||
jsonBody, err := json.Marshal(body)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("marshal request: %w", err)
|
||||
|
||||
@@ -86,18 +86,42 @@ func callDashScopeImageGen(ctx context.Context, apiKey, apiBase, model, prompt s
|
||||
|
||||
endpoint := dashScopeImageEndpoint(apiBase)
|
||||
|
||||
inputBody := map[string]any{
|
||||
"messages": []map[string]any{
|
||||
{"role": "user", "content": prompt},
|
||||
},
|
||||
}
|
||||
parametersBody := map[string]any{
|
||||
"n": 1,
|
||||
"size": size,
|
||||
"prompt_extend": promptExtend,
|
||||
}
|
||||
|
||||
if rawImgs, ok := params["ref_images"]; ok {
|
||||
if refImgs, ok := rawImgs.([]*referenceImage); ok && len(refImgs) > 0 {
|
||||
if len(refImgs) > 1 {
|
||||
slog.Warn("dashscope image gen: provider only supports 1 reference image, using the first one", "count", len(refImgs))
|
||||
}
|
||||
refImg := refImgs[0]
|
||||
var refURL string
|
||||
if refImg.URL != "" {
|
||||
refURL = refImg.URL
|
||||
} else {
|
||||
refMime := refImg.MimeType
|
||||
if refMime == "" {
|
||||
refMime = "image/png"
|
||||
}
|
||||
refURL = fmt.Sprintf("data:%s;base64,%s", refMime, refImg.Base64)
|
||||
}
|
||||
inputBody["ref_img"] = refURL
|
||||
parametersBody["ref_strength"] = refImg.Strength
|
||||
}
|
||||
}
|
||||
|
||||
body := map[string]any{
|
||||
"model": model,
|
||||
"input": map[string]any{
|
||||
"messages": []map[string]any{
|
||||
{"role": "user", "content": prompt},
|
||||
},
|
||||
},
|
||||
"parameters": map[string]any{
|
||||
"n": 1,
|
||||
"size": size,
|
||||
"prompt_extend": promptExtend,
|
||||
},
|
||||
"model": model,
|
||||
"input": inputBody,
|
||||
"parameters": parametersBody,
|
||||
}
|
||||
|
||||
jsonBody, err := json.Marshal(body)
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
@@ -55,6 +56,31 @@ func callMinimaxImageGen(ctx context.Context, apiKey, apiBase, model, prompt str
|
||||
"response_format": "base64",
|
||||
}
|
||||
|
||||
if rawImgs, ok := params["ref_images"]; ok {
|
||||
if refImgs, ok := rawImgs.([]*referenceImage); ok && len(refImgs) > 0 {
|
||||
if len(refImgs) > 1 {
|
||||
slog.Warn("minimax image gen: provider only supports 1 reference image, using the first one", "count", len(refImgs))
|
||||
}
|
||||
refImg := refImgs[0]
|
||||
var refURL string
|
||||
if refImg.URL != "" {
|
||||
refURL = refImg.URL
|
||||
} else {
|
||||
refMime := refImg.MimeType
|
||||
if refMime == "" {
|
||||
refMime = "image/png"
|
||||
}
|
||||
refURL = fmt.Sprintf("data:%s;base64,%s", refMime, refImg.Base64)
|
||||
}
|
||||
body["subject_reference"] = []map[string]any{
|
||||
{
|
||||
"type": "character",
|
||||
"image_file": refURL,
|
||||
},
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
jsonBody, err := json.Marshal(body)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("marshal request: %w", err)
|
||||
|
||||
@@ -2,6 +2,14 @@ package tools
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/nextlevelbuilder/goclaw/internal/providers"
|
||||
@@ -18,7 +26,7 @@ type nativeImageProvider struct {
|
||||
}
|
||||
|
||||
func (p *nativeImageProvider) Name() string { return p.name }
|
||||
func (p *nativeImageProvider) DefaultModel() string { return p.model }
|
||||
func (p *nativeImageProvider) DefaultModel() string { return p.model }
|
||||
func (p *nativeImageProvider) Chat(_ context.Context, _ providers.ChatRequest) (*providers.ChatResponse, error) {
|
||||
return &providers.ChatResponse{}, nil
|
||||
}
|
||||
@@ -256,3 +264,482 @@ func TestCreateImageTool_ThreadsImageModel(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateImageTool_ResolveReferenceImage_Path(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
refFile := filepath.Join(tmpDir, "ref.png")
|
||||
pngBytes := []byte{0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a}
|
||||
if err := os.WriteFile(refFile, pngBytes, 0644); err != nil {
|
||||
t.Fatalf("failed to write temp ref file: %v", err)
|
||||
}
|
||||
|
||||
fakeProvider := &nativeImageProvider{
|
||||
name: "openai-codex",
|
||||
model: "gpt-image-2",
|
||||
returnData: pngBytes,
|
||||
}
|
||||
|
||||
reg := providers.NewRegistry(nil)
|
||||
reg.Register(fakeProvider)
|
||||
|
||||
tool := NewCreateImageTool(reg)
|
||||
ctx := WithToolWorkspace(context.Background(), tmpDir)
|
||||
|
||||
chainJSON := []byte(`{"providers":[{"provider":"openai-codex","model":"gpt-image-2","enabled":true,"timeout":30,"max_retries":1}]}`)
|
||||
settings := BuiltinToolSettings{"create_image": chainJSON}
|
||||
ctx = WithBuiltinToolSettings(ctx, settings)
|
||||
|
||||
result := tool.Execute(ctx, map[string]any{
|
||||
"prompt": "generate a picture",
|
||||
"ref_images": []any{map[string]any{"path": refFile}},
|
||||
})
|
||||
|
||||
if result.IsError {
|
||||
t.Fatalf("Execute returned error: %q", result.ForLLM)
|
||||
}
|
||||
|
||||
if fakeProvider.calledWith == nil {
|
||||
t.Fatal("GenerateImage was not called")
|
||||
}
|
||||
|
||||
if len(fakeProvider.calledWith.RefImages) != 1 {
|
||||
t.Fatalf("expected 1 reference image, got %d", len(fakeProvider.calledWith.RefImages))
|
||||
}
|
||||
if fakeProvider.calledWith.RefImages[0].Base64 == "" {
|
||||
t.Error("RefImages[0].Base64 was not populated")
|
||||
}
|
||||
if fakeProvider.calledWith.RefImages[0].Strength != 0.6 {
|
||||
t.Errorf("RefImages[0].Strength = %f, want 0.6", fakeProvider.calledWith.RefImages[0].Strength)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateImageTool_ResolveReferenceImage_URL(t *testing.T) {
|
||||
fakeProvider := &nativeImageProvider{
|
||||
name: "openai-codex",
|
||||
model: "gpt-image-2",
|
||||
returnData: []byte{0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a},
|
||||
}
|
||||
|
||||
reg := providers.NewRegistry(nil)
|
||||
reg.Register(fakeProvider)
|
||||
|
||||
tool := NewCreateImageTool(reg)
|
||||
ctx := WithToolWorkspace(context.Background(), t.TempDir())
|
||||
|
||||
chainJSON := []byte(`{"providers":[{"provider":"openai-codex","model":"gpt-image-2","enabled":true,"timeout":30,"max_retries":1}]}`)
|
||||
settings := BuiltinToolSettings{"create_image": chainJSON}
|
||||
ctx = WithBuiltinToolSettings(ctx, settings)
|
||||
|
||||
result := tool.Execute(ctx, map[string]any{
|
||||
"prompt": "generate a picture",
|
||||
"ref_images": []any{map[string]any{"url": "https://example.com/ref.png"}},
|
||||
})
|
||||
|
||||
if result.IsError {
|
||||
t.Fatalf("Execute returned error: %q", result.ForLLM)
|
||||
}
|
||||
|
||||
if fakeProvider.calledWith == nil {
|
||||
t.Fatal("GenerateImage was not called")
|
||||
}
|
||||
|
||||
if len(fakeProvider.calledWith.RefImages) != 1 {
|
||||
t.Fatalf("expected 1 reference image, got %d", len(fakeProvider.calledWith.RefImages))
|
||||
}
|
||||
if fakeProvider.calledWith.RefImages[0].URL != "https://example.com/ref.png" {
|
||||
t.Errorf("RefImageUrl = %q, want https://example.com/ref.png", fakeProvider.calledWith.RefImages[0].URL)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateImageTool_ResolveReferenceImage_ID(t *testing.T) {
|
||||
fakeProvider := &nativeImageProvider{
|
||||
name: "openai-codex",
|
||||
model: "gpt-image-2",
|
||||
returnData: []byte{0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a},
|
||||
}
|
||||
|
||||
reg := providers.NewRegistry(nil)
|
||||
reg.Register(fakeProvider)
|
||||
|
||||
tool := NewCreateImageTool(reg)
|
||||
ctx := WithToolWorkspace(context.Background(), t.TempDir())
|
||||
|
||||
// Put mock image content in context
|
||||
mockImg := providers.ImageContent{
|
||||
MimeType: "image/png",
|
||||
Data: "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg==", // 1x1 png base64
|
||||
}
|
||||
ctx = WithMediaImages(ctx, []providers.ImageContent{mockImg})
|
||||
|
||||
chainJSON := []byte(`{"providers":[{"provider":"openai-codex","model":"gpt-image-2","enabled":true,"timeout":30,"max_retries":1}]}`)
|
||||
settings := BuiltinToolSettings{"create_image": chainJSON}
|
||||
ctx = WithBuiltinToolSettings(ctx, settings)
|
||||
|
||||
result := tool.Execute(ctx, map[string]any{
|
||||
"prompt": "generate a picture",
|
||||
"ref_images": []any{map[string]any{"id": "latest"}},
|
||||
})
|
||||
|
||||
if result.IsError {
|
||||
t.Fatalf("Execute returned error: %q", result.ForLLM)
|
||||
}
|
||||
|
||||
if fakeProvider.calledWith == nil {
|
||||
t.Fatal("GenerateImage was not called")
|
||||
}
|
||||
|
||||
if len(fakeProvider.calledWith.RefImages) != 1 {
|
||||
t.Fatalf("expected 1 reference image, got %d", len(fakeProvider.calledWith.RefImages))
|
||||
}
|
||||
if fakeProvider.calledWith.RefImages[0].Base64 != mockImg.Data {
|
||||
t.Errorf("RefImageBase64 = %q, want %q", fakeProvider.calledWith.RefImages[0].Base64, mockImg.Data)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateImageTool_MultipleReferenceImages_MixedSources(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
refFile := filepath.Join(tmpDir, "ref.png")
|
||||
pngBytes := []byte{0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a}
|
||||
if err := os.WriteFile(refFile, pngBytes, 0644); err != nil {
|
||||
t.Fatalf("failed to write temp ref file: %v", err)
|
||||
}
|
||||
|
||||
fakeProvider := &nativeImageProvider{
|
||||
name: "openai-codex",
|
||||
model: "gpt-image-2",
|
||||
returnData: pngBytes,
|
||||
}
|
||||
|
||||
reg := providers.NewRegistry(nil)
|
||||
reg.Register(fakeProvider)
|
||||
|
||||
tool := NewCreateImageTool(reg)
|
||||
ctx := WithToolWorkspace(context.Background(), tmpDir)
|
||||
|
||||
chainJSON := []byte(`{"providers":[{"provider":"openai-codex","model":"gpt-image-2","enabled":true,"timeout":30,"max_retries":1}]}`)
|
||||
settings := BuiltinToolSettings{"create_image": chainJSON}
|
||||
ctx = WithBuiltinToolSettings(ctx, settings)
|
||||
|
||||
result := tool.Execute(ctx, map[string]any{
|
||||
"prompt": "generate a picture",
|
||||
"ref_images": []any{
|
||||
map[string]any{
|
||||
"path": refFile,
|
||||
"strength": 0.9,
|
||||
},
|
||||
map[string]any{
|
||||
"url": "https://example.com/ref2.png",
|
||||
"strength": 0.4,
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
if result.IsError {
|
||||
t.Fatalf("Execute returned error: %q", result.ForLLM)
|
||||
}
|
||||
|
||||
if fakeProvider.calledWith == nil {
|
||||
t.Fatal("GenerateImage was not called")
|
||||
}
|
||||
|
||||
if len(fakeProvider.calledWith.RefImages) != 2 {
|
||||
t.Fatalf("expected 2 reference images, got %d", len(fakeProvider.calledWith.RefImages))
|
||||
}
|
||||
|
||||
// Verify image 1 (path)
|
||||
if fakeProvider.calledWith.RefImages[0].Base64 == "" {
|
||||
t.Error("RefImages[0].Base64 was not populated")
|
||||
}
|
||||
if fakeProvider.calledWith.RefImages[0].Strength != 0.9 {
|
||||
t.Errorf("RefImages[0].Strength = %f, want 0.9", fakeProvider.calledWith.RefImages[0].Strength)
|
||||
}
|
||||
|
||||
// Verify image 2 (url)
|
||||
if fakeProvider.calledWith.RefImages[1].URL != "https://example.com/ref2.png" {
|
||||
t.Errorf("RefImages[1].URL = %q, want https://example.com/ref2.png", fakeProvider.calledWith.RefImages[1].URL)
|
||||
}
|
||||
if fakeProvider.calledWith.RefImages[1].Strength != 0.4 {
|
||||
t.Errorf("RefImages[1].Strength = %f, want 0.4", fakeProvider.calledWith.RefImages[1].Strength)
|
||||
}
|
||||
}
|
||||
|
||||
type fakeCreds struct {
|
||||
key string
|
||||
base string
|
||||
}
|
||||
|
||||
func (c fakeCreds) APIKey() string { return c.key }
|
||||
func (c fakeCreds) APIBase() string { return c.base }
|
||||
|
||||
func TestCreateImageTool_OpenAIEdits_JSON(t *testing.T) {
|
||||
pngBytes := []byte{0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a, 0x00, 0x00, 0x00, 0x0d, 0x49, 0x48, 0x44, 0x52}
|
||||
mockResponse := `{
|
||||
"data": [
|
||||
{
|
||||
"b64_json": "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg=="
|
||||
}
|
||||
]
|
||||
}`
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != "/images/edits" {
|
||||
t.Errorf("expected path /images/edits, got %s", r.URL.Path)
|
||||
}
|
||||
ct := r.Header.Get("Content-Type")
|
||||
if !strings.HasPrefix(ct, "application/json") {
|
||||
t.Errorf("expected application/json content type, got %s", ct)
|
||||
}
|
||||
bodyBytes, err := io.ReadAll(r.Body)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read body: %v", err)
|
||||
}
|
||||
var reqBody struct {
|
||||
Model string `json:"model"`
|
||||
Prompt string `json:"prompt"`
|
||||
ResponseFormat string `json:"response_format"`
|
||||
Images []struct {
|
||||
ImageURL string `json:"image_url"`
|
||||
} `json:"images"`
|
||||
}
|
||||
if err := json.Unmarshal(bodyBytes, &reqBody); err != nil {
|
||||
t.Fatalf("failed to unmarshal body: %v", err)
|
||||
}
|
||||
if reqBody.Prompt != "sunny beach" {
|
||||
t.Errorf("expected prompt 'sunny beach', got %s", reqBody.Prompt)
|
||||
}
|
||||
if reqBody.Model != "gpt-image-2" {
|
||||
t.Errorf("expected model 'gpt-image-2', got %s", reqBody.Model)
|
||||
}
|
||||
if reqBody.ResponseFormat != "b64_json" {
|
||||
t.Errorf("expected response_format 'b64_json', got %s", reqBody.ResponseFormat)
|
||||
}
|
||||
if len(reqBody.Images) != 1 {
|
||||
t.Fatalf("expected 1 image, got %d", len(reqBody.Images))
|
||||
}
|
||||
expectedURL := "data:image/png;base64," + base64.StdEncoding.EncodeToString(pngBytes)
|
||||
if reqBody.Images[0].ImageURL != expectedURL {
|
||||
t.Errorf("image url mismatch")
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(http.StatusOK)
|
||||
w.Write([]byte(mockResponse))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
creds := fakeCreds{key: "fake-key", base: server.URL}
|
||||
reg := providers.NewRegistry(nil)
|
||||
tool := NewCreateImageTool(reg)
|
||||
ctx := WithToolWorkspace(context.Background(), t.TempDir())
|
||||
|
||||
refImg := &referenceImage{
|
||||
Data: pngBytes,
|
||||
MimeType: "image/png",
|
||||
}
|
||||
|
||||
params := map[string]any{
|
||||
"prompt": "sunny beach",
|
||||
"ref_images": []*referenceImage{refImg},
|
||||
"_provider_type": "openai",
|
||||
}
|
||||
|
||||
imageBytes, _, err := tool.callProvider(ctx, creds, "openai", "gpt-image-2", params)
|
||||
if err != nil {
|
||||
t.Fatalf("callProvider returned error: %v", err)
|
||||
}
|
||||
|
||||
if len(imageBytes) == 0 {
|
||||
t.Error("returned imageBytes is empty")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateImageTool_DeduplicateReferenceImages(t *testing.T) {
|
||||
fakeProvider := &nativeImageProvider{
|
||||
name: "openai-codex",
|
||||
model: "gpt-image-2",
|
||||
returnData: []byte{0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a},
|
||||
}
|
||||
|
||||
reg := providers.NewRegistry(nil)
|
||||
reg.Register(fakeProvider)
|
||||
|
||||
tool := NewCreateImageTool(reg)
|
||||
ctx := WithToolWorkspace(context.Background(), t.TempDir())
|
||||
|
||||
chainJSON := []byte(`{"providers":[{"provider":"openai-codex","model":"gpt-image-2","enabled":true,"timeout":30,"max_retries":1}]}`)
|
||||
settings := BuiltinToolSettings{"create_image": chainJSON}
|
||||
ctx = WithBuiltinToolSettings(ctx, settings)
|
||||
|
||||
// Pass duplicate parameters as sent by the user
|
||||
url1 := "https://example.com/image1.png"
|
||||
url2 := "https://example.com/image2.png"
|
||||
|
||||
result := tool.Execute(ctx, map[string]any{
|
||||
"prompt": "generate a picture",
|
||||
"ref_images": []any{
|
||||
map[string]any{
|
||||
"url": url1,
|
||||
"strength": 0.8,
|
||||
},
|
||||
map[string]any{
|
||||
"url": url2,
|
||||
"strength": 0.8,
|
||||
},
|
||||
map[string]any{
|
||||
"url": url1, // Duplicated
|
||||
"strength": 0.8,
|
||||
},
|
||||
map[string]any{
|
||||
"url": url2, // Duplicated
|
||||
"strength": 0.8,
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
if result.IsError {
|
||||
t.Fatalf("Execute returned error: %q", result.ForLLM)
|
||||
}
|
||||
|
||||
if fakeProvider.calledWith == nil {
|
||||
t.Fatal("GenerateImage was not called")
|
||||
}
|
||||
|
||||
// Verify that the length is 2 instead of 4
|
||||
gotLen := len(fakeProvider.calledWith.RefImages)
|
||||
if gotLen != 2 {
|
||||
t.Fatalf("expected 2 unique reference images, got %d (duplicated URL error!)", gotLen)
|
||||
}
|
||||
|
||||
// Verify that the priority order of ref_images is preserved (strength = 0.8)
|
||||
if fakeProvider.calledWith.RefImages[0].Strength != 0.8 {
|
||||
t.Errorf("RefImages[0].Strength = %f, want 0.8", fakeProvider.calledWith.RefImages[0].Strength)
|
||||
}
|
||||
if fakeProvider.calledWith.RefImages[1].Strength != 0.8 {
|
||||
t.Errorf("RefImages[1].Strength = %f, want 0.8", fakeProvider.calledWith.RefImages[1].Strength)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateImageTool_OpenAIEdits_JSON_MultipleImages(t *testing.T) {
|
||||
pngBytes1 := []byte{0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a, 0x00, 0x00, 0x00, 0x0d, 0x49, 0x48, 0x44, 0x52}
|
||||
pngBytes2 := []byte{0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a, 0x00, 0x00, 0x00, 0x0d, 0x49, 0x48, 0x44, 0x53}
|
||||
mockResponse := `{
|
||||
"data": [
|
||||
{
|
||||
"b64_json": "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg=="
|
||||
}
|
||||
]
|
||||
}`
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != "/images/edits" {
|
||||
t.Errorf("expected path /images/edits, got %s", r.URL.Path)
|
||||
}
|
||||
ct := r.Header.Get("Content-Type")
|
||||
if !strings.HasPrefix(ct, "application/json") {
|
||||
t.Errorf("expected application/json content type, got %s", ct)
|
||||
}
|
||||
bodyBytes, err := io.ReadAll(r.Body)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read body: %v", err)
|
||||
}
|
||||
var reqBody struct {
|
||||
Prompt string `json:"prompt"`
|
||||
Images []struct {
|
||||
ImageURL string `json:"image_url"`
|
||||
} `json:"images"`
|
||||
}
|
||||
if err := json.Unmarshal(bodyBytes, &reqBody); err != nil {
|
||||
t.Fatalf("failed to unmarshal body: %v", err)
|
||||
}
|
||||
if reqBody.Prompt != "two captains" {
|
||||
t.Errorf("expected prompt 'two captains', got %s", reqBody.Prompt)
|
||||
}
|
||||
if len(reqBody.Images) != 2 {
|
||||
t.Fatalf("expected 2 images, got %d", len(reqBody.Images))
|
||||
}
|
||||
expectedURL1 := "data:image/png;base64," + base64.StdEncoding.EncodeToString(pngBytes1)
|
||||
expectedURL2 := "data:image/png;base64," + base64.StdEncoding.EncodeToString(pngBytes2)
|
||||
if reqBody.Images[0].ImageURL != expectedURL1 || reqBody.Images[1].ImageURL != expectedURL2 {
|
||||
t.Error("images mismatch")
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(http.StatusOK)
|
||||
w.Write([]byte(mockResponse))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
creds := fakeCreds{key: "fake-key", base: server.URL}
|
||||
reg := providers.NewRegistry(nil)
|
||||
tool := NewCreateImageTool(reg)
|
||||
ctx := WithToolWorkspace(context.Background(), t.TempDir())
|
||||
|
||||
refImg1 := &referenceImage{Data: pngBytes1, MimeType: "image/png"}
|
||||
refImg2 := &referenceImage{Data: pngBytes2, MimeType: "image/png"}
|
||||
|
||||
params := map[string]any{
|
||||
"prompt": "two captains",
|
||||
"ref_images": []*referenceImage{refImg1, refImg2},
|
||||
"_provider_type": "openai",
|
||||
}
|
||||
|
||||
imageBytes, _, err := tool.callProvider(ctx, creds, "openai", "gpt-image-2", params)
|
||||
if err != nil {
|
||||
t.Fatalf("callProvider returned error: %v", err)
|
||||
}
|
||||
if len(imageBytes) == 0 {
|
||||
t.Error("returned imageBytes is empty")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateImageTool_OpenAIEdits_JSON_WithDescription(t *testing.T) {
|
||||
pngBytes := []byte{0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a, 0x00, 0x00, 0x00, 0x0d, 0x49, 0x48, 0x44, 0x52}
|
||||
mockResponse := `{
|
||||
"data": [
|
||||
{
|
||||
"b64_json": "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg=="
|
||||
}
|
||||
]
|
||||
}`
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
bodyBytes, err := io.ReadAll(r.Body)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read body: %v", err)
|
||||
}
|
||||
var reqBody struct {
|
||||
Prompt string `json:"prompt"`
|
||||
}
|
||||
if err := json.Unmarshal(bodyBytes, &reqBody); err != nil {
|
||||
t.Fatalf("failed to unmarshal body: %v", err)
|
||||
}
|
||||
expectedPrompt := "two captains\n\n[Reference Image Roles]\n- image_1.png: Lâm\n- image_2.png: Quân"
|
||||
if reqBody.Prompt != expectedPrompt {
|
||||
t.Errorf("expected prompt %q, got %q", expectedPrompt, reqBody.Prompt)
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(http.StatusOK)
|
||||
w.Write([]byte(mockResponse))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
creds := fakeCreds{key: "fake-key", base: server.URL}
|
||||
reg := providers.NewRegistry(nil)
|
||||
tool := NewCreateImageTool(reg)
|
||||
ctx := WithToolWorkspace(context.Background(), t.TempDir())
|
||||
|
||||
refImg1 := &referenceImage{Data: pngBytes, MimeType: "image/png", Description: "Lâm"}
|
||||
refImg2 := &referenceImage{Data: pngBytes, MimeType: "image/png", Description: "Quân"}
|
||||
|
||||
params := map[string]any{
|
||||
"prompt": "two captains",
|
||||
"ref_images": []*referenceImage{refImg1, refImg2},
|
||||
"_provider_type": "openai",
|
||||
}
|
||||
|
||||
_, _, err := tool.callProvider(ctx, creds, "openai", "gpt-image-2", params)
|
||||
if err != nil {
|
||||
t.Fatalf("callProvider returned error: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,130 @@
|
||||
package tools
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/nextlevelbuilder/goclaw/internal/security"
|
||||
)
|
||||
|
||||
// downloadImageBytes fetches caller-supplied reference image URLs server-side, so
|
||||
// it must reject SSRF targets and cap the response size. These tests lock that down.
|
||||
|
||||
func TestDownloadImageBytes_BlocksSSRFTargets(t *testing.T) {
|
||||
// Default: loopback/private/link-local are blocked by the SSRF guard.
|
||||
tool := NewCreateImageTool(nil)
|
||||
cases := []struct {
|
||||
name string
|
||||
url string
|
||||
}{
|
||||
{"loopback", "http://127.0.0.1:9/x.png"},
|
||||
{"private", "http://10.0.0.5/x.png"},
|
||||
{"link_local_metadata", "http://169.254.169.254/latest/meta-data/"},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
_, _, err := tool.downloadImageBytes(context.Background(), tc.url)
|
||||
if err == nil {
|
||||
t.Fatalf("expected SSRF guard to block %s, got nil error", tc.url)
|
||||
}
|
||||
if !strings.Contains(err.Error(), "invalid reference image URL") {
|
||||
t.Errorf("expected SSRF validation error, got: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestDownloadImageBytes_OversizedRejected(t *testing.T) {
|
||||
security.SetAllowLoopbackForTest(true)
|
||||
defer security.SetAllowLoopbackForTest(false)
|
||||
|
||||
// Shrink the cap so we don't have to transfer 20 MB.
|
||||
orig := refImageMaxBytes
|
||||
refImageMaxBytes = 16
|
||||
defer func() { refImageMaxBytes = orig }()
|
||||
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.Header().Set("Content-Type", "image/png")
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write([]byte(strings.Repeat("A", 1024))) // well over 16 bytes
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
tool := NewCreateImageTool(nil)
|
||||
_, _, err := tool.downloadImageBytes(context.Background(), srv.URL)
|
||||
if err == nil {
|
||||
t.Fatal("expected oversized reference image to be rejected, got nil error")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "exceeds maximum size") {
|
||||
t.Errorf("expected size-limit error, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDownloadImageBytes_HappyPath(t *testing.T) {
|
||||
security.SetAllowLoopbackForTest(true)
|
||||
defer security.SetAllowLoopbackForTest(false)
|
||||
|
||||
want := []byte("\x89PNG\r\n\x1a\nfake-image-bytes")
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.Header().Set("Content-Type", "image/png")
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write(want)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
tool := NewCreateImageTool(nil)
|
||||
got, contentType, err := tool.downloadImageBytes(context.Background(), srv.URL)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if string(got) != string(want) {
|
||||
t.Errorf("body = %q, want %q", got, want)
|
||||
}
|
||||
if contentType != "image/png" {
|
||||
t.Errorf("contentType = %q, want image/png", contentType)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDownloadImageBytes_RejectsRedirect(t *testing.T) {
|
||||
security.SetAllowLoopbackForTest(true)
|
||||
defer security.SetAllowLoopbackForTest(false)
|
||||
|
||||
// SafeClient never follows redirects: a 3xx is returned as-is and rejected
|
||||
// by the non-200 status check (prevents redirect-to-internal SSRF).
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
http.Redirect(w, &http.Request{}, "http://169.254.169.254/", http.StatusFound)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
tool := NewCreateImageTool(nil)
|
||||
_, _, err := tool.downloadImageBytes(context.Background(), srv.URL)
|
||||
if err == nil {
|
||||
t.Fatal("expected redirect to be rejected, got nil error")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "HTTP error") {
|
||||
t.Errorf("expected non-200 status error from unfollowed redirect, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveReferenceImages_RejectsNonHTTPScheme(t *testing.T) {
|
||||
tool := NewCreateImageTool(nil)
|
||||
for _, bad := range []string{"file:///etc/passwd", "gopher://127.0.0.1/", "data:text/plain,hi"} {
|
||||
t.Run(bad, func(t *testing.T) {
|
||||
args := map[string]any{
|
||||
"ref_images": []any{
|
||||
map[string]any{"url": bad},
|
||||
},
|
||||
}
|
||||
_, err := tool.resolveReferenceImages(context.Background(), args)
|
||||
if err == nil {
|
||||
t.Fatalf("expected %q to be rejected as non-http(s), got nil error", bad)
|
||||
}
|
||||
if !strings.Contains(err.Error(), "must be http(s)") {
|
||||
t.Errorf("expected scheme error, got: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
Reference in new issue
Block a user