Files
2026-08-27 14:11:50 +08:00

250 lines
7.6 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package common
import (
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"net/url"
"strings"
"time"
"github.com/gogf/gf/v2/errors/gerror"
"github.com/gogf/gf/v2/frame/g"
)
// ImageGenProvider 图像生成 provider 抽象:管理端「AI 生成图片」按 config.yml imageGen.provider 选择实现。
type ImageGenProvider interface {
// Generate 生成一张图片,返回图片字节;prompt 由调用方保证不含位置描述(项目提示词规范)。
Generate(ctx context.Context, prompt, size string) ([]byte, error)
}
// ImageGen 当前配置的图像生成 provider 单例(provider 必填项缺失返回 nil,调用方判 CodeImageGenNotConfigured)。
func ImageGen(ctx context.Context) ImageGenProvider {
switch g.Cfg().MustGet(ctx, "imageGen.provider", "dashscope").String() {
case "localai":
if g.Cfg().MustGet(ctx, "imageGen.baseUrl").String() == "" {
return nil
}
return &localAiImageGen{
baseUrl: g.Cfg().MustGet(ctx, "imageGen.baseUrl").String(),
model: g.Cfg().MustGet(ctx, "imageGen.model", "qwen-image").String(),
}
case "dashscope":
if g.Cfg().MustGet(ctx, "imageGen.apiKey").String() == "" {
return nil
}
return &dashScopeImageGen{apiKey: g.Cfg().MustGet(ctx, "imageGen.apiKey").String()}
default:
return nil
}
}
// dashScopeImageGen 通义万相(DashScope)实现:提交异步任务 → 轮询 task 状态 → 下载产物图片。
// 文档:https://help.aliyun.com/zh/model-studio/text-to-image-api-reference
type dashScopeImageGen struct {
apiKey string
}
const (
dashScopeSynthUrl = "https://dashscope.aliyuncs.com/api/v1/services/aigc/text2image/image-synthesis"
dashScopeTaskUrlFmt = "https://dashscope.aliyuncs.com/api/v1/tasks/%s"
)
func (p *dashScopeImageGen) Generate(ctx context.Context, prompt, size string) ([]byte, error) {
model := g.Cfg().MustGet(ctx, "imageGen.model", "qwen-image-3.0").String()
// 前端尺寸格式 1152x2048 → API 规格 1152*2048
apiSize := strings.ReplaceAll(size, "x", "*")
body, err := json.Marshal(map[string]any{
"model": model,
"input": map[string]string{"prompt": prompt, "size": apiSize},
"parameters": map[string]any{"n": 1, "watermark": false},
})
if err != nil {
return nil, err
}
resp, err := p.post(ctx, dashScopeSynthUrl, body)
if err != nil {
return nil, err
}
taskId, ok := resp["task_id"].(string)
if !ok || taskId == "" {
return nil, fmt.Errorf("DashScope 提交失败: %v", resp)
}
// 轮询任务结果:异步生成通常 10~60s,上限 120s(与设计一致:超时 2min/张)
deadline := time.Now().Add(120 * time.Second)
for {
if ctx.Err() != nil {
return nil, ctx.Err()
}
if time.Now().After(deadline) {
return nil, fmt.Errorf("DashScope 生成超时(120s)")
}
task, err := p.post(ctx, fmt.Sprintf(dashScopeTaskUrlFmt, taskId), nil)
if err != nil {
return nil, err
}
status, _ := task["task_status"].(string)
switch status {
case "SUCCEEDED":
results, _ := task["results"].([]any)
if len(results) == 0 {
return nil, fmt.Errorf("DashScope 成功但无产物图片")
}
url, _ := results[0].(map[string]any)["url"].(string)
if url == "" {
return nil, fmt.Errorf("DashScope 成功但无图片 URL")
}
return p.download(ctx, url)
case "FAILED", "CANCELED":
msg, _ := task["message"].(string)
return nil, fmt.Errorf("DashScope 生成失败: %s", msg)
}
select {
case <-time.After(2 * time.Second):
case <-ctx.Done():
return nil, ctx.Err()
}
}
}
// post 调用 DashScope HTTP 接口并解析统一 JSON(无 body 时为空 GET,轮询任务用)
func (p *dashScopeImageGen) post(ctx context.Context, url string, body []byte) (map[string]any, error) {
method := http.MethodPost
var rd io.Reader
if body == nil {
method = http.MethodGet
} else {
rd = strings.NewReader(string(body))
}
req, err := http.NewRequestWithContext(ctx, method, url, rd)
if err != nil {
return nil, err
}
req.Header.Set("Authorization", "Bearer "+p.apiKey)
if body != nil {
req.Header.Set("Content-Type", "application/json")
}
resp, err := http.DefaultClient.Do(req)
if err != nil {
return nil, err
}
defer func() { _ = resp.Body.Close() }()
raw, err := io.ReadAll(io.LimitReader(resp.Body, 8<<20))
if err != nil {
return nil, err
}
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("DashScope HTTP %d: %s", resp.StatusCode, truncateStr(string(raw), 200))
}
var m map[string]any
if err := json.Unmarshal(raw, &m); err != nil {
return nil, err
}
return m, nil
}
// download 下载生成产物图片(存于阿里云 OSS,无需鉴权)
func (p *dashScopeImageGen) download(ctx context.Context, url string) ([]byte, error) {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return nil, err
}
resp, err := http.DefaultClient.Do(req)
if err != nil {
return nil, err
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("下载生成图片失败: HTTP %d", resp.StatusCode)
}
return io.ReadAll(io.LimitReader(resp.Body, 16<<20))
}
func truncateStr(s string, n int) string {
if len(s) <= n {
return s
}
return s[:n] + "..."
}
// localAiImageGen local-aiOpenAI 兼容 /v1/images/generations)实现:同步生成 → 下载产物图片。
// 训练机 local-ai 已加载 qwen-image 文生图;返回 url 为容器内 localhost 地址,需替换为配置 baseUrl 下载。
type localAiImageGen struct {
baseUrl string
model string
}
func (p *localAiImageGen) Generate(ctx context.Context, prompt, size string) ([]byte, error) {
body, err := json.Marshal(map[string]any{
"model": p.model,
"prompt": prompt,
"size": size,
"n": 1,
})
if err != nil {
return nil, err
}
url := strings.TrimSuffix(p.baseUrl, "/") + "/v1/images/generations"
req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, strings.NewReader(string(body)))
if err != nil {
return nil, err
}
req.Header.Set("Content-Type", "application/json")
resp, err := http.DefaultClient.Do(req)
if err != nil {
return nil, gerror.Wrap(err, "local-ai 生成请求失败")
}
defer func() { _ = resp.Body.Close() }()
raw, err := io.ReadAll(io.LimitReader(resp.Body, 8<<20))
if err != nil {
return nil, err
}
if resp.StatusCode != http.StatusOK {
msg := truncateStr(string(raw), 200)
if msg == "" {
msg = resp.Status
}
return nil, fmt.Errorf("local-ai 生成失败: HTTP %d: %s", resp.StatusCode, msg)
}
var m struct {
Data []struct {
Url string `json:"url"`
} `json:"data"`
}
if err := json.Unmarshal(raw, &m); err != nil {
return nil, fmt.Errorf("local-ai 响应解析失败: %v", err)
}
if len(m.Data) == 0 || m.Data[0].Url == "" {
return nil, fmt.Errorf("local-ai 成功但无图片 URL")
}
return p.download(ctx, m.Data[0].Url)
}
// download 下载生成产物图片:url 为容器内地址(localhost),替换 scheme/host 为配置 baseUrl 后下载
func (p *localAiImageGen) download(ctx context.Context, imageUrl string) ([]byte, error) {
base, err := url.Parse(p.baseUrl)
if err != nil {
return nil, err
}
u, err := url.Parse(imageUrl)
if err != nil {
return nil, err
}
u.Scheme, u.Host = base.Scheme, base.Host
req, err := http.NewRequestWithContext(ctx, http.MethodGet, u.String(), nil)
if err != nil {
return nil, err
}
resp, err := http.DefaultClient.Do(req)
if err != nil {
return nil, gerror.Wrap(err, "下载生成图片失败")
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("下载生成图片失败: HTTP %d", resp.StatusCode)
}
return io.ReadAll(io.LimitReader(resp.Body, 16<<20))
}