1
This commit is contained in:
@@ -6,9 +6,11 @@ import (
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gogf/gf/v2/errors/gerror"
|
||||
"github.com/gogf/gf/v2/frame/g"
|
||||
)
|
||||
|
||||
@@ -18,13 +20,21 @@ type ImageGenProvider interface {
|
||||
Generate(ctx context.Context, prompt, size string) ([]byte, error)
|
||||
}
|
||||
|
||||
// ImageGen 当前配置的图像生成 provider 单例(未配置返回 nil,调用方判 CodeImageGenNotConfigured)。
|
||||
// ImageGen 当前配置的图像生成 provider 单例(provider 必填项缺失返回 nil,调用方判 CodeImageGenNotConfigured)。
|
||||
func ImageGen(ctx context.Context) ImageGenProvider {
|
||||
if g.Cfg().MustGet(ctx, "imageGen.apiKey").String() == "" {
|
||||
return nil
|
||||
}
|
||||
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
|
||||
@@ -158,3 +168,82 @@ func truncateStr(s string, n int) string {
|
||||
}
|
||||
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))
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user