Cloudflare Workers AI changed the API for flux-2-dev, flux-2-klein-4b,
and flux-2-klein-9b to require multipart/form-data (instead of JSON) and
now returns {"image":"<base64>"} instead of raw PNG bytes.
- Add requiresMultipart() helper for the three FLUX.2 models
- callImageAPI builds multipart body for those models, JSON for others
- Parse {"image":"<base64>"} JSON response; fall back to raw bytes for legacy models
- Use "steps" field name (not "num_steps") in multipart forms per CF docs
- Book page: capture and display actual backend error message instead of blank 'Error'
476 lines
15 KiB
Go
476 lines
15 KiB
Go
// Image generation via Cloudflare Workers AI text-to-image models.
|
||
//
|
||
// API reference:
|
||
//
|
||
// POST https://api.cloudflare.com/client/v4/accounts/{accountID}/ai/run/{model}
|
||
// Authorization: Bearer {apiToken}
|
||
//
|
||
// FLUX.2 models (flux-2-dev, flux-2-klein-4b, flux-2-klein-9b):
|
||
//
|
||
// Content-Type: multipart/form-data
|
||
// Fields: prompt, num_steps, width, height, guidance, image_b64 (optional)
|
||
// Response: { "image": "<base64 JPEG>" }
|
||
//
|
||
// Other models (flux-1-schnell, SDXL, SD 1.5):
|
||
//
|
||
// Content-Type: application/json
|
||
// Body: { "prompt": "...", "num_steps": 20 }
|
||
// Response: { "image": "<base64>" } or raw bytes depending on model
|
||
//
|
||
// Reference-image request (FLUX.2):
|
||
//
|
||
// Same multipart form; include image_b64 field with base64-encoded reference.
|
||
//
|
||
// Reference-image request (SD img2img):
|
||
//
|
||
// JSON body: { "prompt": "...", "image": [r,g,b,a,...], "strength": 0.75 }
|
||
//
|
||
// Recommended models for LibNovel:
|
||
// - Book covers (no reference): flux-2-dev, flux-2-klein-9b, lucid-origin
|
||
// - Chapter images (speed): flux-2-klein-4b, flux-1-schnell
|
||
// - With reference image: flux-2-dev, flux-2-klein-9b, sd-v1-5-img2img
|
||
package cfai
|
||
|
||
import (
|
||
"bytes"
|
||
"context"
|
||
"encoding/base64"
|
||
"encoding/json"
|
||
"fmt"
|
||
"image"
|
||
"image/draw"
|
||
"image/jpeg"
|
||
_ "image/jpeg" // register JPEG decoder
|
||
"image/png"
|
||
_ "image/png" // register PNG decoder
|
||
"io"
|
||
"mime/multipart"
|
||
"net/http"
|
||
"strings"
|
||
"time"
|
||
)
|
||
|
||
// ImageModel identifies a Cloudflare Workers AI text-to-image model.
|
||
type ImageModel string
|
||
|
||
const (
|
||
// ImageModelFlux2Dev — best quality, multi-reference. Recommended for covers.
|
||
ImageModelFlux2Dev ImageModel = "@cf/black-forest-labs/flux-2-dev"
|
||
// ImageModelFlux2Klein9B — 9B params, multi-reference. Good for covers.
|
||
ImageModelFlux2Klein9B ImageModel = "@cf/black-forest-labs/flux-2-klein-9b"
|
||
// ImageModelFlux2Klein4B — ultra-fast, unified gen+edit. Recommended for chapters.
|
||
ImageModelFlux2Klein4B ImageModel = "@cf/black-forest-labs/flux-2-klein-4b"
|
||
// ImageModelFlux1Schnell — fastest, text-only. Good for quick illustrations.
|
||
ImageModelFlux1Schnell ImageModel = "@cf/black-forest-labs/flux-1-schnell"
|
||
// ImageModelSDXLLightning — fast 1024px generation.
|
||
ImageModelSDXLLightning ImageModel = "@cf/bytedance/stable-diffusion-xl-lightning"
|
||
// ImageModelSD15Img2Img — explicit img2img with flat RGBA reference.
|
||
ImageModelSD15Img2Img ImageModel = "@cf/runwayml/stable-diffusion-v1-5-img2img"
|
||
// ImageModelSDXLBase — Stability AI SDXL base.
|
||
ImageModelSDXLBase ImageModel = "@cf/stabilityai/stable-diffusion-xl-base-1.0"
|
||
// ImageModelLucidOrigin — Leonardo AI; strong prompt adherence.
|
||
ImageModelLucidOrigin ImageModel = "@cf/leonardo/lucid-origin"
|
||
// ImageModelPhoenix10 — Leonardo AI; accurate text rendering.
|
||
ImageModelPhoenix10 ImageModel = "@cf/leonardo/phoenix-1.0"
|
||
|
||
// DefaultImageModel is the default model for book-cover generation.
|
||
DefaultImageModel = ImageModelFlux2Dev
|
||
)
|
||
|
||
// ImageModelInfo describes a single image generation model.
|
||
type ImageModelInfo struct {
|
||
ID string `json:"id"`
|
||
Label string `json:"label"`
|
||
Provider string `json:"provider"`
|
||
SupportsRef bool `json:"supports_ref"`
|
||
RecommendedFor []string `json:"recommended_for"` // "cover" and/or "chapter"
|
||
Description string `json:"description"`
|
||
}
|
||
|
||
// AllImageModels returns metadata about every supported image model.
|
||
func AllImageModels() []ImageModelInfo {
|
||
return []ImageModelInfo{
|
||
{
|
||
ID: string(ImageModelFlux2Dev), Label: "FLUX.2 Dev", Provider: "Black Forest Labs",
|
||
SupportsRef: true, RecommendedFor: []string{"cover"},
|
||
Description: "Best quality; multi-reference editing. Recommended for book covers.",
|
||
},
|
||
{
|
||
ID: string(ImageModelFlux2Klein9B), Label: "FLUX.2 Klein 9B", Provider: "Black Forest Labs",
|
||
SupportsRef: true, RecommendedFor: []string{"cover"},
|
||
Description: "9B parameters with multi-reference support.",
|
||
},
|
||
{
|
||
ID: string(ImageModelFlux2Klein4B), Label: "FLUX.2 Klein 4B", Provider: "Black Forest Labs",
|
||
SupportsRef: true, RecommendedFor: []string{"chapter"},
|
||
Description: "Ultra-fast unified gen+edit. Recommended for chapter images.",
|
||
},
|
||
{
|
||
ID: string(ImageModelFlux1Schnell), Label: "FLUX.1 Schnell", Provider: "Black Forest Labs",
|
||
SupportsRef: false, RecommendedFor: []string{"chapter"},
|
||
Description: "Fastest inference. Good for quick chapter illustrations.",
|
||
},
|
||
{
|
||
ID: string(ImageModelSDXLLightning), Label: "SDXL Lightning", Provider: "ByteDance",
|
||
SupportsRef: false, RecommendedFor: []string{"chapter"},
|
||
Description: "Lightning-fast 1024px images in a few steps.",
|
||
},
|
||
{
|
||
ID: string(ImageModelSD15Img2Img), Label: "SD 1.5 img2img", Provider: "RunwayML",
|
||
SupportsRef: true, RecommendedFor: []string{"cover", "chapter"},
|
||
Description: "Explicit img2img: generates from a reference image + prompt.",
|
||
},
|
||
{
|
||
ID: string(ImageModelSDXLBase), Label: "SDXL Base 1.0", Provider: "Stability AI",
|
||
SupportsRef: false, RecommendedFor: []string{"cover"},
|
||
Description: "Stable Diffusion XL base model.",
|
||
},
|
||
{
|
||
ID: string(ImageModelLucidOrigin), Label: "Lucid Origin", Provider: "Leonardo AI",
|
||
SupportsRef: false, RecommendedFor: []string{"cover"},
|
||
Description: "Highly prompt-responsive; strong graphic design and HD renders.",
|
||
},
|
||
{
|
||
ID: string(ImageModelPhoenix10), Label: "Phoenix 1.0", Provider: "Leonardo AI",
|
||
SupportsRef: false, RecommendedFor: []string{"cover"},
|
||
Description: "Exceptional prompt adherence; accurate text rendering.",
|
||
},
|
||
}
|
||
}
|
||
|
||
// ImageRequest is the input to GenerateImage / GenerateImageFromReference.
|
||
type ImageRequest struct {
|
||
// Prompt is the text description of the desired image.
|
||
Prompt string
|
||
// Model is the CF Workers AI model. Defaults to DefaultImageModel when empty.
|
||
Model ImageModel
|
||
// NumSteps controls inference quality (default 20). Range: 1–20.
|
||
NumSteps int
|
||
// Width and Height in pixels. 0 = model default (typically 1024x1024).
|
||
Width, Height int
|
||
// Guidance controls prompt adherence (default 7.5).
|
||
Guidance float64
|
||
// Strength for img2img: 0.0 = copy reference, 1.0 = ignore reference (default 0.75).
|
||
Strength float64
|
||
}
|
||
|
||
// ImageGenClient generates images via Cloudflare Workers AI.
|
||
type ImageGenClient interface {
|
||
// GenerateImage creates an image from a text prompt only.
|
||
// Returns raw PNG bytes.
|
||
GenerateImage(ctx context.Context, req ImageRequest) ([]byte, error)
|
||
|
||
// GenerateImageFromReference creates an image from a text prompt + reference image.
|
||
// refImage should be PNG or JPEG bytes. Returns raw PNG bytes.
|
||
GenerateImageFromReference(ctx context.Context, req ImageRequest, refImage []byte) ([]byte, error)
|
||
|
||
// Models returns metadata about all supported image models.
|
||
Models() []ImageModelInfo
|
||
}
|
||
|
||
// imageGenHTTPClient is the concrete CF AI image generation client.
|
||
type imageGenHTTPClient struct {
|
||
accountID string
|
||
apiToken string
|
||
http *http.Client
|
||
}
|
||
|
||
// NewImageGen returns an ImageGenClient for the given Cloudflare account.
|
||
func NewImageGen(accountID, apiToken string) ImageGenClient {
|
||
return &imageGenHTTPClient{
|
||
accountID: accountID,
|
||
apiToken: apiToken,
|
||
http: &http.Client{Timeout: 5 * time.Minute},
|
||
}
|
||
}
|
||
|
||
// requiresMultipart reports whether the model requires a multipart/form-data
|
||
// request body instead of JSON. FLUX.2 models on Cloudflare Workers AI changed
|
||
// their API to require multipart and return {"image":"<base64>"} instead of
|
||
// raw image bytes.
|
||
func requiresMultipart(model ImageModel) bool {
|
||
switch model {
|
||
case ImageModelFlux2Dev, ImageModelFlux2Klein4B, ImageModelFlux2Klein9B:
|
||
return true
|
||
default:
|
||
return false
|
||
}
|
||
}
|
||
|
||
// GenerateImage generates an image from text only.
|
||
func (c *imageGenHTTPClient) GenerateImage(ctx context.Context, req ImageRequest) ([]byte, error) {
|
||
req = applyImageDefaults(req)
|
||
|
||
// FLUX.2 multipart models use "steps"; JSON models use "num_steps".
|
||
stepsKey := "num_steps"
|
||
if requiresMultipart(req.Model) {
|
||
stepsKey = "steps"
|
||
}
|
||
|
||
fields := map[string]any{
|
||
"prompt": req.Prompt,
|
||
stepsKey: req.NumSteps,
|
||
}
|
||
if req.Width > 0 {
|
||
fields["width"] = req.Width
|
||
}
|
||
if req.Height > 0 {
|
||
fields["height"] = req.Height
|
||
}
|
||
if req.Guidance > 0 {
|
||
fields["guidance"] = req.Guidance
|
||
}
|
||
return c.callImageAPI(ctx, req.Model, fields, nil)
|
||
}
|
||
|
||
// refImageMaxDim is the maximum dimension (width or height) for reference images
|
||
// sent to Cloudflare Workers AI. CF's JSON body limit is ~4 MB; a 768px JPEG
|
||
// stays well under that while preserving enough detail for img2img guidance.
|
||
const refImageMaxDim = 768
|
||
|
||
// GenerateImageFromReference generates an image from a text prompt + reference image.
|
||
func (c *imageGenHTTPClient) GenerateImageFromReference(ctx context.Context, req ImageRequest, refImage []byte) ([]byte, error) {
|
||
if len(refImage) == 0 {
|
||
return c.GenerateImage(ctx, req)
|
||
}
|
||
req = applyImageDefaults(req)
|
||
|
||
// Shrink the reference image if it exceeds the safe payload size.
|
||
refImage = resizeRefImage(refImage, refImageMaxDim)
|
||
|
||
// FLUX.2 multipart models use "steps"; JSON models use "num_steps".
|
||
stepsKey := "num_steps"
|
||
if requiresMultipart(req.Model) {
|
||
stepsKey = "steps"
|
||
}
|
||
|
||
fields := map[string]any{
|
||
"prompt": req.Prompt,
|
||
stepsKey: req.NumSteps,
|
||
}
|
||
if req.Width > 0 {
|
||
fields["width"] = req.Width
|
||
}
|
||
if req.Height > 0 {
|
||
fields["height"] = req.Height
|
||
}
|
||
if req.Guidance > 0 {
|
||
fields["guidance"] = req.Guidance
|
||
}
|
||
|
||
if requiresMultipart(req.Model) {
|
||
// FLUX.2: reference image sent as base64 form field "image_b64".
|
||
fields["image_b64"] = base64.StdEncoding.EncodeToString(refImage)
|
||
if req.Strength > 0 {
|
||
fields["strength"] = req.Strength
|
||
}
|
||
return c.callImageAPI(ctx, req.Model, fields, nil)
|
||
}
|
||
|
||
if req.Model == ImageModelSD15Img2Img {
|
||
pixels, err := decodeImageToRGBA(refImage)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("cfai/image: decode reference: %w", err)
|
||
}
|
||
strength := req.Strength
|
||
if strength <= 0 {
|
||
strength = 0.75
|
||
}
|
||
fields["image"] = pixels
|
||
fields["strength"] = strength
|
||
return c.callImageAPI(ctx, req.Model, fields, nil)
|
||
}
|
||
|
||
// Other FLUX models: image_b64 JSON field.
|
||
fields["image_b64"] = base64.StdEncoding.EncodeToString(refImage)
|
||
if req.Strength > 0 {
|
||
fields["strength"] = req.Strength
|
||
}
|
||
return c.callImageAPI(ctx, req.Model, fields, nil)
|
||
}
|
||
|
||
// Models returns all supported image model metadata.
|
||
func (c *imageGenHTTPClient) Models() []ImageModelInfo {
|
||
return AllImageModels()
|
||
}
|
||
|
||
func (c *imageGenHTTPClient) callImageAPI(ctx context.Context, model ImageModel, fields map[string]any, _ []byte) ([]byte, error) {
|
||
cfURL := fmt.Sprintf("https://api.cloudflare.com/client/v4/accounts/%s/ai/run/%s",
|
||
c.accountID, string(model))
|
||
|
||
var (
|
||
bodyReader io.Reader
|
||
contentType string
|
||
)
|
||
|
||
if requiresMultipart(model) {
|
||
// Build a multipart/form-data body from the fields map.
|
||
// All values are serialised to their string representation.
|
||
var buf bytes.Buffer
|
||
mw := multipart.NewWriter(&buf)
|
||
for k, v := range fields {
|
||
var strVal string
|
||
switch tv := v.(type) {
|
||
case string:
|
||
strVal = tv
|
||
default:
|
||
encoded, merr := json.Marshal(tv)
|
||
if merr != nil {
|
||
return nil, fmt.Errorf("cfai/image: marshal field %q: %w", k, merr)
|
||
}
|
||
strVal = strings.Trim(string(encoded), `"`)
|
||
}
|
||
if werr := mw.WriteField(k, strVal); werr != nil {
|
||
return nil, fmt.Errorf("cfai/image: write field %q: %w", k, werr)
|
||
}
|
||
}
|
||
if cerr := mw.Close(); cerr != nil {
|
||
return nil, fmt.Errorf("cfai/image: close multipart writer: %w", cerr)
|
||
}
|
||
bodyReader = &buf
|
||
contentType = mw.FormDataContentType()
|
||
} else {
|
||
encoded, merr := json.Marshal(fields)
|
||
if merr != nil {
|
||
return nil, fmt.Errorf("cfai/image: marshal: %w", merr)
|
||
}
|
||
bodyReader = bytes.NewReader(encoded)
|
||
contentType = "application/json"
|
||
}
|
||
|
||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, cfURL, bodyReader)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("cfai/image: build request: %w", err)
|
||
}
|
||
req.Header.Set("Authorization", "Bearer "+c.apiToken)
|
||
req.Header.Set("Content-Type", contentType)
|
||
|
||
resp, err := c.http.Do(req)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("cfai/image: http: %w", err)
|
||
}
|
||
defer resp.Body.Close()
|
||
|
||
respBody, err := io.ReadAll(resp.Body)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("cfai/image: read response: %w", err)
|
||
}
|
||
|
||
if resp.StatusCode != http.StatusOK {
|
||
msg := string(respBody)
|
||
if len(msg) > 300 {
|
||
msg = msg[:300]
|
||
}
|
||
return nil, fmt.Errorf("cfai/image: model %s returned %d: %s", model, resp.StatusCode, msg)
|
||
}
|
||
|
||
// Try to parse as {"image": "<base64>"} first (FLUX.2 and newer models).
|
||
// Fall back to treating the body as raw image bytes for legacy models.
|
||
var jsonResp struct {
|
||
Image string `json:"image"`
|
||
}
|
||
if jerr := json.Unmarshal(respBody, &jsonResp); jerr == nil && jsonResp.Image != "" {
|
||
imgBytes, decErr := base64.StdEncoding.DecodeString(jsonResp.Image)
|
||
if decErr != nil {
|
||
// Try raw (no padding) base64
|
||
imgBytes, decErr = base64.RawStdEncoding.DecodeString(jsonResp.Image)
|
||
if decErr != nil {
|
||
return nil, fmt.Errorf("cfai/image: decode base64 response: %w", decErr)
|
||
}
|
||
}
|
||
return imgBytes, nil
|
||
}
|
||
|
||
// Legacy: model returned raw image bytes directly.
|
||
return respBody, nil
|
||
}
|
||
|
||
func applyImageDefaults(req ImageRequest) ImageRequest {
|
||
if req.Model == "" {
|
||
req.Model = DefaultImageModel
|
||
}
|
||
if req.NumSteps <= 0 {
|
||
req.NumSteps = 20
|
||
}
|
||
return req
|
||
}
|
||
|
||
// resizeRefImage down-scales an image so that its longest side is at most maxDim
|
||
// pixels, then re-encodes it as JPEG (quality 85). If the image is already small
|
||
// enough, or if decoding fails, the original bytes are returned unchanged.
|
||
// This keeps the JSON payload well under Cloudflare Workers AI's 4 MB body limit.
|
||
func resizeRefImage(data []byte, maxDim int) []byte {
|
||
src, format, err := image.Decode(bytes.NewReader(data))
|
||
if err != nil {
|
||
return data
|
||
}
|
||
b := src.Bounds()
|
||
w, h := b.Dx(), b.Dy()
|
||
|
||
longest := w
|
||
if h > longest {
|
||
longest = h
|
||
}
|
||
if longest <= maxDim {
|
||
return data // already fits
|
||
}
|
||
|
||
// Compute target dimensions preserving aspect ratio.
|
||
scale := float64(maxDim) / float64(longest)
|
||
newW := int(float64(w)*scale + 0.5)
|
||
newH := int(float64(h)*scale + 0.5)
|
||
if newW < 1 {
|
||
newW = 1
|
||
}
|
||
if newH < 1 {
|
||
newH = 1
|
||
}
|
||
|
||
// Nearest-neighbour downsample (no extra deps, sufficient for reference guidance).
|
||
dst := image.NewRGBA(image.Rect(0, 0, newW, newH))
|
||
for y := 0; y < newH; y++ {
|
||
for x := 0; x < newW; x++ {
|
||
srcX := b.Min.X + int(float64(x)/scale)
|
||
srcY := b.Min.Y + int(float64(y)/scale)
|
||
draw.Draw(dst, image.Rect(x, y, x+1, y+1), src, image.Pt(srcX, srcY), draw.Src)
|
||
}
|
||
}
|
||
|
||
var buf bytes.Buffer
|
||
if format == "jpeg" {
|
||
if encErr := jpeg.Encode(&buf, dst, &jpeg.Options{Quality: 85}); encErr != nil {
|
||
return data
|
||
}
|
||
} else {
|
||
if encErr := png.Encode(&buf, dst); encErr != nil {
|
||
return data
|
||
}
|
||
}
|
||
return buf.Bytes()
|
||
}
|
||
|
||
// decodeImageToRGBA decodes PNG/JPEG bytes to a flat []uint8 RGBA pixel array
|
||
// required by the stable-diffusion-v1-5-img2img model.
|
||
func decodeImageToRGBA(data []byte) ([]uint8, error) {
|
||
img, _, err := image.Decode(bytes.NewReader(data))
|
||
if err != nil {
|
||
return nil, fmt.Errorf("decode image: %w", err)
|
||
}
|
||
bounds := img.Bounds()
|
||
w := bounds.Max.X - bounds.Min.X
|
||
h := bounds.Max.Y - bounds.Min.Y
|
||
pixels := make([]uint8, w*h*4)
|
||
idx := 0
|
||
for y := bounds.Min.Y; y < bounds.Max.Y; y++ {
|
||
for x := bounds.Min.X; x < bounds.Max.X; x++ {
|
||
r, g, b, a := img.At(x, y).RGBA()
|
||
pixels[idx] = uint8(r >> 8)
|
||
pixels[idx+1] = uint8(g >> 8)
|
||
pixels[idx+2] = uint8(b >> 8)
|
||
pixels[idx+3] = uint8(a >> 8)
|
||
idx += 4
|
||
}
|
||
}
|
||
return pixels, nil
|
||
}
|