250 lines
7.6 KiB
Go
250 lines
7.6 KiB
Go
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-ai(OpenAI 兼容 /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))
|
||
}
|