package service import ( "bytes" "context" "crypto/rand" "encoding/json" "fmt" "image" "image/color" "image/png" "io" "log" "net/http" "strings" "time" "gen2d/internal/config" ) // imgCfg 保存文生图配置,由 main 通过 InitImageGenConfig 注入。 var imgCfg config.ImageGenConfig // InitImageGenConfig 注入文生图配置。 func InitImageGenConfig(cfg config.ImageGenConfig) { imgCfg = cfg } // ======================== GPT Image 2 API 类型 ======================== // genSubmitReq 提交生图任务请求体。 type genSubmitReq struct { Prompt string `json:"prompt"` AspectRatio string `json:"aspect_ratio,omitempty"` ReferenceImages []string `json:"reference_images,omitempty"` XChannel string `json:"x_channel,omitempty"` } // genSubmitResp 提交生图任务响应体。 type genSubmitResp struct { Code int `json:"code"` Message string `json:"message"` Data struct { TaskID string `json:"task_id"` Status string `json:"status"` } `json:"data"` } // genStatusReq 查询任务状态请求体。 type genStatusReq struct { TaskID string `json:"task_id"` } // genStatusResp 查询任务状态响应体。 type genStatusResp struct { Code int `json:"code"` Message string `json:"message"` Data struct { TaskID string `json:"task_id"` Status string `json:"status"` Progress int `json:"progress"` ResultImageURL string `json:"result_image_url"` ErrorMessage string `json:"error_message"` } `json:"data"` } // ======================== 文生图 ======================== // GenerateImages 通过 GPT Image 2 异步 API 生成图片,未配置 key 时回退到 mock。 func GenerateImages(ctx context.Context, prompt string, params AssetParams) ([]GeneratedImage, error) { count := 1 if params.Frames.Directions > 0 && params.Frames.FramesPerDirection > 0 { count = params.Frames.Directions * params.Frames.FramesPerDirection } if imgCfg.APIKey != "" { log.Println("[inference] using image gen async API") return generateAsync(ctx, prompt, count, nil) } size := params.Resolution if size <= 0 { size = 64 } log.Println("[inference] image API key not configured, using mock") return generateMockImages(size, count) } // ======================== 图片编辑 ======================== // EditImages 图片编辑接口,以参考图模式提交 GPT Image 2 编辑任务。 func EditImages(ctx context.Context, imageData []byte, prompt string, count int) ([]GeneratedImage, error) { if imgCfg.APIKey == "" { return nil, fmt.Errorf("image API key not configured") } log.Println("[inference] using image gen async API for edit") // 暂不支持参考图 URL 模式,提示用户 return nil, fmt.Errorf("image edit with reference image: not yet implemented for async API") } // ======================== 异步核心 ======================== // generateAsync 并行提交 count 个任务,轮询完成后下载图片。 func generateAsync(ctx context.Context, prompt string, count int, refImages []string) ([]GeneratedImage, error) { type submitResult struct { index int taskID string err error } results := make(chan submitResult, count) for i := 0; i < count; i++ { go func(idx int) { taskID, err := submitTask(ctx, prompt, refImages) results <- submitResult{index: idx, taskID: taskID, err: err} }(i) } taskIDs := make([]string, count) for i := 0; i < count; i++ { r := <-results if r.err != nil { return nil, fmt.Errorf("submit task %d: %w", r.index, r.err) } taskIDs[r.index] = r.taskID } type imageResult struct { index int image GeneratedImage err error } imgResults := make(chan imageResult, count) for i, tid := range taskIDs { go func(idx int, taskID string) { img, err := pollAndDownload(ctx, taskID) imgResults <- imageResult{index: idx, image: img, err: err} }(i, tid) } images := make([]GeneratedImage, count) for i := 0; i < count; i++ { r := <-imgResults if r.err != nil { return nil, fmt.Errorf("task %d: %w", r.index, r.err) } images[r.index] = r.image } return images, nil } // submitTask 提交生图任务,返回 taskID。 func submitTask(ctx context.Context, prompt string, refImages []string) (string, error) { reqBody := genSubmitReq{ Prompt: prompt, AspectRatio: imgCfg.AspectRatio, ReferenceImages: refImages, XChannel: imgCfg.XChannel, } body, err := json.Marshal(reqBody) if err != nil { return "", fmt.Errorf("marshal: %w", err) } url := strings.TrimRight(imgCfg.BaseURL, "/") + "/gpt-image2/generate" req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(body)) if err != nil { return "", fmt.Errorf("create request: %w", err) } req.Header.Set("Content-Type", "application/json") req.Header.Set("Authorization", "Bearer "+imgCfg.APIKey) resp, err := http.DefaultClient.Do(req) if err != nil { return "", fmt.Errorf("send: %w", err) } defer resp.Body.Close() var sr genSubmitResp if err := json.NewDecoder(resp.Body).Decode(&sr); err != nil { return "", fmt.Errorf("decode: %w", err) } if sr.Code != 200 { return "", fmt.Errorf("submit failed: %s", sr.Message) } log.Printf("[inference] task submitted: %s", sr.Data.TaskID) return sr.Data.TaskID, nil } // pollAndDownload 轮询任务直到完成,下载图片。 func pollAndDownload(ctx context.Context, taskID string) (GeneratedImage, error) { pollInterval := imgCfg.PollInterval if pollInterval <= 0 { pollInterval = 3 } maxWait := imgCfg.PollMaxWait if maxWait <= 0 { maxWait = 120 } deadline := time.Now().Add(time.Duration(maxWait) * time.Second) ticker := time.NewTicker(time.Duration(pollInterval) * time.Second) defer ticker.Stop() for { select { case <-ctx.Done(): return GeneratedImage{}, ctx.Err() case <-ticker.C: st, err := queryStatus(ctx, taskID) if err != nil { return GeneratedImage{}, fmt.Errorf("query status: %w", err) } switch st.Data.Status { case "completed": log.Printf("[inference] task %s completed, downloading", taskID) data, err := downloadResult(ctx, st.Data.ResultImageURL) if err != nil { return GeneratedImage{}, fmt.Errorf("download: %w", err) } return GeneratedImage{Data: data, Format: "png"}, nil case "failed": errMsg := st.Data.ErrorMessage if errMsg == "" { errMsg = "unknown error" } return GeneratedImage{}, fmt.Errorf("generation failed: %s", errMsg) default: log.Printf("[inference] task %s status=%s progress=%d", taskID, st.Data.Status, st.Data.Progress) } if time.Now().After(deadline) { return GeneratedImage{}, fmt.Errorf("poll timeout after %ds", maxWait) } } } } // queryStatus 查询任务状态。 func queryStatus(ctx context.Context, taskID string) (*genStatusResp, error) { body, err := json.Marshal(genStatusReq{TaskID: taskID}) if err != nil { return nil, fmt.Errorf("marshal: %w", err) } url := strings.TrimRight(imgCfg.BaseURL, "/") + "/gpt-image2/status" req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(body)) if err != nil { return nil, fmt.Errorf("create request: %w", err) } req.Header.Set("Content-Type", "application/json") req.Header.Set("Authorization", "Bearer "+imgCfg.APIKey) resp, err := http.DefaultClient.Do(req) if err != nil { return nil, fmt.Errorf("send: %w", err) } defer resp.Body.Close() var sr genStatusResp if err := json.NewDecoder(resp.Body).Decode(&sr); err != nil { return nil, fmt.Errorf("decode: %w", err) } if sr.Code != 200 { return nil, fmt.Errorf("status query failed: %s", sr.Message) } return &sr, nil } // downloadResult 下载生成的图片。 func downloadResult(ctx context.Context, imageURL string) ([]byte, error) { req, err := http.NewRequestWithContext(ctx, http.MethodGet, imageURL, nil) if err != nil { return nil, fmt.Errorf("create download request: %w", err) } resp, err := http.DefaultClient.Do(req) if err != nil { return nil, fmt.Errorf("download: %w", err) } defer resp.Body.Close() if resp.StatusCode != http.StatusOK { return nil, fmt.Errorf("download status %d", resp.StatusCode) } return io.ReadAll(resp.Body) } // ======================== 质检 ======================== // QualityChecker 质检函数,可替换用于测试。 var QualityChecker = defaultCheckQuality // CheckQuality 调用当前 QualityChecker。 func CheckQuality(ctx context.Context, images []GeneratedImage, style map[string]string) (bool, string, error) { return QualityChecker(ctx, images, style) } func defaultCheckQuality(ctx context.Context, images []GeneratedImage, style map[string]string) (bool, string, error) { return true, "", nil } func NewCountedQualityChecker(passOnRetry int) func(context.Context, []GeneratedImage, map[string]string) (bool, string, error) { var callCount int return func(_ context.Context, _ []GeneratedImage, _ map[string]string) (bool, string, error) { callCount++ if callCount >= passOnRetry { return true, "", nil } return false, fmt.Sprintf("style inconsistent (attempt %d)", callCount), nil } } func AlwaysFailQualityChecker() func(context.Context, []GeneratedImage, map[string]string) (bool, string, error) { return func(_ context.Context, _ []GeneratedImage, _ map[string]string) (bool, string, error) { return false, "style inconsistent", nil } } // ======================== Mock 回退 ======================== func generateMockImages(size, count int) ([]GeneratedImage, error) { images := make([]GeneratedImage, count) for i := 0; i < count; i++ { data, err := generateMockImage(size, i) if err != nil { return nil, fmt.Errorf("generate mock image %d: %w", i, err) } images[i] = GeneratedImage{Data: data, Width: size, Height: size, Format: "png"} } return images, nil } func generateMockImage(size int, seed int) ([]byte, error) { img := image.NewRGBA(image.Rect(0, 0, size, size)) r := uint8((seed*47 + 13) % 256) g := uint8((seed*83 + 37) % 256) b := uint8((seed*61 + 71) % 256) for y := 0; y < size; y++ { for x := 0; x < size; x++ { img.Set(x, y, color.RGBA{R: r, G: g, B: b, A: 255}) } } var buf bytes.Buffer if err := png.Encode(&buf, img); err != nil { return nil, err } return buf.Bytes(), nil } func generateRandomBytes(n int) ([]byte, error) { b := make([]byte, n) _, err := rand.Read(b) return b, err }