mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-09-27 12:53:35 +00:00
把"能不能并发"从内核硬编码名单改成**工具自己的声明项**,形态照 SDK 的
NoMemory 走。
## ★ 起因:提示词在跟内核不一致
阶段 2.5 写进提示词的「内核默认并行执行」当时是**假的**:toolParallelSafe
只查 stageHost 与 io 两个来源,而全仓 ParallelSafe:true 的生产代码数量
是 **0**。于是除碰巧只发一个工具外,每一批都整批串行回退,而提示词正教
模型把多个查询放同一轮。**内核行为与提示词不一致 = 对模型说谎。**
并发面:0 → 37 个工具(18 插件 ParallelSafe + 19 插件 Serial + 9 内置只读)。
## 声明形态(照 SDK,不自创)
### 插件:结构体字段
s.RegisterTool("config_get", sdk.ToolDef{
Name: ..., Description: ...,
Parameters: map[string]interface{}{...},
// 已核实只读:…
ParallelSafe: true, ← 插在 Parameters 之后、handler 之前
}, p.handleGet(s))
位置与 SDK 的 NoMemory/ContextPolicy/RecallPolicy 一致:Name 在首位,
声明项在末尾,不打散 gofmt 对齐。
### 新增 SDK 声明项:ToolDef.Serial
ParallelSafe 的**反向**标记,判据优先级高于 ParallelSafe。
为什么需要:ParallelSafe 零值 false 已表达"安全",插件无法区分"我没想过"
与"我确认过必须串行"。没有这个区分,工具作者只能靠命名约定传递意图。
内核已消费它(io.ToolDef 同步加字段对齐),并有判据守"Serial 胜出"。
### 内置工具:toolDef 的 toolParallel 选项
内置工具以裸 schema map 下发,没有 ToolDef 结构,所以用变参选项:
toolDef(名字, 描述, 属性) // 默认串行
toolDef(名字, 描述, 属性, "toolParallel") // 已核实只读,可并发
读工具表的老调用点一行不用动,声明就写在工具定义那一行。
## ★ 走过的弯路(都留了判据)
1. **硬编码白名单**:先在 toolParallelSafe 里查一张
builtinParallelSafeTools map。那把声明从"工具自己"搬回了内核 ——
工具改名/新增不会自动跟着变,得靠一条 grep 源码的判据才能发现漂移,
而判据一改就忘。已删,改为从定义读。
2. **判据前提错(同一个坑踩了两次)**:拿裸 &Agent{} 的 buildToolDefs 输出
当"实际可见工具",但这 9 个内置工具全在条件分支里(a.knowledge != nil /
a.social != nil / a.parentID != ""…),裸 Agent 一个都不产出 ⇒ 全部误报
"声明形同虚设"。第一次叫它"幽灵条目",没认出是同一个坑。
3. **注释模仿真实签名污染判据**:toolParallel 的用法注释写着
`toolDef("knowledge_search", ...)`,判据按文本匹配先撞上注释。
4. **buildToolDefs 的 nil 不一致**:开头判了 a.io != nil,末尾却无条件
a.io.ListChannels()。任何无 IO 的 Agent 调它都 panic —— 而 panic 报在
io 包里,根因在 tooldefs.go。已补。
5. **插入脚本用正则找"最后一个顶层字段"**:被嵌套 map 里的同形文本骗到,
823 处错误重排把文件改坏。改用括号深度 + 记录进入深度 3 的行号
(空 properties 会让深度在同一行进出平衡,只判 depth==2 不够)。
工具在 SDK 仓 tools/annotate_parallel/,复用时用绝对路径。
## 提示词措辞同步修正
「默认并行执行」→「尽量并发执行,但这是**逐工具判断**的」,并教模型
**把查询类放同一轮、写操作单独发一轮**(写和查混在一批,整批都串行)。
## 判据
- TestSerialOverridesParallelSafe Serial 优先于 ParallelSafe
- TestToolParallelDeclarationsAudit 并发面不许再归零
- TestNoToolDeclaresBothParallelAndSerial 两者同标即谎话
- TestBuiltinParallelDeclaredWhereDefined 声明写在定义处、且内核真读到
- TestStoreListIgnoresForeignJSON 压测抓到的 List() 缺陷
301 lines
8.8 KiB
Go
301 lines
8.8 KiB
Go
package aiimage
|
||
|
||
import (
|
||
"bytes"
|
||
"encoding/json"
|
||
"fmt"
|
||
"io"
|
||
"net/http"
|
||
"os"
|
||
"path/filepath"
|
||
"strings"
|
||
"time"
|
||
|
||
"gitcode.com/JianFeeeee/HomeAgent/internal/plugin"
|
||
sdk "gitcode.com/JianFeeeee/HomeAgent/internal/sdk"
|
||
)
|
||
|
||
func init() {
|
||
plugin.RegisterPluginMeta("ai_image", "AI 生图", "AI Image")
|
||
plugin.RegisterFactory("ai_image", func(name string, config map[string]interface{}) (sdk.Plugin, error) {
|
||
return New(name), nil
|
||
})
|
||
}
|
||
|
||
type Plugin struct {
|
||
name string
|
||
client *http.Client
|
||
apiKey string
|
||
baseURL string
|
||
model string
|
||
size string
|
||
dataDir string
|
||
}
|
||
|
||
func New(name string) *Plugin {
|
||
return &Plugin{
|
||
name: name,
|
||
client: &http.Client{Timeout: 120 * time.Second},
|
||
}
|
||
}
|
||
|
||
func (p *Plugin) Name() string { return p.name }
|
||
|
||
func getSettingString(s sdk.SettingsAPI, key, def string) string {
|
||
if s == nil {
|
||
return def
|
||
}
|
||
v, err := s.Get(key)
|
||
if err != nil || v == nil {
|
||
return def
|
||
}
|
||
if sv, ok := v.(string); ok && sv != "" {
|
||
return sv
|
||
}
|
||
return def
|
||
}
|
||
|
||
func (p *Plugin) Start(s *sdk.PluginSDK) error {
|
||
s.SetAutoRestart(true)
|
||
|
||
// 插件专属数据目录(SDK DataDir API,内核保证存在)
|
||
if dd := s.Settings().DataDir(); dd != "" {
|
||
p.dataDir = dd
|
||
}
|
||
if p.dataDir != "" {
|
||
os.MkdirAll(p.dataDir, 0755)
|
||
}
|
||
|
||
s.Settings().RegisterDef(sdk.ConfigDef{
|
||
Key: "base_url", Default: "http://127.0.0.1:8081/v1", Type: "string",
|
||
DisplayName: "Base URL", Description: "OpenAI 兼容生图网关地址(默认本机 llmsproxy)",
|
||
Category: "ai_image",
|
||
})
|
||
s.Settings().RegisterDef(sdk.ConfigDef{
|
||
Key: "api_key", Default: "", Type: "password",
|
||
DisplayName: "API Key", Description: "生图网关 API Key(默认使用内核 LLM key,留空则取 core.llm.api_key)",
|
||
Category: "ai_image",
|
||
})
|
||
s.Settings().RegisterDef(sdk.ConfigDef{
|
||
Key: "model", Default: "", Type: "string",
|
||
DisplayName: "Model", Description: "生图模型 ID,留空由网关 AUTO 决定(如 flux-1)",
|
||
Category: "ai_image",
|
||
})
|
||
s.Settings().RegisterDef(sdk.ConfigDef{
|
||
Key: "size", Default: "1024x1024", Type: "string",
|
||
DisplayName: "Size", Description: "默认图片尺寸,如 1024x1024",
|
||
Category: "ai_image",
|
||
})
|
||
|
||
p.apiKey = getSettingString(s.Settings(), "api_key", "")
|
||
p.baseURL = strings.TrimRight(getSettingString(s.Settings(), "base_url", "http://127.0.0.1:8081/v1"), "/")
|
||
p.model = getSettingString(s.Settings(), "model", "")
|
||
p.size = getSettingString(s.Settings(), "size", "1024x1024")
|
||
|
||
if p.apiKey == "" {
|
||
if v, _ := s.Settings().GetCore("llm.api_key"); v != nil {
|
||
if sv, ok := v.(string); ok && sv != "" {
|
||
p.apiKey = sv
|
||
}
|
||
}
|
||
}
|
||
|
||
s.RegisterTool("ai_image_generate", sdk.ToolDef{
|
||
Name: "ai_image_generate", Description: "Generate image from text prompt using AI. Returns image URL.",
|
||
Parameters: map[string]interface{}{
|
||
"type": "object",
|
||
"properties": map[string]interface{}{
|
||
"prompt": map[string]interface{}{"type": "string", "description": "Text description of the image to generate"},
|
||
"size": map[string]interface{}{"type": "string", "description": "Image size (1024x1024, etc.), default from config"},
|
||
"model": map[string]interface{}{"type": "string", "description": "Model override (e.g. flux-1)"},
|
||
"n": map[string]interface{}{"type": "integer", "description": "Number of images to generate (1-10), default 1"},
|
||
},
|
||
"required": []string{"prompt"},
|
||
},
|
||
// 外部调用,插件内无共享可变状态
|
||
ParallelSafe: true,
|
||
}, p.handleGenerate)
|
||
|
||
return nil
|
||
}
|
||
|
||
func (p *Plugin) Stop() error { return nil }
|
||
|
||
type genRequest struct {
|
||
Model string `json:"model"`
|
||
Prompt string `json:"prompt"`
|
||
N int `json:"n"`
|
||
Size string `json:"size"`
|
||
ResponseFormat string `json:"response_format"`
|
||
}
|
||
|
||
type genResp struct {
|
||
Created int64 `json:"created"`
|
||
Data []struct {
|
||
URL string `json:"url"`
|
||
B64JSON string `json:"b64_json"`
|
||
RevisedPrompt string `json:"revised_prompt"`
|
||
} `json:"data"`
|
||
Error *struct {
|
||
Message string `json:"message"`
|
||
Type string `json:"type"`
|
||
} `json:"error"`
|
||
}
|
||
|
||
func (p *Plugin) handleGenerate(args map[string]interface{}) (interface{}, error) {
|
||
prompt, _ := args["prompt"].(string)
|
||
if strings.TrimSpace(prompt) == "" {
|
||
return map[string]interface{}{"isError": true, "content": "prompt is required"}, nil
|
||
}
|
||
|
||
key := p.apiKey
|
||
if key == "" {
|
||
return map[string]interface{}{"isError": true, "content": "生图 API key 未配置(plugin.ai_image.api_key 或 core.llm.api_key)"}, nil
|
||
}
|
||
|
||
model := p.model
|
||
if m, ok := args["model"].(string); ok && m != "" {
|
||
model = m
|
||
}
|
||
size := p.size
|
||
if sz, ok := args["size"].(string); ok && sz != "" {
|
||
size = sz
|
||
}
|
||
n := 1
|
||
if nv, ok := args["n"].(float64); ok {
|
||
n = int(nv)
|
||
if n < 1 {
|
||
n = 1
|
||
}
|
||
if n > 10 {
|
||
n = 10
|
||
}
|
||
}
|
||
|
||
body := genRequest{Model: model, Prompt: prompt, N: n, Size: size, ResponseFormat: "url"}
|
||
raw, _ := json.Marshal(body)
|
||
|
||
url := p.baseURL + "/images/generations"
|
||
req, err := http.NewRequest(http.MethodPost, url, bytes.NewReader(raw))
|
||
if err != nil {
|
||
return map[string]interface{}{"isError": true, "content": "构造请求失败: " + err.Error()}, nil
|
||
}
|
||
req.Header.Set("Content-Type", "application/json")
|
||
req.Header.Set("Authorization", "Bearer "+key)
|
||
|
||
resp, err := p.client.Do(req)
|
||
if err != nil {
|
||
return map[string]interface{}{"isError": true, "content": "生图请求失败: " + err.Error()}, nil
|
||
}
|
||
defer resp.Body.Close()
|
||
|
||
respBody, _ := io.ReadAll(resp.Body)
|
||
if resp.StatusCode != 200 {
|
||
return map[string]interface{}{"isError": true, "content": fmt.Sprintf("生图 API error (status %d): %s", resp.StatusCode, string(respBody))}, nil
|
||
}
|
||
|
||
var result genResp
|
||
if err := json.Unmarshal(respBody, &result); err != nil {
|
||
return map[string]interface{}{"isError": true, "content": "解析生图响应失败: " + err.Error()}, nil
|
||
}
|
||
if result.Error != nil {
|
||
return map[string]interface{}{"isError": true, "content": "生图 API error: " + result.Error.Message}, nil
|
||
}
|
||
if len(result.Data) == 0 {
|
||
return map[string]interface{}{"isError": true, "content": "no images returned"}, nil
|
||
}
|
||
|
||
var urls []string
|
||
for _, d := range result.Data {
|
||
imgURL := d.URL
|
||
if imgURL == "" && d.B64JSON != "" {
|
||
imgURL = "data:image/png;base64," + d.B64JSON
|
||
}
|
||
if imgURL != "" {
|
||
urls = append(urls, imgURL)
|
||
}
|
||
}
|
||
if len(urls) == 0 {
|
||
return map[string]interface{}{"isError": true, "content": "生图响应中没有可用图片"}, nil
|
||
}
|
||
|
||
// 下载到插件数据目录,返回本地文件路径(而非临时 S3 URL):
|
||
// - S3 临时 URL 约 1 小时过期,且对无浏览器 UA 客户端拒绝访问(agent 裸 curl 验证必败)
|
||
// - 本地路径经 webui /files/ 永久下发,支持 output_send type=image
|
||
var localPaths []string
|
||
var dlErrs []string
|
||
for i, u := range urls {
|
||
if strings.HasPrefix(u, "data:") {
|
||
continue // base64 内联图不落盘
|
||
}
|
||
path, err := p.downloadImage(u, fmt.Sprintf("ai_%d_%d", time.Now().UnixNano(), i))
|
||
if err != nil {
|
||
dlErrs = append(dlErrs, fmt.Sprintf("第%d张保存失败: %v", i+1, err))
|
||
continue
|
||
}
|
||
localPaths = append(localPaths, path)
|
||
}
|
||
|
||
content := fmt.Sprintf("Generated %d image(s) with model %s", len(urls), model)
|
||
if len(localPaths) > 0 {
|
||
content += "\n本地文件:\n" + strings.Join(localPaths, "\n")
|
||
content += "\n\n图片已保存到本地(不会过期)。如需展示请用 output_send__webui(payload=本地路径, type=image)。"
|
||
}
|
||
if len(dlErrs) > 0 {
|
||
content += "\n\n" + strings.Join(dlErrs, "\n")
|
||
}
|
||
if len(localPaths) < len(urls) {
|
||
content += "\n原始 URL(1小时内有效):\n" + strings.Join(urls, "\n")
|
||
}
|
||
|
||
return map[string]interface{}{
|
||
"content": content,
|
||
"images": urls,
|
||
"local_paths": localPaths,
|
||
"prompt": prompt,
|
||
"model": model,
|
||
}, nil
|
||
}
|
||
|
||
// downloadImage 把生图返回的临时 URL 下载为本地文件,返回路径。
|
||
// 带浏览器 UA 规避图床对无 UA 客户端的拦截。
|
||
func (p *Plugin) downloadImage(imgURL, baseName string) (string, error) {
|
||
if p.dataDir == "" {
|
||
return "", fmt.Errorf("data dir unavailable")
|
||
}
|
||
dl := &http.Client{Timeout: 60 * time.Second}
|
||
req, err := http.NewRequest(http.MethodGet, imgURL, nil)
|
||
if err != nil {
|
||
return "", err
|
||
}
|
||
req.Header.Set("User-Agent", "Mozilla/5.0 (compatible; HomeAgent/1.0)")
|
||
resp, err := dl.Do(req)
|
||
if err != nil {
|
||
return "", err
|
||
}
|
||
defer resp.Body.Close()
|
||
if resp.StatusCode != http.StatusOK {
|
||
b, _ := io.ReadAll(resp.Body)
|
||
if len(b) > 200 {
|
||
b = b[:200]
|
||
}
|
||
return "", fmt.Errorf("HTTP %d: %s", resp.StatusCode, strings.TrimSpace(string(b)))
|
||
}
|
||
data, err := io.ReadAll(resp.Body)
|
||
if err != nil {
|
||
return "", err
|
||
}
|
||
ext := ".png"
|
||
switch ct := resp.Header.Get("Content-Type"); {
|
||
case strings.Contains(ct, "jpeg"), strings.Contains(ct, "jpg"):
|
||
ext = ".jpg"
|
||
case strings.Contains(ct, "webp"):
|
||
ext = ".webp"
|
||
}
|
||
path := filepath.Join(p.dataDir, baseName+ext)
|
||
if err := os.WriteFile(path, data, 0644); err != nil {
|
||
return "", err
|
||
}
|
||
return path, nil
|
||
}
|