diff --git a/DESIGN.md b/DESIGN.md new file mode 100644 index 0000000..93e836a --- /dev/null +++ b/DESIGN.md @@ -0,0 +1,257 @@ +# HomeAgent 架构设计 v2 + +## 一、核心理念 + +24 小时陪伴用户、随时待命的智能管家。**单会话·单 Agent**(可分身/调用其他 Agent),身份不漂移。 + +### 设计原则 +- 所有输入走 IO 抽象层(中断模式),不直调 agent 方法 +- DeepSeek v4 flash 为默认 LLM,thinking 模式关闭 +- 人格固定(personal.md),记忆分层管理防止性格突变 +- 知识独立于记忆,agent 主动学习 + +--- + +## 二、人格内核(Personality Core) + +``` +personal.md ──→ 固定注入 system prompt +``` + +- 纯文本 Markdown 文件,定义 agent 的人格、行为准则 +- 加载一次永不改变(不随对话漂移) +- 放在 prompt 最前面,优先级最高 + +--- + +## 三、记忆体系(Memory System) + +三层分级架构,自顶向下逐渐持久化、抽象化: + +``` + 用户输入 + │ + ▼ + ┌─────────────────────┐ + │ Layer 1: 上下 文 │ RelevanceContext + │ 基于相关性保留 │ 30 条活跃,TF-IDF 排序 + │ 最不相关的→文档记忆 │ + └────────┬────────────┘ + │ 相关性裁剪(每次响应后) + ▼ + ┌─────────────────────┐ + │ Layer 2: 文档记忆 │ document.Store + │ 全文 + TF-IDF 向量 │ JSON 持久化 + │ 冷文档→图数据库 │ 72h 未访问→Graph + └────────┬────────────┘ + │ 心跳蒸馏(30min) + ▼ + ┌─────────────────────┐ + │ Layer 3: 图数据库 │ memory.GraphDB (SQLite) + │ 实体-关系三元组 │ 向量索引(实体名) + │ 定期重整+同义合并 │ bigram Jaccard + └─────────────────────┘ +``` + +### Layer 1: 上下文(RelevanceContext) + +**不再使用固定条数裁剪。改为:** + +1. 每次用户输入后,计算**模型输出**与每条上下文的 TF-IDF 余弦相似度 +2. 按相关性从高到低排序,保留 topK(默认 30) +3. 最不相关的上下文 → 归档到文档记忆(Layer 2),保留全文+向量 +4. 活跃上下文中只保留最近且相关性高的内容 + +**相关文件:** `internal/agent/core/context.go` + +### Layer 2: 文档记忆(Document Store) + +**作用:** +- 存储被上下文裁剪下来的事件摘要、agent 主动提交的文档、图记忆同步的索引 +- 每个文档包含:摘要、全文、标签、实体、来源、访问计数、最后访问时间 +- TF-IDF 向量索引(字符 bigram),支持向量相似度搜索 + +**冷文档归化:** +- 心跳检测(每 30min):找出 72h 未访问且访问次数 ≤ 2 的文档 +- 转为图数据库三元组(文档→包含内容、提及实体、标签、来源) +- 清理已归化的冷文档 + +**相关文件:** `internal/memory/document/document.go` + +### Layer 3: 图数据库(Graph Memory) + +**存储:** SQLite 三元组(实体-关系-实体),每个关系带置信度、会话 ID、日期桶 + +**注入方式(区分于文档注入):** +- 用户消息到达时,先对实体名做**向量相似度搜索**(TF-IDF) +- 找到相关实体名 → 查询图数据库 +- 只注入**节点索引**(实体名+类型+提及次数+关系类型)到 prompt,不注入全文 +- 需要更多细节时,agent 调用 `memory_recall` 工具查询 + +**定期重整(心跳触发,每 30min):** +1. 同步实体名到向量索引(`Indexer.Sync()`) +2. 文档记忆向量索引重建 +3. 冷文档→图归化 +4. 实体同义合并(bigram Jaccard > 0.5) + +**相关文件:** `internal/memory/graph.go`, `internal/memory/indexer.go` + +### Agent 可用记忆工具 + +| 工具 | 作用 | 操作对象 | +|------|------|----------| +| memory_recall | 检索图记忆 | GraphDB | +| memory_commit | 写入三元组 | GraphDB | +| memory_introspect | 查看记忆统计 | GraphDB | +| doc_query | 向量查询文档记忆 | Document Store | +| doc_commit | 提交文档 | Document Store | + +--- + +## 四、知识体系(Knowledge System) + +独立于记忆系统,用于 agent 学习知识: + +``` +knowledge/ + smart_home/ + content.md ← 原始知识文件 + cooking/ + content.md + ... + │ + ▼ + TF-IDF 向量索引 ← 知识目录扫描时自动构建 + │ + ▼ + knowledge_search(query) → 返回相关内容 +``` + +**知识来源:** +1. **agent 主动学习**:调用 `knowledge_create` 工具,生成知识→写入目录+向量化 +2. **用户上传**:HTTP 文件上传端点 → 写入目录+向量化 +3. **预置知识**:`knowledge/` 目录下的 content.md + +**相关文件:** `internal/knowledge/knowledge.go` + +### Agent 可用知识工具 + +| 工具 | 作用 | +|------|------| +| knowledge_search | 向量搜索知识库 | +| knowledge_list | 列出知识分类 | +| knowledge_create | agent 主动创建知识 | + +--- + +## 五、数据流总览 + +``` +用户消息 + │ + ▼ +IO 输入中断 (channel.go) + │ + ▼ +eventLoop → processTextInput + │ + ├─ 1. 追加上下文 (RelevanceContext.Append) + │ └─ 计算向量,缓存 + │ + ├─ 2. 构建 prompt: + │ ├─ 人格设定 (personal.md) + │ ├─ 图记忆索引 (Indexer.BuildContext → 向量搜索实体名 → Recall → 摘要注入) + │ ├─ 文档记忆摘要 (DocStore.Query → 注入前3条摘要) + │ ├─ 工作记忆上下文 (RelevanceContext.Format) + │ ├─ 技能注入 (skills) + │ └─ 工具说明 (indexer + doc + knowledge tools) + │ + ├─ 3. 工具循环 (process) + │ ├─ 调用 LLM (DeepSeek v4 flash) + │ ├─ 解析 tool_calls + │ ├─ 执行工具 (memory_*/knowledge_*/doc_*/设备工具) + │ └─ 返回结果,循环直到无 tool_calls + │ + ├─ 4. 追加响应到上下文 + │ + ├─ 5. 相关性裁剪 (RelevanceContext.Prune) + │ └─ 最不相关的 → 文档记忆归档 + │ + ├─ 6. IO 输出响应 + │ + └─ 7. 记忆候选事件 + └─ → TextMemory (JSONL 持久化) + └─ → Distiller → GraphMemory (三元组萃取) + +心跳线程 (每 30min): + ├─ distillContext: 安全裁剪兜底 + ├─ syncGraphToDocs: 图→文档索引同步 + ├─ reorgGraph: + │ ├─ Indexer.Sync(): 实体名→向量索引 + │ ├─ DocStore.Reindex(): 文档向量重建 + │ ├─ 冷文档→图归化 + │ └─ 同义实体合并 + └─ (后续) GraphDB 向量索引生成/重整 +``` + +--- + +## 六、关键文件 + +``` +cmd/homed/main.go — 入口:组装所有子系统 +internal/agent/core/agent.go — Agent 核心:事件循环、工具循环、心跳 +internal/agent/core/context.go — RelevanceContext:基于 TF-IDF 的上下文管理 +internal/agent/personal.go — Personality:personal.md 加载 +internal/agent/api/provider.go — LLM Provider:DeepSeek API 封装 +internal/agent/io/channel.go — IO 抽象层:中断输入、设备注册 +internal/api/handler.go — HTTP API:REST + OpenAI 兼容端点 +internal/knowledge/knowledge.go — 知识系统:目录扫描、向量索引、搜索 +internal/memory/graph.go — 图数据库:SQLite 三元组 CRUD +internal/memory/indexer.go — 图索引器:向量搜索实体名、摘要注入 +internal/memory/vector/store.go — 向量存储:TF-IDF + 倒排索引 + 余弦相似度 +internal/memory/document/doc.go — 文档记忆:存储、查询、冷文档检测 +internal/memory/text/text.go — 文本记忆:JSONL 原始日志持久化 +internal/memory/pipeline/ — 蒸馏器:原始日志→图记忆 +internal/tracker/tracker.go — Change Tracker:overlayfs 文件变更追踪 +internal/plugin/plugin.go — 插件平台:SKILL.md 加载 +internal/supervisor/daemon.go — 守护进程:健康检查、自动 rollback +internal/lua/vm.go — Lua 适配器 VM +config/config.go — 配置加载 +pkg/types/ — 类型定义 +``` + +--- + +## 七、配置示例 + +```yaml +daemon: + listen_addr: ":8080" + data_dir: "/var/lib/homeagent" + heartbeat_interval: 15s + check_interval: 30s + +llm: + provider: "openai" + model: "deepseek-v4-flash" + base_url: "https://api.deepseek.com" + api_key: "${DEEPSEEK_API_KEY}" + temperature: 0.7 + max_tokens: 4096 +``` + +--- + +## 八、对比原设计 + +| 维度 | 旧设计 | 新设计 | +|------|--------|--------| +| 上下文裁剪 | 固定 20 条 FIFO | TF-IDF 相关性排序,保留最相关 | +| 上下文→文档 | 蒸馏器每周期 flush | 每次响应后按相关性裁剪归档 | +| 图记忆注入 | 关键词 LIKE 查询 | 向量搜索实体名,只注索引 | +| 文档→图 | 无 | 72h 冷文档自动归化 | +| 图重整 | 无 | 心跳:向量同步+同义合并 | +| 知识系统 | 无 | `knowledge/` 目录+向量索引+工具 | +| 人格 | 无 | `personal.md` 固定注入 | +| 向量引擎 | 无 | 自研 TF-IDF + 倒排索引 | diff --git a/Makefile b/Makefile new file mode 100644 index 0000000..45b2b65 --- /dev/null +++ b/Makefile @@ -0,0 +1,47 @@ +.PHONY: all build clean install test run + +BINARY=homed +GO=go +GOCACHE=/tmp/gocache +GOPATH=$(shell go env GOPATH) +BUILD_DIR=build + +all: build + +build: + @mkdir -p $(BUILD_DIR) + CGO_ENABLED=1 $(GO) build -o $(BUILD_DIR)/$(BINARY) ./cmd/homed/ + @echo "Built: $(BUILD_DIR)/$(BINARY)" + +build-static: + @mkdir -p $(BUILD_DIR) + CGO_ENABLED=1 $(GO) build -tags netgo -ldflags '-extldflags "-static"' -o $(BUILD_DIR)/$(BINARY)-static ./cmd/homed/ + @echo "Built (static): $(BUILD_DIR)/$(BINARY)-static" + +build-linux-arm64: + @mkdir -p $(BUILD_DIR) + GOOS=linux GOARCH=arm64 CGO_ENABLED=1 $(GO) build -o $(BUILD_DIR)/$(BINARY)-arm64 ./cmd/homed/ + @echo "Built (arm64): $(BUILD_DIR)/$(BINARY)-arm64" + +clean: + rm -rf $(BUILD_DIR) $(BINARY) + +install: build + cp $(BUILD_DIR)/$(BINARY) /usr/local/bin/$(BINARY) + mkdir -p /etc/homeagent /var/lib/homeagent + cp config/config.yaml /etc/homeagent/ + cp deploy/homeagent.service /etc/systemd/system/ + systemctl daemon-reload + @echo "Installed. Run: systemctl enable --now homeagent" + +test: + $(GO) test ./... + +run: build + ./$(BUILD_DIR)/$(BINARY) -data /tmp/homeagent + +fmt: + $(GO) fmt ./... + +lint: + $(GO) vet ./... diff --git a/cmd/homed/main.go b/cmd/homed/main.go new file mode 100644 index 0000000..f16af4e --- /dev/null +++ b/cmd/homed/main.go @@ -0,0 +1,303 @@ +package main + +import ( + "flag" + "log" + "net/http" + "os" + "os/signal" + "path/filepath" + "syscall" + "time" + + agentAPI "gitcode.com/JianFeeeee/HomeAgent/internal/agent/api" + agentCore "gitcode.com/JianFeeeee/HomeAgent/internal/agent/core" + agentIO "gitcode.com/JianFeeeee/HomeAgent/internal/agent/io" + agentPkg "gitcode.com/JianFeeeee/HomeAgent/internal/agent" + "gitcode.com/JianFeeeee/HomeAgent/internal/api" + "gitcode.com/JianFeeeee/HomeAgent/config" + "gitcode.com/JianFeeeee/HomeAgent/internal/knowledge" + luapkg "gitcode.com/JianFeeeee/HomeAgent/internal/lua" + "gitcode.com/JianFeeeee/HomeAgent/internal/memory" + "gitcode.com/JianFeeeee/HomeAgent/internal/memory/document" + "gitcode.com/JianFeeeee/HomeAgent/internal/memory/pipeline" + "gitcode.com/JianFeeeee/HomeAgent/internal/memory/text" + "gitcode.com/JianFeeeee/HomeAgent/internal/onebot" + "gitcode.com/JianFeeeee/HomeAgent/internal/plugin" + "gitcode.com/JianFeeeee/HomeAgent/internal/skill" + "gitcode.com/JianFeeeee/HomeAgent/internal/supervisor" + "gitcode.com/JianFeeeee/HomeAgent/internal/tracker" +) + +func main() { + configPath := flag.String("config", config.DefaultConfigPath, "path to config file") + dataDir := flag.String("data", "/var/lib/homeagent", "data directory") + flag.Parse() + + log.SetFlags(log.Ldate | log.Ltime | log.Lshortfile) + log.Printf("[homed] starting HomeAgent v0.1.0") + + cfg, err := config.Load(*configPath) + if err != nil { + log.Fatalf("load config: %v", err) + } + + cfg.Daemon.DataDir = *dataDir + + agentWorkDir := filepath.Join(cfg.Daemon.DataDir, "agentfs") + dirs := []string{ + cfg.Daemon.DataDir, + filepath.Join(cfg.Daemon.DataDir, "snapshots"), + filepath.Join(cfg.Daemon.DataDir, "skills"), + filepath.Join(cfg.Daemon.DataDir, "plugins"), + filepath.Join(cfg.Daemon.DataDir, "changesets"), + filepath.Join(cfg.Daemon.DataDir, "memory"), + filepath.Join(cfg.Daemon.DataDir, "memory", "raw"), + filepath.Join(cfg.Daemon.DataDir, "adapters"), + agentWorkDir, + } + for _, d := range dirs { + if err := os.MkdirAll(d, 0755); err != nil { + log.Fatalf("create dir %s: %v", d, err) + } + } + + // === Graph Memory === + memDB, err := memory.NewGraphDB(filepath.Join(cfg.Daemon.DataDir, "memory", "graph.db")) + if err != nil { + log.Printf("[homed] warning: memory init failed: %v", err) + memDB = nil + } else { + log.Printf("[homed] graph memory initialized") + } + if memDB != nil { + defer memDB.Close() + } + + memIdx := memory.NewIndexer(memDB) + + // === Memory Pipeline === + distiller := pipeline.NewDistiller(memDB, cfg.Daemon.DataDir, pipeline.DistillerConfig{ + Interval: 10 * time.Minute, + RetentionDays: 7, + BatchSize: 50, + }) + if memDB != nil { + distiller.Start() + defer distiller.Stop() + } + + // === Skills === + skMgr := skill.NewManager(filepath.Join(cfg.Daemon.DataDir, "skills")) + if err := skMgr.Init(); err != nil { + log.Printf("[homed] warning: skill init failed: %v", err) + } + + // === Plugin Registry (OpenClaw SKILL.md compatible) === + pluginReg := plugin.NewRegistry() + // 注册内置原生插件工厂 + pluginReg.RegisterNative("qq", func(name string, config map[string]interface{}, iom *agentIO.IOManager) (agentIO.Device, error) { + wsURL, _ := config["entry"].(string) + if wsURL == "" { + wsURL = "ws://127.0.0.1:6700" + } + accessToken, _ := config["access_token"].(string) + return onebot.NewDevice(name, wsURL, accessToken, iom), nil + }) + + // === Supervisor === + sup := supervisor.New(cfg) + if err := sup.Start(); err != nil { + log.Fatalf("start supervisor: %v", err) + } + + // === Lua VM === + luaVM := luapkg.NewVM(filepath.Join(cfg.Daemon.DataDir, "adapters")) + if err := luaVM.Start(); err != nil { + log.Printf("[homed] warning: lua vm init failed: %v", err) + } else { + defer luaVM.Stop() + } + + // === IO Abstraction Layer (唯一输入路径) === + iom := agentIO.NewIOManager() + iom.RegisterDevice(agentIO.NewMicrophone("mic", 16000, iom)) + iom.RegisterDevice(agentIO.NewSpeaker("speaker", iom)) + // mic 输入 → speaker 输出(语音 I/O 配对) + iom.RegisterOutputRoute("mic", "speaker") + iom.RegisterOutputRoute("voice", "speaker") + iom.RegisterDevice(agentIO.NewCamera("camera", iom)) + iom.RegisterDevice(agentIO.NewRobotArm("arm", iom)) + iom.RegisterDevice(agentIO.NewGPIODevice("gpio", []int{2, 3, 4, 17}, iom)) + iom.StartAll() + defer iom.StopAll() + + // 插件绑定 IO 管理器 → 插件自动注册为 IO 设备 + pluginReg.SetIOManager(iom) + // 首次加载插件 + if result, err := pluginReg.Reload(filepath.Join(cfg.Daemon.DataDir, "plugins")); err != nil { + log.Printf("[homed] warning: load plugins: %v", err) + } else { + log.Printf("[homed] %s", result) + } + + // === Text Memory (三层记忆: Context → Text → Graph) === + textMem := text.New(filepath.Join(cfg.Daemon.DataDir, "memory", "text")) + if err := textMem.Start(); err != nil { + log.Printf("[homed] warning: text memory start: %v", err) + } else { + defer textMem.Stop() + log.Printf("[homed] text memory active at %s", filepath.Join(cfg.Daemon.DataDir, "memory", "text")) + } + + // Wire IO output events → TextMemory + distiller → GraphMemory + if distiller != nil { + go func() { + for evt := range iom.OutputChan() { + if evt.Target == "memory" && evt.Type == "memory_candidate" { + source, _ := evt.Payload["source"].(string) + input, _ := evt.Payload["input"].(string) + response, _ := evt.Payload["response"].(string) + toolsUsed, _ := evt.Payload["tools_used"].([]string) + agentID, _ := evt.Payload["agent_id"].(string) + + // 1. 写文本记忆(持久化原始日志) + if input != "" && textMem != nil { + te := text.Event{ + Timestamp: time.Now().Unix(), + Source: source, + Input: input, + Response: response, + ToolsUsed: toolsUsed, + AgentID: agentID, + } + if err := textMem.Append(te); err != nil { + log.Printf("[homed] text memory append: %v", err) + } + } + + // 2. 喂蒸馏器(生成三元组 → 图记忆) + if input != "" { + distiller.Append("agent", "user", input) + } + if response != "" { + distiller.Append("agent", "assistant", response) + } + } + } + }() + } + + // === Change Tracker (overlayfs-based, 追踪所有修改) === + trk := tracker.NewTracker(cfg.Daemon.DataDir, agentWorkDir) + if err := trk.Init(); err != nil { + log.Printf("[homed] warning: tracker init: %v", err) + } else { + if err := trk.Start(); err != nil { + log.Printf("[homed] warning: tracker mount overlay: %v (non-fatal: no overlayfs support?)", err) + } else { + log.Printf("[homed] change tracker active at %s", trk.MergeDir()) + } + } + + // === API Provider (唯一输出路径) === + apiKey := cfg.LLM.APIKey + if apiKey == "" { + apiKey = os.Getenv("DEEPSEEK_API_KEY") + } + provider := agentAPI.NewOpenAIProvider(agentAPI.BaseConfig{ + Model: cfg.LLM.Model, + BaseURL: cfg.LLM.BaseURL, + APIKey: apiKey, + Temperature: cfg.LLM.Temperature, + MaxTokens: cfg.LLM.MaxTokens, + }) + + // === Personality (固定人格内核) === + personalPath := filepath.Join(cfg.Daemon.DataDir, "personal", "personal.md") + personality, err := agentPkg.LoadPersonality(personalPath) + if err != nil { + log.Printf("[homed] warning: load personality: %v", err) + } + if personality != nil && personality.Content != "" { + log.Printf("[homed] personality loaded (%d bytes)", len(personality.Content)) + } + + // === Document Memory (第二层记忆:上下文→文档) === + docStore := document.NewStore(filepath.Join(cfg.Daemon.DataDir, "memory", "documents")) + if err := docStore.Start(); err != nil { + log.Printf("[homed] warning: document store: %v", err) + } + + // === Knowledge Store (知识库) === + ks := knowledge.NewStore(filepath.Join(cfg.Daemon.DataDir, "knowledge")) + if err := ks.Start(); err != nil { + log.Printf("[homed] warning: knowledge store: %v", err) + } else { + log.Printf("[homed] knowledge store active with %d items", len(ks.List())) + } + + // === Single Agent Core === + agent := agentCore.New(agentCore.AgentConfig{ + ID: "main", + SystemPrompt: "你是一个智能家庭管家,持续运行。你有以下工具:\n1. memory_recall — 查询图记忆\n2. memory_commit — 写入图记忆\n3. memory_introspect — 查看记忆统计\n4. knowledge_search — 搜索知识库\n5. doc_query — 查询文档记忆\n6. doc_commit — 写入文档记忆\n\n当用户问及个人信息或历史时,调用 memory_recall 工具来查询。当用户告诉了你新的个人信息时,调用 memory_commit 来记住。需要查询知识时使用 knowledge_search。", + Provider: provider, + IO: iom, + Memory: memDB, + Indexer: memIdx, + Skills: skMgr, + Tracker: trk, + MaxToolTurns: 10, + DocStore: docStore, + Knowledge: ks, + Personality: personality, + }) + agent.Start() + defer agent.Stop() + + // Wire supervisor with tracker + agent registration (after both exist) + sup.SetTracker(trk) + sup.RegisterAgent("main") + + log.Printf("[homed] main agent started, model=%s base=%s", cfg.LLM.Model, cfg.LLM.BaseURL) + + // === HTTP API === + handler := api.NewHandler(sup, memDB, skMgr, luaVM, cfg, iom, textMem, ks) + + mux := http.NewServeMux() + handler.RegisterRoutes(mux) + + server := &http.Server{ + Addr: cfg.Daemon.ListenAddr, + Handler: mux, + } + + go func() { + log.Printf("[homed] HTTP API listening on %s", cfg.Daemon.ListenAddr) + if err := server.ListenAndServe(); err != nil && err != http.ErrServerClosed { + log.Fatalf("http server: %v", err) + } + }() + + sigCh := make(chan os.Signal, 1) + signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM) + <-sigCh + + log.Printf("[homed] shutting down...") + if trk != nil { + trk.Stop() + } + sup.Shutdown() + server.Close() + log.Printf("[homed] stopped") +} + +func countIOPlugins(r *plugin.Registry) int { + n := 0 + for _, p := range r.List() { + if p.IOConfig() != nil { + n++ + } + } + return n +} diff --git a/config/config.go b/config/config.go new file mode 100644 index 0000000..bd184da --- /dev/null +++ b/config/config.go @@ -0,0 +1,91 @@ +package config + +import ( + "fmt" + "os" + "path/filepath" + "time" + + "gitcode.com/JianFeeeee/HomeAgent/pkg/types" + "gopkg.in/yaml.v3" +) + +const DefaultConfigPath = "/etc/homeagent/config.yaml" + +func DefaultConfig() types.Config { + return types.Config{ + Daemon: types.DaemonConfig{ + ListenAddr: ":8080", + DataDir: "/var/lib/homeagent", + HeartbeatInterval: 15 * time.Second, + CheckInterval: 30 * time.Second, + LogLevel: "info", + }, + LLM: types.LLMConfig{ + Provider: "openai", + Model: "deepseek-v4-flash", + BaseURL: "https://api.deepseek.com", + Temperature: 0.7, + MaxTokens: 4096, + }, + Defaults: types.AgentConfig{ + Image: "homeagent/agent-base:latest", + LLMEndpoints: []string{"https://api.openai.com/v1"}, + SnapshotPolicy: types.SnapshotPolicy{ + Interval: 10 * time.Minute, + MaxSnapshots: 20, + PreAction: true, + PostAction: false, + }, + RollbackPolicy: types.RollbackPolicy{ + MaxRetries: 3, + HealthThreshold: types.HealthDown, + CooldownPeriod: 30 * time.Second, + AutoRollback: true, + }, + ResourceLimit: types.ResourceLimit{ + CPU: "2", + Memory: "2g", + Disk: "10g", + Network: true, + }, + OpenClawEnabled: true, + }, + } +} + +func Load(path string) (*types.Config, error) { + cfg := DefaultConfig() + + data, err := os.ReadFile(path) + if err != nil { + if os.IsNotExist(err) { + return &cfg, nil + } + return nil, fmt.Errorf("read config: %w", err) + } + + if err := yaml.Unmarshal(data, &cfg); err != nil { + return nil, fmt.Errorf("parse config: %w", err) + } + + return &cfg, nil +} + +func Save(path string, cfg *types.Config) error { + dir := filepath.Dir(path) + if err := os.MkdirAll(dir, 0755); err != nil { + return fmt.Errorf("create config dir: %w", err) + } + + data, err := yaml.Marshal(cfg) + if err != nil { + return fmt.Errorf("marshal config: %w", err) + } + + if err := os.WriteFile(path, data, 0644); err != nil { + return fmt.Errorf("write config: %w", err) + } + + return nil +} diff --git a/config/config.yaml b/config/config.yaml new file mode 100644 index 0000000..4f2ede1 --- /dev/null +++ b/config/config.yaml @@ -0,0 +1,34 @@ +daemon: + listen_addr: ":8080" + data_dir: "/var/lib/homeagent" + heartbeat_interval: 15s + check_interval: 30s + log_level: "info" + +defaults: + image: "homeagent/agent-base:latest" + llm_endpoints: + - "https://api.openai.com/v1" + snapshot_policy: + interval: 10m + max_snapshots: 20 + pre_action: true + post_action: false + rollback_policy: + max_retries: 3 + health_threshold: 3 + cooldown_period: 30s + auto_rollback: true + resource_limit: + cpu: "2" + memory: "2g" + disk: "10g" + network: true + openclaw_enabled: true + +agents: + - id: "default" + name: "Default Agent" + image: "homeagent/agent-base:latest" + llm_endpoints: + - "https://api.openai.com/v1" diff --git a/deploy/homeagent.service b/deploy/homeagent.service new file mode 100644 index 0000000..6d3486d --- /dev/null +++ b/deploy/homeagent.service @@ -0,0 +1,33 @@ +[Unit] +Description=HomeAgent - Personal AI Home Assistant +Documentation=https://github.com/trueagent/HomeAgent +After=network-online.target docker.service +Wants=network-online.target + +[Service] +Type=simple +ExecStart=/usr/local/bin/homed -config /etc/homeagent/config.yaml -data /var/lib/homeagent +Restart=always +RestartSec=10 +StartLimitBurst=3 +StartLimitInterval=60s + +# Security hardening +NoNewPrivileges=true +ProtectSystem=strict +ProtectHome=true +PrivateTmp=true +PrivateDevices=false +DevicePolicy=closed +# Audio devices +DeviceAllow=/dev/snd rw +DeviceAllow=/dev/dsp rw + +# Resource limits +LimitNOFILE=65536 +LimitNPROC=256 +MemoryMax=2G +CPUQuota=100% + +[Install] +WantedBy=multi-user.target diff --git a/docs/ARCHITECTURE.md b/docs/ARCHITECTURE.md new file mode 100644 index 0000000..feb5710 --- /dev/null +++ b/docs/ARCHITECTURE.md @@ -0,0 +1,495 @@ +# HomeAgent 架构设计文档 + +> 基于 NextAgent 认知解耦架构,结合图记忆与工具调用系统的单二进制 AI 家庭助手。 +> 参考设计:NextAgent — 严格的边界划分 + 工具调用;TrulyMEM — 自主图记忆系统 + +--- + +## 一、核心设计原则 + +### 1. 认知解耦架构(源自 NextAgent) + +系统分为三个严格边界层: + +``` +┌──────────────────────────────────────────────────────────────────┐ +│ 外层:IO 抽象层(唯一输入路径) │ +│ IOManager ── Microphone / Camera / GPIO / HTTP / Sensors │ +│ 所有外部输入 → InputEvent → inputCh │ +└──────────────────────────┬───────────────────────────────────────┘ + │ +┌──────────────────────────▼───────────────────────────────────────┐ +│ 中层:Agent 操作层(核心编排器) │ +│ Agent Core ── 事件循环 consume inputCh │ +│ ├─ 构建上下文(记忆索引 + 技能注入 + ToolDef) │ +│ ├─ Provider.Chat() → CompletionResponse │ +│ └─ 通过 IO 层输出 text + memory_candidate 事件 │ +└──────────────────────────┬───────────────────────────────────────┘ + │ +┌──────────────────────────▼───────────────────────────────────────┐ +│ 内层:API 抽象层(唯一输出路径) │ +│ Provider ── OpenAI / Ollama / LuaAdaptedProvider │ +│ 所有 LLM 调用通过 Provider.Chat() / ChatStream() │ +└──────────────────────────────────────────────────────────────────┘ +``` + +**三条核心规则:** +1. 所有外部输入 → 必须通过 `IOManager.InjectInput()` / `InjectText()` 注入 +2. 所有 LLM 调用 → 必须通过 `Provider.Chat()` / `ChatStream()` 发出 +3. Agent Core 不直接操作记忆系统,只发射 `memory_candidate` 事件,由 Memory Pipeline 异步消费 + +### 2. 自主记忆系统(源自 TrulyMEM) + +``` +User Input ──→ Agent Core ──→ Provider ──→ LLM Response + │ + ▼ + IO 层发射 memory_candidate 事件 + │ + ▼ + Memory Pipeline Distiller 消费 + │ + ▼ + 周期蒸馏 → Graph Memory (SQLite) + │ + ▼ + 删除原始会话文件 +``` + +**关键约束:** +- Agent 上下文中只注入**记忆索引 + 摘要**,**永远不**注入原始文本 +- Agent 必须通过 `memory_recall` tool call **主动查询** 获取完整记忆细节 +- 记忆蒸馏完全异步、自主运行,不阻塞主流程 + +--- + +## 二、系统层次详解 + +### 2.1 IO 抽象层(`internal/agent/io/channel.go`) + +**唯一输入路径。** 所有外部输入必须通过此层进入系统。 + +``` +Device(Microphone) ─┐ +Device(Camera) ─┤ +Device(GPIO) ─┤──→ IOManager.InjectInput() → InputEvent → inputCh → Agent Core +Device(RobotArm) ─┤ +HTTP API ─┘ +``` + +**核心类型:** + +| 类型 | 说明 | +|------|------| +| `InputEvent` | Source + Type + Payload — 所有外部输入的标准化格式 | +| `OutputEvent` | Target + Type + Payload — 所有输出的标准化格式 | +| `Device` | 接口:Name() / Type() / Tools() / Execute() | +| `DeviceType` | Input / Output / IO | +| `ToolDef` | Name + Description + Parameters — 与 LLM Function Calling 同构 | + +**内置设备(目前为桩实现,待真正硬件接入):** + +| 设备 | 方向 | ToolDef 暴露 | +|------|------|-------------| +| Microphone | Input | `{name}_capture` — 录音 | +| Speaker | Output | `{name}_speak` — 语音播放 | +| Camera | Input | `{name}_capture` 拍照 + `{name}_stream` 视频流 | +| RobotArm | IO | `{name}_move` 移动 + `{name}_grip` 夹爪 | +| GPIO | IO | `{name}_gpio_write` + `{name}_gpio_read` | + +**Device 与 Tool 同构原则:** +- 所有 Device 的 `Tools()` 返回 `[]ToolDef`,格式与 LLM Function Calling 完全一致 +- Agent Core 自动收集所有 Device 的 ToolDef 合并到请求的 `tools` 字段 +- Agent 通过 Function Calling 调用设备 capability + +### 2.2 API 抽象层(`internal/agent/api/provider.go`) + +**唯一输出路径。** 所有 LLM 请求通过此层发出。 + +``` +Agent Core ──→ ProviderManager ──→ Provider.Chat() + │ + ┌─────────┼─────────┐ + ▼ ▼ ▼ + OpenAI Ollama LuaAdaptedProvider + │ + ┌───────┴───────┐ + ▼ ▼ + Lua Adapter Base Provider + (transform) (OpenAI/Ollama) +``` + +**Provider 接口:** + +```go +type Provider interface { + Name() string + Chat(ctx, req) → (*CompletionResponse, error) + ChatStream(ctx, req) → (<-chan StreamChunk, error) +} +``` + +| 实现 | 说明 | +|------|------| +| `OpenAIProvider` | 标准 OpenAI API 格式,支持 /chat/completions | +| `OllamaProvider` | Ollama /api/chat 格式,本地部署 | +| `LuaAdaptedProvider` | 通过 Lua 脚本转换请求/响应的适配 wrapper | + +**Lua 适配器机制:** +- 适配器文件位于 `{dataDir}/adapters/*.lua` +- 每个适配器返回 Lua table 包含 `name` + `transform_request` + `transform_response` +- 首次运行时自动从 embed.FS 复制捆绑适配器(openai / deepseek / ollama / custom) +- 支持热重载(`POST /api/v1/adapters`) +- 用于兼容不同 API 供应商的请求/响应格式差异 + +### 2.3 Agent 操作层(`internal/agent/core/agent.go`) + +**纯编排器,不涉及张量运算。** + +``` +eventLoop() + │ + ▼ select on inputCh +handleInput(evt) + │ + ▼ +process(input) + ├─ indexer.BuildContext(input) → 仅摘要+索引 + ├─ buildSystemPrompt() → 拼接 system prompt + ├─ buildToolDefs() → 收集 IO 工具 + 记忆工具 + ├─ provider.Chat(req) → LLM 调用 + ├─ 记忆候选事件 → IO 层 EmitOutput("memory", "memory_candidate", ...) + └─ 输出 → IO 层 EmitOutput(source, "text", ...) +``` + +**关键设计决策:** +- Agent 不直接持有 `memory.GraphDB` 引用 — 只通过 `Indexer` 构建上下文 +- Agent 不直接调用 `memory.Commit()` — 只发射事件让 Pipeline 异步处理 +- `Indexer.BuildContext()` 只返回内存索引摘要,不返回原始数据 +- `buildSystemPrompt()` 中注入 `memory_recall` / `memory_commit` / `memory_introspect` 工具说明 +- **所有输出走 IO 层** — 文本输出和记忆事件都通过 EmitOutput + +### 2.4 图记忆系统(`internal/memory/graph.go`) + +SQLite 三元组存储,支持实体-关系-实体的图遍历。 + +**数据库 Schema:** + +```sql +entities(id PK, name UNIQUE, type, mention_count, created_at, updated_at) +relations(id PK, source_id FK→entities, target_id FK→entities, + relation_type, confidence, status, session_id, turn_id, + created_at, updated_at, date_bucket) +``` + +**核心操作:** + +| 操作 | 说明 | +|------|------| +| `Commit(triples, sessionID, turnID)` | 写入三元组 → 自动 upsert 实体 + 插入关系 | +| `Recall(keywords, seedEntities, depth, session)` | 关键词搜索 → 图遍历 → 返回实体+关系 | +| `Purge(criteria, mode)` | 软/硬删除匹配的关系 | +| `Introspect()` | 统计信息:实体数、关系数、热点实体 | +| `Archive(days)` | 归档超过指定天数的关系 | + +**Context Injection 机制(`internal/memory/indexer.go`):** +- `BuildContext(userInput)` → 关键词提取 → `GraphDB.Recall()` → 构建摘要 +- `FormatContext(context)` → 输出格式如: + `【记忆索引】关联 N 个记忆实体,高频:A、B、C 索引: A, B, C | 需更多细节请用 memory_recall 查询` +- `BuildToolPrompt()` → 生成 `memory_recall/commit/introspect/purge` 工具的 prompt 说明 +- `GetToolDefinitions()` → 返回 LLM Function Calling 格式的工具定义 + +### 2.5 记忆管道(`internal/memory/pipeline/pipeline.go`) + +自主异步蒸馏管线。 + +``` +Agent Core → EmitOutput("memory", "memory_candidate", {input, response}) + │ + ▼ +onMemory(input, response) 回调 + │ + ▼ +Append() → records[] 内存缓冲区 + │ + ▼ 每 10 分钟触发 +distillLoop() + │ + ▼ +extractKeyTriples() → GraphDB.Commit() + │ + ▼ +cleanupRawFiles() 删除超期原始文件 +``` + +**配置:** +- `Interval: 10m` — 每 10 分钟蒸馏一次 +- `RetentionDays: 7` — 原始记录保留 7 天 +- `BatchSize: 50` — 每批处理 50 条 + +### 2.6 技能系统(`internal/skill/manager.go`) + +兼容 OpenClaw 格式的技能管理。 + +**技能格式:** +- `{skillsDir}/{name}/SKILL.md` — Markdown 描述文件 +- `{skillsDir}/{name}/skill.json` — 可选的元数据文件 + +**注入机制:** +- `GetInjectedPrompt()` → 收集所有已启用的技能内容注入到 system prompt +- 支持安装/卸载/启用/禁用 + +### 2.7 Supervisor 守护进程(`internal/supervisor/daemon.go`) + +```go +type Daemon struct { + cfg *types.Config + cm *container.Manager // Docker 容器管理 + nm *network.Monitor // 网络监控 + sm *snapshot.Manager // 快照管理 + agents map[AgentID]*agentInstance +} +``` + +**职责:** +- Agent 生命周期管理(launch / restart / shutdown) +- 健康检查循环(心跳间隔 15s) +- 自动快照循环(快照间隔 10m) +- 故障恢复:失败 MaxRetries(3) 次后自动回滚到最近快照 +- 网络监控:检查 LLM API 可达性,影响 Agent 健康状态 + +### 2.8 快照与回滚(`internal/snapshot/manager.go`) + +基于 Docker commit/save/load 的版本管理。 + +**快照流程:** +`Docker Commit(container → image) → SaveImage(image → .tar) → 记录快照元数据` + +**回滚流程:** +`Stop(container) → Remove(container) → LoadImage(.tar) → 创建新容器 → Start` + +**策略:** +- 定时快照(每 10m) +- 操作前快照(PreAction — 可选) +- 限制最大保留(默认 20 个) +- 溢出时自动删除最旧的 + +### 2.9 网络监控(`internal/network/monitor.go`) + +异步定时检查 LLM API 端点可达性。 + +**监控方式:** +- HTTP HEAD 请求到配置的 endpoints +- 并发检查(goroutine per endpoint) +- DNS 解析检查(fallback: google.com → baidu.com) +- 结果聚合:`LLMAPIReachable` + `DNSResolving` + 平均延迟 + +**影响:** +- 网络不可达 → Agent 状态变为 `HealthDegraded` +- 持续不可达 → 触发回滚策略 + +### 2.10 Lua 虚拟机(`internal/lua/vm.go`) + +纯 Go 的 gopher-lua 5.1 VM,用于 API 格式适配器。 + +**功能:** +- 加载 `{adapterDir}/*.lua` 适配器脚本 +- `CallTransform(name, input)` — 调用 adapter 的 transform_request +- `CallResponseTransform(name, raw)` — 调用 adapter 的 transform_response +- `ReloadAll()` — 热重载所有适配器 +- 内置 mock 函数:`log()`, `json_encode()`, `http_get()`, `http_post()` + +### 2.11 Embedder(`internal/embed/embedder.go`) + +向量嵌入接口,支持文本相似度计算。 + +| 实现 | 说明 | +|------|------| +| `OllamaEmbedder` | 通过 Ollama API 获取嵌入向量(默认: nomic-embed-text, 768d) | +| `HashEmbedder` | 基于字符哈希的本地嵌入(无需外部依赖),用于备选方案 | + +**可用性:** Embedder 已定义但尚未集成到记忆系统中。 + +### 2.12 Tokenizer(`internal/tokenizer/jieba.go`) + +中文分词工具,基于 gojieba。全局单例,线程安全。 + +**功能:** +- `ExtractKeywords(text, topK)` — 提取关键词(TF-IDF 加权) +- `Cut(text)` — 分词 +- `Tag(text)` — 词性标注 + +--- + +## 三、数据流全景 + +### 3.1 正常交互流程 + +``` +外部输入(HTTP POST / voice / GPIO 事件) + │ + ▼ +IOManager.InjectInput() + → InputEvent{Source, Type, Payload} + → 推入 inputCh + │ + ▼ +Agent Core eventLoop() + 1. consume InputEvent + 2. Indexer.BuildContext(input) → 记忆摘要(仅索引+摘要) + 3. buildSystemPrompt() → 拼接 system + 记忆 + 技能 + 工具 + 4. buildToolDefs() → IO 工具 + 记忆工具 + 5. Provider.Chat(req) → LLM 响应 + 6. EmitOutput(target, "text", response) + 7. EmitOutput("memory", "memory_candidate", {input, response}) + │ + ┌─────┴─────┐ + ▼ ▼ +IO 输出 Memory Pipeline(异步) + │ │ + ▼ ▼ +HTTP Append() → 周期蒸馏 +Response → GraphDB.Commit() +Speaker → 清理原始文件 +``` + +### 3.2 记忆查询流程 + +``` +Agent 推理中决定调用 memory_recall + │ + ▼ +LLM 返回 tool_call: {name: "memory_recall", args: {query_intent: "..."}} + │ + ▼ +Agent Core 解析 tool_call → 调用 GraphDB.Recall(keywords) + │ + ▼ +返回实体+关系数据 → 注入后续 LLM 请求上下文 +``` + +--- + +## 四、配置系统 + +配置文件: `/etc/homeagent/config.yaml`(YAML)。若文件不存在则使用默认配置。 + +```yaml +daemon: + listen_addr: ":8080" + data_dir: "/var/lib/homeagent" + heartbeat_interval: 15s + check_interval: 30s + log_level: "info" + +defaults: + image: "homeagent/agent-base:latest" + llm_endpoints: + - "https://api.openai.com/v1" + snapshot_policy: + interval: 10m + max_snapshots: 20 + pre_action: true + post_action: false + rollback_policy: + max_retries: 3 + health_threshold: 3 + cooldown_period: 30s + auto_rollback: true + openclaw_enabled: true + +agents: [] +``` + +--- + +## 五、HTTP API + +| 方法 | 路径 | 说明 | +|------|------|------| +| GET | `/api/v1/status` | 系统状态 | +| GET | `/api/v1/agents` | Agent 列表 | +| POST | `/api/v1/agents` | 创建 Agent | +| GET | `/api/v1/agents/{id}` | Agent 详情 | +| POST | `/api/v1/agents/{id}/start/stop/restart` | 操作 | +| GET | `/api/v1/agents/{id}/snapshots` | 快照列表 | +| POST | `/api/v1/agents/{id}/rollback/{snap}` | 回滚 | +| GET | `/api/v1/memory?q=关键词` | 记忆检索 | +| POST | `/api/v1/memory` | 写入三元组 | +| DELETE | `/api/v1/memory` | 删除记忆 | +| GET | `/api/v1/memory/context?q=...` | 获取上下文注入 | +| GET | `/api/v1/memory/tools` | 记忆工具定义 | +| GET | `/api/v1/skills` | 技能列表 | +| POST | `/api/v1/skills` | 安装技能 | +| DELETE | `/api/v1/skills?name=...` | 卸载技能 | +| GET | `/api/v1/adapters` | 适配器列表 | +| POST | `/api/v1/adapters` | 安装适配器 | +| DELETE | `/api/v1/adapters/{name}` | 删除适配器 | +| GET | `/api/v1/network` | 网络状态 | +| GET | `/api/v1/config` | 配置查看 | +| PUT | `/api/v1/config` | 配置更新 | +| GET | `/` | WebUI 仪表盘 | + +--- + +## 六、数据目录结构 + +``` +{dataDir}/ +├── config.yaml # 系统配置 +├── memory/ +│ ├── graph.db # SQLite 图记忆数据库 +│ └── raw/ # 原始会话记录文件 +│ └── raw_*.jsonl +├── skills/ # 安装的技能 +│ └── {name}/ +│ ├── SKILL.md +│ └── skill.json +├── adapters/ # Lua API 格式适配器 +│ ├── openai.lua +│ ├── deepseek.lua +│ ├── ollama.lua +│ └── custom.lua +└── snapshots/ # Docker 快照 + └── {agent_id}/ + └── snap_*.tar +``` + +--- + +## 七、构建与部署 + +**构建:** + +```bash +make build # 编译主二进制(~15MB) +make install # 编译 + 安装到 /usr/local/bin +make test # 运行测试 +make fmt # gofmt +make lint # golangci-lint +``` + +**部署:** + +- 单二进制:`homed -config /etc/homeagent/config.yaml -data /var/lib/homeagent` +- systemd:`deploy/homeagent.service` +- 依赖:Go 1.19+(CGo enabled,用于 go-sqlite3);Docker 可选(快照/回滚) + +--- + +## 八、设计限制与后续计划 + +### 已知限制 +1. Agent Core 目前是单轮 tool_call 处理,尚未实现完整的多轮 tool 执行循环 +2. Embedder 已定义但未接入图记忆 — 缺失向量相似度排序和 ANN 索引 +3. 记忆蒸馏器使用简单启发式三元组提取,生产环境应调用 LLM 进行结构化抽取 +4. Linux namespace 隔离(overlayfs)尚未实现作为 Docker 替代方案 +5. Speaker / Microphone / Camera 均为桩实现,无实际 ALSA/PulseAudio/Video4Linux 驱动 +6. 无真正的唤醒词检测 +7. LuaAdaptedProvider 失败时无重试/降级逻辑 +8. 无持久化消息历史管理(目前仅内存中保留最近 50 条) + +### 路线图 +- **近期**:完成 tool_call 执行循环 → 设备真正驱动 → 向量记忆增强 +- **中期**:单二进制 namespace 隔离 → 唤醒词检测 → CI/CD 流水线 +- **远期**:多 Agent 协作 → 分布式部署 → 联邦记忆 diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..e27f7ce --- /dev/null +++ b/go.mod @@ -0,0 +1,14 @@ +module gitcode.com/JianFeeeee/HomeAgent + +go 1.19 + +require ( + github.com/mattn/go-sqlite3 v1.14.22 + gopkg.in/yaml.v3 v3.0.1 +) + +require ( + github.com/gorilla/websocket v1.5.3 // indirect + github.com/yanyiwu/gojieba v1.4.7 // indirect + github.com/yuin/gopher-lua v1.1.2 // indirect +) diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..d1b2f0d --- /dev/null +++ b/go.sum @@ -0,0 +1,12 @@ +github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg= +github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE= +github.com/mattn/go-sqlite3 v1.14.22 h1:2gZY6PC6kBnID23Tichd1K+Z0oS6nE/XwU+Vz/5o4kU= +github.com/mattn/go-sqlite3 v1.14.22/go.mod h1:Uh1q+B4BYcTPb+yiD3kU8Ct7aC0hY9fxUwlHK0RXw+Y= +github.com/yanyiwu/gojieba v1.4.7 h1:2YkXELcYLTE0SJetq6xv4MjpEikWga6VpFn4jIFFQ/k= +github.com/yanyiwu/gojieba v1.4.7/go.mod h1:JUq4DddFVGdHXJHxxepxRmhrKlDpaBxR8O28v6fKYLY= +github.com/yuin/gopher-lua v1.1.2 h1:yF/FjE3hD65tBbt0VXLE13HWS9h34fdzJmrWRXwobGA= +github.com/yuin/gopher-lua v1.1.2/go.mod h1:7aRmXIWl37SqRf0koeyylBEzJ+aPt8A+mmkQ4f1ntR8= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/homed b/homed new file mode 100755 index 0000000..090ec33 Binary files /dev/null and b/homed differ diff --git a/internal/agent/api/provider.go b/internal/agent/api/provider.go new file mode 100644 index 0000000..338d604 --- /dev/null +++ b/internal/agent/api/provider.go @@ -0,0 +1,509 @@ +package api + +import ( + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" + "sync" + "time" + + luaVM "gitcode.com/JianFeeeee/HomeAgent/internal/lua" +) + +type Message struct { + Role string `json:"role"` + Content string `json:"content"` + ReasoningContent string `json:"reasoning_content,omitempty"` + ToolCallID string `json:"tool_call_id,omitempty"` + ToolCalls []ToolCall `json:"-"` +} + +func (m Message) MarshalJSON() ([]byte, error) { + raw := map[string]interface{}{ + "role": m.Role, + "content": m.Content, + } + if m.ReasoningContent != "" { + raw["reasoning_content"] = m.ReasoningContent + } + if m.ToolCallID != "" { + raw["tool_call_id"] = m.ToolCallID + } + if len(m.ToolCalls) > 0 { + apiTCs := make([]apiToolCall, len(m.ToolCalls)) + for i, tc := range m.ToolCalls { + argsBytes, _ := json.Marshal(tc.Arguments) + apiTCs[i] = apiToolCall{ + ID: tc.ID, + Type: "function", + Function: apiFunction{ + Name: tc.Name, + Arguments: string(argsBytes), + }, + } + } + raw["tool_calls"] = apiTCs + } + return json.Marshal(raw) +} + +type CompletionRequest struct { + Model string `json:"model,omitempty"` + Messages []Message `json:"messages"` + Temperature float64 `json:"temperature,omitempty"` + MaxTokens int `json:"max_tokens,omitempty"` + Stream bool `json:"stream,omitempty"` + Tools []interface{} `json:"tools,omitempty"` + ToolChoice interface{} `json:"tool_choice,omitempty"` + ExtraBody map[string]interface{} `json:"-"` +} + +func (r *CompletionRequest) MarshalJSON() ([]byte, error) { + type Alias CompletionRequest + data, err := json.Marshal((*Alias)(r)) + if err != nil { + return nil, err + } + if len(r.ExtraBody) == 0 { + return data, nil + } + var raw map[string]interface{} + if err := json.Unmarshal(data, &raw); err != nil { + return nil, err + } + for k, v := range r.ExtraBody { + raw[k] = v + } + return json.Marshal(raw) +} + +type CompletionResponse struct { + Content string `json:"content"` + FinishReason string `json:"finish_reason,omitempty"` + TokenUsage TokenUsage `json:"token_usage,omitempty"` + ToolCalls []ToolCall `json:"tool_calls,omitempty"` +} + +type TokenUsage struct { + Prompt int `json:"prompt"` + Completion int `json:"completion"` + Total int `json:"total"` +} + +type ToolCall struct { + ID string `json:"id"` + Type string `json:"type"` + Name string `json:"name"` + Arguments map[string]interface{} `json:"arguments"` +} + +type apiToolCall struct { + ID string `json:"id"` + Type string `json:"type"` + Function apiFunction `json:"function"` +} + +type apiFunction struct { + Name string `json:"name"` + Arguments string `json:"arguments"` +} + +type StreamChunk struct { + Content string `json:"content"` + Done bool `json:"done"` + ToolCall *ToolCall `json:"tool_call,omitempty"` +} + +type Provider interface { + Name() string + Chat(ctx context.Context, req *CompletionRequest) (*CompletionResponse, error) + ChatStream(ctx context.Context, req *CompletionRequest) (<-chan StreamChunk, error) +} + +type BaseConfig struct { + Model string `json:"model"` + BaseURL string `json:"base_url"` + APIKey string `json:"api_key"` + Temperature float64 `json:"temperature"` + MaxTokens int `json:"max_tokens"` +} + +type OpenAIProvider struct { + cfg BaseConfig + client *http.Client +} + +func NewOpenAIProvider(cfg BaseConfig) *OpenAIProvider { + if cfg.BaseURL == "" { + cfg.BaseURL = "https://api.openai.com/v1" + } + if cfg.Temperature == 0 { + cfg.Temperature = 0.7 + } + if cfg.MaxTokens == 0 { + cfg.MaxTokens = 4096 + } + return &OpenAIProvider{ + cfg: cfg, + client: &http.Client{Timeout: 60 * time.Second}, + } +} + +func (p *OpenAIProvider) Name() string { return "openai" } + +func (p *OpenAIProvider) Chat(ctx context.Context, req *CompletionRequest) (*CompletionResponse, error) { + if req.Model == "" { + req.Model = p.cfg.Model + } + + body, _ := json.Marshal(req) + httpReq, _ := http.NewRequestWithContext(ctx, "POST", p.cfg.BaseURL+"/chat/completions", strings.NewReader(string(body))) + httpReq.Header.Set("Content-Type", "application/json") + httpReq.Header.Set("Authorization", "Bearer "+p.cfg.APIKey) + + resp, err := p.client.Do(httpReq) + if err != nil { + return nil, fmt.Errorf("api call: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != 200 { + respBody, _ := io.ReadAll(resp.Body) + return nil, fmt.Errorf("api error %d: %s", resp.StatusCode, string(respBody)) + } + + var rawResult struct { + Choices []struct { + Message struct { + Content *string `json:"content"` + ReasoningContent *string `json:"reasoning_content"` + ToolCalls []rawToolCall `json:"tool_calls"` + Role string `json:"role"` + } `json:"message"` + FinishReason string `json:"finish_reason"` + } `json:"choices"` + Usage struct { + PromptTokens int `json:"prompt_tokens"` + CompletionTokens int `json:"completion_tokens"` + TotalTokens int `json:"total_tokens"` + } `json:"usage"` + } + + if err := json.NewDecoder(resp.Body).Decode(&rawResult); err != nil { + return nil, fmt.Errorf("decode: %w", err) + } + + if len(rawResult.Choices) == 0 { + return nil, fmt.Errorf("no choices returned") + } + + ch := rawResult.Choices[0] + content := "" + if ch.Message.Content != nil { + content = *ch.Message.Content + } + + var toolCalls []ToolCall + for _, tc := range ch.Message.ToolCalls { + tc := tc + args := make(map[string]interface{}) + if err := json.Unmarshal([]byte(tc.Function.Arguments), &args); err != nil { + args["_raw"] = tc.Function.Arguments + } + toolCalls = append(toolCalls, ToolCall{ + ID: tc.ID, + Type: tc.Type, + Name: tc.Function.Name, + Arguments: args, + }) + } + + return &CompletionResponse{ + Content: content, + FinishReason: ch.FinishReason, + TokenUsage: TokenUsage{ + Prompt: rawResult.Usage.PromptTokens, + Completion: rawResult.Usage.CompletionTokens, + Total: rawResult.Usage.TotalTokens, + }, + ToolCalls: toolCalls, + }, nil +} + +func (p *OpenAIProvider) ChatStream(ctx context.Context, req *CompletionRequest) (<-chan StreamChunk, error) { + req.Stream = true + ch := make(chan StreamChunk, 64) + + body, _ := json.Marshal(req) + httpReq, _ := http.NewRequestWithContext(ctx, "POST", p.cfg.BaseURL+"/chat/completions", strings.NewReader(string(body))) + httpReq.Header.Set("Content-Type", "application/json") + httpReq.Header.Set("Authorization", "Bearer "+p.cfg.APIKey) + + resp, err := p.client.Do(httpReq) + if err != nil { + return nil, fmt.Errorf("stream api: %w", err) + } + + go func() { + defer resp.Body.Close() + defer close(ch) + + decoder := json.NewDecoder(resp.Body) + for { + var line struct { + Choices []struct { + Delta struct { + Content string `json:"content"` + } `json:"delta"` + FinishReason *string `json:"finish_reason"` + } `json:"choices"` + } + + if err := decoder.Decode(&line); err != nil { + break + } + + if len(line.Choices) > 0 { + ch <- StreamChunk{ + Content: line.Choices[0].Delta.Content, + Done: line.Choices[0].FinishReason != nil, + } + } + } + }() + + return ch, nil +} + +type OllamaProvider struct { + cfg BaseConfig + client *http.Client +} + +func NewOllamaProvider(cfg BaseConfig) *OllamaProvider { + if cfg.BaseURL == "" { + cfg.BaseURL = "http://localhost:11434" + } + if cfg.Temperature == 0 { + cfg.Temperature = 0.7 + } + if cfg.MaxTokens == 0 { + cfg.MaxTokens = 4096 + } + return &OllamaProvider{ + cfg: cfg, + client: &http.Client{Timeout: 120 * time.Second}, + } +} + +func (p *OllamaProvider) Name() string { return "ollama" } + +func (p *OllamaProvider) Chat(ctx context.Context, req *CompletionRequest) (*CompletionResponse, error) { + ollamaReq := map[string]interface{}{ + "model": req.Model, + "messages": req.Messages, + "stream": false, + "options": map[string]interface{}{ + "temperature": req.Temperature, + "num_predict": req.MaxTokens, + }, + } + + body, _ := json.Marshal(ollamaReq) + httpReq, _ := http.NewRequestWithContext(ctx, "POST", p.cfg.BaseURL+"/api/chat", strings.NewReader(string(body))) + httpReq.Header.Set("Content-Type", "application/json") + + resp, err := p.client.Do(httpReq) + if err != nil { + return nil, fmt.Errorf("ollama chat: %w", err) + } + defer resp.Body.Close() + + var result struct { + Message struct { + Content string `json:"content"` + } `json:"message"` + DoneReason string `json:"done_reason"` + } + if err := json.NewDecoder(resp.Body).Decode(&result); err != nil { + return nil, fmt.Errorf("decode: %w", err) + } + + return &CompletionResponse{ + Content: result.Message.Content, + FinishReason: result.DoneReason, + }, nil +} + +func (p *OllamaProvider) ChatStream(ctx context.Context, req *CompletionRequest) (<-chan StreamChunk, error) { + ch := make(chan StreamChunk, 64) + + go func() { + defer close(ch) + ch <- StreamChunk{Done: true} + }() + + return ch, nil +} + +type LuaAdaptedProvider struct { + name string + base Provider + vm *luaVM.VM + adapter string +} + +func NewLuaAdaptedProvider(base Provider, vm *luaVM.VM, adapter string) *LuaAdaptedProvider { + return &LuaAdaptedProvider{ + name: fmt.Sprintf("lua_%s", adapter), + base: base, + vm: vm, + adapter: adapter, + } +} + +func (p *LuaAdaptedProvider) Name() string { return p.name } + +func (p *LuaAdaptedProvider) Chat(ctx context.Context, req *CompletionRequest) (*CompletionResponse, error) { + inputMap := map[string]interface{}{ + "model": req.Model, + "messages": messagesToMap(req.Messages), + "temperature": req.Temperature, + "max_tokens": req.MaxTokens, + "stream": false, + } + + transformed, err := p.vm.CallTransform(p.adapter, inputMap) + if err != nil { + return nil, fmt.Errorf("lua transform: %w", err) + } + + transformedReq := &CompletionRequest{ + Model: getString(transformed, "model"), + Temperature: getFloat(transformed, "temperature"), + MaxTokens: int(getFloat(transformed, "max_tokens")), + Stream: false, + } + + if msgs, ok := transformed["messages"].([]interface{}); ok { + for _, m := range msgs { + if mm, ok := m.(map[string]interface{}); ok { + transformedReq.Messages = append(transformedReq.Messages, Message{ + Role: getString(mm, "role"), + Content: getString(mm, "content"), + }) + } + } + } + + resp, err := p.base.Chat(ctx, transformedReq) + if err != nil { + return nil, err + } + + return resp, nil +} + +func (p *LuaAdaptedProvider) ChatStream(ctx context.Context, req *CompletionRequest) (<-chan StreamChunk, error) { + return p.base.ChatStream(ctx, req) +} + +type ProviderManager struct { + mu sync.RWMutex + providers map[string]Provider + default_ string +} + +func NewProviderManager() *ProviderManager { + return &ProviderManager{ + providers: make(map[string]Provider), + } +} + +func (m *ProviderManager) Register(name string, p Provider) { + m.mu.Lock() + defer m.mu.Unlock() + m.providers[name] = p + if m.default_ == "" { + m.default_ = name + } +} + +func (m *ProviderManager) SetDefault(name string) error { + m.mu.Lock() + defer m.mu.Unlock() + if _, ok := m.providers[name]; !ok { + return fmt.Errorf("provider %s not found", name) + } + m.default_ = name + return nil +} + +func (m *ProviderManager) Get(name string) Provider { + m.mu.RLock() + defer m.mu.RUnlock() + if name == "" { + name = m.default_ + } + return m.providers[name] +} + +func (m *ProviderManager) Default() Provider { + m.mu.RLock() + defer m.mu.RUnlock() + return m.providers[m.default_] +} + +func (m *ProviderManager) List() []string { + m.mu.RLock() + defer m.mu.RUnlock() + var names []string + for n := range m.providers { + names = append(names, n) + } + return names +} + +type rawToolCall struct { + ID string `json:"id"` + Type string `json:"type"` + Function struct { + Name string `json:"name"` + Arguments string `json:"arguments"` + } `json:"function"` +} + +func messagesToMap(msgs []Message) []interface{} { + result := make([]interface{}, len(msgs)) + for i, m := range msgs { + result[i] = map[string]interface{}{ + "role": m.Role, + "content": m.Content, + } + } + return result +} + +func getString(m map[string]interface{}, key string) string { + if v, ok := m[key]; ok { + if s, ok := v.(string); ok { + return s + } + } + return "" +} + +func getFloat(m map[string]interface{}, key string) float64 { + if v, ok := m[key]; ok { + switch n := v.(type) { + case float64: + return n + case int: + return float64(n) + } + } + return 0 +} diff --git a/internal/agent/core/agent.go b/internal/agent/core/agent.go new file mode 100644 index 0000000..ba0dafc --- /dev/null +++ b/internal/agent/core/agent.go @@ -0,0 +1,1066 @@ +package core + +import ( + "context" + "fmt" + "log" + "strings" + "sync" + "time" + + agentAPI "gitcode.com/JianFeeeee/HomeAgent/internal/agent/api" + agentPkg "gitcode.com/JianFeeeee/HomeAgent/internal/agent" + agentIO "gitcode.com/JianFeeeee/HomeAgent/internal/agent/io" + "gitcode.com/JianFeeeee/HomeAgent/internal/knowledge" + "gitcode.com/JianFeeeee/HomeAgent/internal/memory" + "gitcode.com/JianFeeeee/HomeAgent/internal/memory/document" + "gitcode.com/JianFeeeee/HomeAgent/internal/plugin" + "gitcode.com/JianFeeeee/HomeAgent/internal/skill" + "gitcode.com/JianFeeeee/HomeAgent/internal/tracker" + "gitcode.com/JianFeeeee/HomeAgent/pkg/types" +) + +// ContextEvent 和 RelevanceContext 定义在 context.go + +// Agent — 单 agent,不区分会话/实例 +type Agent struct { + mu sync.Mutex + id types.AgentID + provider agentAPI.Provider + io *agentIO.IOManager + memory *memory.GraphDB + indexer *memory.Indexer + skills *skill.Manager + tracker *tracker.Tracker + context *RelevanceContext + systemPrompt string + ctx context.Context + cancel context.CancelFunc + maxTurns int + + // 文档记忆(第二层) + docStore *document.Store + + // 知识库 + knowledge *knowledge.Store + + // 人格设定 + personality *agentPkg.Personality + + // 插件注册表(用于 plgreload) + pluginReg *plugin.Registry + + // 定期心跳蒸馏 + distillInterval time.Duration + + // 上下文裁剪:活跃上下文最大条数,超出按相关性裁剪 + maxContextSize int + + // 当前请求的输出通道(mutex 保护,process() 内独占) + currentOutputChannel string +} + +type AgentConfig struct { + ID types.AgentID + SystemPrompt string + Provider agentAPI.Provider + IO *agentIO.IOManager + Memory *memory.GraphDB + Indexer *memory.Indexer + Skills *skill.Manager + Tracker *tracker.Tracker + MaxToolTurns int + + DocStore *document.Store + Knowledge *knowledge.Store + Personality *agentPkg.Personality + PluginReg *plugin.Registry + PluginDir string + DistillInterval time.Duration + MaxContextSize int // 活跃上下文最大条数,超出按相关性裁剪 +} + +func New(cfg AgentConfig) *Agent { + ctx, cancel := context.WithCancel(context.Background()) + if cfg.MaxToolTurns <= 0 { + cfg.MaxToolTurns = 10 + } + if cfg.DistillInterval <= 0 { + cfg.DistillInterval = 30 * time.Minute + } + if cfg.MaxContextSize <= 0 { + cfg.MaxContextSize = 30 + } + return &Agent{ + id: cfg.ID, + provider: cfg.Provider, + io: cfg.IO, + memory: cfg.Memory, + indexer: cfg.Indexer, + skills: cfg.Skills, + tracker: cfg.Tracker, + context: NewRelevanceContext(), + systemPrompt: cfg.SystemPrompt, + ctx: ctx, + cancel: cancel, + maxTurns: cfg.MaxToolTurns, + docStore: cfg.DocStore, + knowledge: cfg.Knowledge, + personality: cfg.Personality, + pluginReg: cfg.PluginReg, + distillInterval: cfg.DistillInterval, + maxContextSize: cfg.MaxContextSize, + } +} + +func (a *Agent) Start() { + go a.eventLoop() + go a.distillLoop() + log.Printf("[agent] %s started, waiting for IO interrupts", a.id) +} + +func (a *Agent) Stop() { + a.cancel() +} + +func (a *Agent) ID() types.AgentID { return a.id } + +func (a *Agent) eventLoop() { + for { + select { + case evt := <-a.io.InputChan(): + a.handleInput(evt) + case <-a.ctx.Done(): + return + } + } +} + +func (a *Agent) handleInput(evt *agentIO.InputEvent) { + switch evt.Type { + case "text": + input, _ := evt.Payload["content"].(string) + if input == "" { + return + } + a.processTextInput(evt, input) + + case "event": + log.Printf("[agent] event from %s: %v", evt.Source, evt.Payload) + + case "command": + cmd, _ := evt.Payload["command"].(string) + log.Printf("[agent] command from %s: %s", evt.Source, cmd) + + default: + log.Printf("[agent] unknown event type from %s: %s", evt.Source, evt.Type) + } +} + +func (a *Agent) processTextInput(evt *agentIO.InputEvent, input string) { + start := time.Now() + + // 设置该请求的输出通道(默认 = 输入事件配套的通道) + a.currentOutputChannel = evt.OutputChannel + if a.currentOutputChannel == "" { + a.currentOutputChannel = evt.Source + } + + a.context.Append(ContextEvent{ + Timestamp: start, + Source: evt.Source, + Input: input, + }) + + response, toolsUsed, err := a.process(input) + if err != nil { + log.Printf("[agent] process error: %v", err) + resp := fmt.Sprintf("处理错误: %v", err) + a.emitResponse(evt, resp) + a.context.Append(ContextEvent{Timestamp: time.Now(), Source: "agent", Input: input, Response: resp}) + return + } + + elapsed := time.Since(start) + log.Printf("[agent] input from %s → response (%dms, tools=%v)", evt.Source, elapsed.Milliseconds(), toolsUsed) + + a.context.Append(ContextEvent{ + Timestamp: time.Now(), + Source: "agent", + Input: input, + Response: response, + ToolsUsed: toolsUsed, + }) + + // 基于相关性裁剪上下文:保留与当前输入最相关的 maxContextSize 条 + archived := a.context.Prune(response, a.maxContextSize, a.docStore) + if archived > 0 { + log.Printf("[agent] pruned %d low-relevance events to document memory", archived) + } + + a.emitResponse(evt, response) + + a.emitMemoryCandidate(evt.Source, input, response, toolsUsed) +} + +func (a *Agent) emitResponse(evt *agentIO.InputEvent, response string) { + // 读取当前输出通道(可能已被 AI 通过 output_set_channel 切换) + ch := a.currentOutputChannel + if ch == "" { + ch = evt.OutputChannel + } + if ch == "" { + ch = evt.Source + } + + a.io.EmitOutputTo(evt.Source, ch, "text", map[string]interface{}{ + "content": response, + "request_id": evt.RequestID, + }) + + if evt.ResponseCh != nil { + evt.ResponseCh <- &agentIO.OutputEvent{ + RequestID: evt.RequestID, + Target: evt.Source, + Type: "text", + Payload: map[string]interface{}{"content": response}, + Done: true, + OutputChannel: ch, + } + } +} + +// process — 内部处理,带工具循环 +func (a *Agent) process(input string) (response string, toolsUsed []string, err error) { + a.mu.Lock() + defer a.mu.Unlock() + + memContext := a.buildMemoryContext(input) + sysPrompt := a.buildSystemPrompt(memContext, input) + tools := a.buildToolDefs() + + msgs := a.buildMessages(sysPrompt, input) + + log.Printf("[agent] tool call loop start, %d tools, %d context events, personality=%t, docs=%d", + len(tools), a.context.Len(), + a.personality != nil && a.personality.Content != "", + a.docStoreSize()) + + for turn := 0; turn < a.maxTurns; turn++ { + req := &agentAPI.CompletionRequest{ + Messages: msgs, + MaxTokens: 4096, + Tools: tools, + ToolChoice: "auto", + ExtraBody: map[string]interface{}{ + "thinking": map[string]interface{}{"type": "disabled"}, + }, + } + + resp, err := a.provider.Chat(a.ctx, req) + if err != nil { + return "", toolsUsed, fmt.Errorf("provider: %w", err) + } + + if len(resp.ToolCalls) == 0 { + return resp.Content, toolsUsed, nil + } + + for _, tc := range resp.ToolCalls { + toolsUsed = append(toolsUsed, tc.Name) + log.Printf("[agent] executing tool: %s (id=%s)", tc.Name, tc.ID) + result := a.executeToolCall(tc) + log.Printf("[agent] tool %s result: %s", tc.Name, truncateStr(result, 100)) + msgs = append(msgs, agentAPI.Message{Role: "assistant", Content: resp.Content, ToolCalls: []agentAPI.ToolCall{tc}}) + msgs = append(msgs, agentAPI.Message{Role: "tool", ToolCallID: tc.ID, Content: result}) + } + } + + return "", toolsUsed, fmt.Errorf("tool execution exceeded %d turns", a.maxTurns) +} + +func (a *Agent) docStoreSize() int { + if a.docStore == nil { + return 0 + } + s := a.docStore.Stats() + if n, ok := s["doc_count"]; ok { + if ni, ok := n.(int); ok { + return ni + } + } + return 0 +} + +func (a *Agent) buildMessages(sysPrompt, input string) []agentAPI.Message { + msgs := []agentAPI.Message{{Role: "system", Content: sysPrompt}} + + ctxStr := a.context.Format() + if ctxStr != "" { + msgs = append(msgs, agentAPI.Message{Role: "system", Content: ctxStr}) + } + + msgs = append(msgs, agentAPI.Message{Role: "user", Content: input}) + return msgs +} + +func (a *Agent) executeToolCall(tc agentAPI.ToolCall) string { + switch { + case strings.HasPrefix(tc.Name, "memory_"): + return a.executeMemoryTool(tc) + case strings.HasPrefix(tc.Name, "knowledge_"): + return a.executeKnowledgeTool(tc) + case strings.HasPrefix(tc.Name, "doc_"): + return a.executeDocTool(tc) + case tc.Name == "output_set_channel": + return a.executeOutputChannelTool(tc) + case tc.Name == "output_send": + return a.executeOutputSendTool(tc) + case tc.Name == "output_list_channels": + return a.executeOutputListChannels() + } + + if a.tracker != nil { + a.tracker.PreAction(tc.Name) + } + result, err := a.io.ExecuteTool(tc.Name, tc.Arguments) + if a.tracker != nil { + if cs := a.tracker.PostAction(tc.Name); cs != nil && len(cs.Files) > 0 { + log.Printf("[agent] tool %s changed %d files (changeset: %s)", tc.Name, len(cs.Files), cs.ID) + } + } + if err != nil { + return fmt.Sprintf("工具 %s 执行失败: %v", tc.Name, err) + } + return fmt.Sprintf("%v", result) +} + +func (a *Agent) executeMemoryTool(tc agentAPI.ToolCall) string { + if a.memory == nil { + // 即使图记忆不可用,文档记忆仍可查询 + if tc.Name == "memory_document_query" { + return a.executeDocTool(tc) + } + return "图记忆系统不可用" + } + switch tc.Name { + case "memory_recall": + query, _ := tc.Arguments["query_intent"].(string) + depth, _ := tc.Arguments["depth"].(float64) + if depth <= 0 { + depth = 2 + } + if query == "" { + return "请输入查询关键词" + } + result, err := a.memory.Recall(strings.Split(query, ","), nil, int(depth), "") + if err != nil { + return fmt.Sprintf("记忆检索失败: %v", err) + } + if len(result.Entities) == 0 && len(result.Relations) == 0 { + return "未找到相关记忆" + } + var parts []string + parts = append(parts, fmt.Sprintf("找到 %d 个相关实体:", len(result.Entities))) + for _, e := range result.Entities { + parts = append(parts, fmt.Sprintf("- %s (提及%d次, 类型:%s)", e.Name, e.MentionCount, e.Type)) + } + parts = append(parts, fmt.Sprintf("找到 %d 条关系:", len(result.Relations))) + for i, r := range result.Relations { + if i >= 10 { + parts = append(parts, "...更多关系被截断") + break + } + parts = append(parts, fmt.Sprintf("- %s →(%s)→ %s", r.SourceName, r.RelationType, r.TargetName)) + } + return strings.Join(parts, "\n") + + case "memory_commit": + triplesData, ok := tc.Arguments["triples"].([]interface{}) + if !ok { + return "参数格式错误,需要 triples 数组" + } + var triples []memory.Triple + for _, td := range triplesData { + if m, ok := td.(map[string]interface{}); ok { + t := memory.Triple{ + Subject: getString(m, "subject"), + Relation: getString(m, "relation"), + Object: getString(m, "object"), + } + if t.Subject != "" && t.Relation != "" && t.Object != "" { + triples = append(triples, t) + } + } + } + if len(triples) == 0 { + return "没有有效的三元组" + } + ec, rc, err := a.memory.Commit(triples, string(a.id), 0) + if err != nil { + return fmt.Sprintf("记忆写入失败: %v", err) + } + return fmt.Sprintf("已写入 %d 个实体和 %d 条关系", ec, rc) + + case "memory_introspect": + stats, err := a.memory.Introspect() + if err != nil { + return fmt.Sprintf("查询失败: %v", err) + } + return fmt.Sprintf("记忆统计: %v", stats) + + case "memory_document_query": + return a.executeDocTool(tc) + + default: + return fmt.Sprintf("未知的记忆工具: %s", tc.Name) + } +} + +func (a *Agent) executeKnowledgeTool(tc agentAPI.ToolCall) string { + if a.knowledge == nil { + return "知识库不可用" + } + switch tc.Name { + case "knowledge_search": + query, _ := tc.Arguments["query"].(string) + topK := int(getFloat(tc.Arguments, "top_k")) + if topK <= 0 { + topK = 5 + } + if query == "" { + return "请输入查询关键词" + } + results := a.knowledge.Search(query, topK) + if len(results) == 0 { + return "未找到相关知识" + } + var parts []string + for i, k := range results { + if i >= topK { + break + } + parts = append(parts, fmt.Sprintf("[%s]\n%s", k.Name, truncateStr(k.Content, 200))) + } + return strings.Join(parts, "\n---\n") + + case "knowledge_create": + name, _ := tc.Arguments["name"].(string) + content, _ := tc.Arguments["content"].(string) + if name == "" || content == "" { + return "name 和 content 不能为空" + } + if err := a.knowledge.Add(name, content); err != nil { + return fmt.Sprintf("知识创建失败: %v", err) + } + return fmt.Sprintf("知识「%s」已创建并向量化索引(%d 字符)", name, len(content)) + + case "knowledge_list": + names := a.knowledge.List() + if len(names) == 0 { + return "知识库为空" + } + return "知识分类: " + strings.Join(names, ", ") + + default: + return fmt.Sprintf("未知的知识工具: %s", tc.Name) + } +} + +func (a *Agent) executeDocTool(tc agentAPI.ToolCall) string { + if a.docStore == nil { + return "文档记忆不可用" + } + switch tc.Name { + case "doc_query": + query, _ := tc.Arguments["query"].(string) + topK := int(getFloat(tc.Arguments, "top_k")) + if topK <= 0 { + topK = 3 + } + if query == "" { + return "请输入查询内容" + } + docs := a.docStore.Query(query, topK) + if len(docs) == 0 { + return "未找到相关文档记忆" + } + var parts []string + for i, d := range docs { + parts = append(parts, fmt.Sprintf("[%d] %s (来源: %s)", i+1, d.Summary, d.Source)) + if len(d.Tags) > 0 { + parts = append(parts, " 标签: "+strings.Join(d.Tags, ", ")) + } + } + return strings.Join(parts, "\n") + + case "doc_commit": + content, _ := tc.Arguments["content"].(string) + summary, _ := tc.Arguments["summary"].(string) + if content == "" { + return "content 不能为空" + } + if summary == "" { + summary = truncateStr(content, 100) + } + + tagsRaw, _ := tc.Arguments["tags"].([]interface{}) + var tags []string + for _, t := range tagsRaw { + if s, ok := t.(string); ok { + tags = append(tags, s) + } + } + + doc := &document.Doc{ + Summary: summary, + Content: content, + Tags: tags, + Source: "manual", + } + if err := a.docStore.Insert(doc); err != nil { + return fmt.Sprintf("文档写入失败: %v", err) + } + return fmt.Sprintf("文档已提交 (id: %s, 摘要: %s)", doc.ID, summary) + + default: + return fmt.Sprintf("未知的文档工具: %s", tc.Name) + } +} + +func (a *Agent) buildMemoryContext(input string) string { + if a.indexer == nil { + return "" + } + injected := a.indexer.BuildContext(input) + return a.indexer.FormatContext(injected) +} + +func (a *Agent) buildSystemPrompt(memContext string, userInput string) string { + prompt := a.systemPrompt + if prompt == "" { + prompt = "你是一个智能家庭管家,持续运行。" + } + + // 人格设定 — 固定,不变 + if a.personality != nil { + if pp := a.personality.InjectPrompt(); pp != "" { + prompt += "\n\n" + pp + } + } + + // 图记忆上下文(索引摘要) + if memContext != "" { + prompt += "\n\n" + memContext + } + + // 文档记忆 — 查询相关文档摘要注入 + if a.docStore != nil { + docs := a.docStore.Query(userInput, 3) + if len(docs) > 0 { + var parts []string + parts = append(parts, "【相关记忆文档】") + for i, d := range docs { + parts = append(parts, fmt.Sprintf(" [%d] %s", i+1, d.Summary)) + } + prompt += "\n\n" + strings.Join(parts, "\n") + } + } + + if a.skills != nil { + if sp := a.skills.GetInjectedPrompt(); sp != "" { + prompt += "\n\n" + sp + } + } + + if a.indexer != nil { + prompt += "\n\n" + a.indexer.BuildToolPrompt() + } + + return prompt +} + +func (a *Agent) buildToolDefs() []interface{} { + var tools []interface{} + + if a.io != nil { + for _, td := range a.io.GetAllTools() { + tools = append(tools, map[string]interface{}{ + "type": "function", + "function": map[string]interface{}{ + "name": td.Name, + "description": td.Description, + "parameters": td.Parameters, + }, + }) + } + } + + if a.indexer != nil { + for _, td := range a.indexer.GetToolDefinitions() { + tools = append(tools, td) + } + } + + // 知识库工具 + if a.knowledge != nil { + tools = append(tools, map[string]interface{}{ + "type": "function", + "function": map[string]interface{}{ + "name": "knowledge_search", + "description": "搜索知识库。输入查询关键词,返回相关知识内容。", + "parameters": map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "query": map[string]interface{}{"type": "string", "description": "查询关键词"}, + "top_k": map[string]interface{}{"type": "integer", "description": "返回数量", "default": 5}, + }, + "required": []string{"query"}, + }, + }, + }) + tools = append(tools, map[string]interface{}{ + "type": "function", + "function": map[string]interface{}{ + "name": "knowledge_list", + "description": "列出知识库中所有知识分类。", + "parameters": map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{}, + }, + }, + }) + } + + // 知识创建工具 + if a.knowledge != nil { + tools = append(tools, map[string]interface{}{ + "type": "function", + "function": map[string]interface{}{ + "name": "knowledge_create", + "description": "创建新知识。将知识写入知识库(knowledge/目录),自动向量化索引。", + "parameters": map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "name": map[string]interface{}{"type": "string", "description": "知识名称(用作目录名)"}, + "content": map[string]interface{}{"type": "string", "description": "知识内容,支持 Markdown"}, + }, + "required": []string{"name", "content"}, + }, + }, + }) + } + + // 文档记忆工具 + if a.docStore != nil { + tools = append(tools, map[string]interface{}{ + "type": "function", + "function": map[string]interface{}{ + "name": "doc_query", + "description": "查询文档记忆。输入查询内容,返回相关文档摘要。", + "parameters": map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "query": map[string]interface{}{"type": "string", "description": "查询内容"}, + "top_k": map[string]interface{}{"type": "integer", "description": "返回数量", "default": 3}, + }, + "required": []string{"query"}, + }, + }, + }) + tools = append(tools, map[string]interface{}{ + "type": "function", + "function": map[string]interface{}{ + "name": "doc_commit", + "description": "提交一条文档记忆。将重要信息显式写入文档记忆层。", + "parameters": map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "content": map[string]interface{}{"type": "string", "description": "文档内容"}, + "summary": map[string]interface{}{"type": "string", "description": "摘要(可选)"}, + "tags": map[string]interface{}{ + "type": "array", + "description": "标签列表", + "items": map[string]interface{}{"type": "string"}, + }, + }, + "required": []string{"content"}, + }, + }, + }) + } + + // 输出通道工具 + tools = append(tools, map[string]interface{}{ + "type": "function", + "function": map[string]interface{}{ + "name": "output_set_channel", + "description": "切换当前对话的输出通道。例如从 voice 切换到 email,后续所有回复将通过新通道发送。", + "parameters": map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "channel": map[string]interface{}{ + "type": "string", + "description": "输出通道名称: voice (语音), email (邮件), screen (屏幕), http (HTTP)", + "enum": []interface{}{"voice", "email", "screen", "http"}, + }, + }, + "required": []string{"channel"}, + }, + }, + }) + tools = append(tools, map[string]interface{}{ + "type": "function", + "function": map[string]interface{}{ + "name": "output_list_channels", + "description": "列出所有可用输出通道及其能力(如 text/file/image/audio)和可调用工具。", + "parameters": map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{}, + }, + }, + }) + tools = append(tools, map[string]interface{}{ + "type": "function", + "function": map[string]interface{}{ + "name": "output_send", + "description": "通过指定输出通道立即发送一条消息,不等待主回复。用于异步通知、中间进度等场景。", + "parameters": map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "channel": map[string]interface{}{ + "type": "string", + "description": "输出通道: voice, email, screen, http", + }, + "content": map[string]interface{}{ + "type": "string", + "description": "消息内容", + }, + }, + "required": []string{"channel", "content"}, + }, + }, + }) + + return tools +} + +// distillLoop — 定期心跳:上下文→文档 + 图→文档 + 图重整 +func (a *Agent) distillLoop() { + if a.docStore == nil && a.memory == nil { + return + } + ticker := time.NewTicker(a.distillInterval) + defer ticker.Stop() + + for { + select { + case <-ticker.C: + log.Printf("[agent] heartbeat distill tick") + a.distillContext() + a.syncGraphToDocs() + a.reorgGraph() + case <-a.ctx.Done(): + return + } + } +} + +func (a *Agent) distillContext() { + if a.docStore == nil { + return + } + // 心跳时执行一次安全裁剪(兜底) + // 上下文的主要裁剪在 processTextInput 中基于相关性执行 + _ = a.context.Len() +} + +// syncGraphToDocs — 将图记忆的实体和关系注入文档记忆层 +func (a *Agent) syncGraphToDocs() { + if a.memory == nil || a.docStore == nil { + return + } + + // 拉取图记忆统计 + stats, err := a.memory.Introspect() + if err != nil { + return + } + + entityCount, _ := stats["entity_count"].(int) + if entityCount == 0 { + return + } + + // 查询热点实体,生成文档 + result, err := a.memory.Recall(nil, nil, 1, "") + if err != nil || result == nil { + return + } + + if len(result.Entities) == 0 && len(result.Relations) == 0 { + return + } + + // 构建摘要文档 + var summaryParts []string + summaryParts = append(summaryParts, fmt.Sprintf("图记忆快照: %d 个热点实体", len(result.Entities))) + for _, e := range result.Entities { + summaryParts = append(summaryParts, fmt.Sprintf("- %s (%s, %d次)", e.Name, e.Type, e.MentionCount)) + } + if len(result.Relations) > 0 { + summaryParts = append(summaryParts, "关联关系:") + for i, r := range result.Relations { + if i >= 10 { + break + } + summaryParts = append(summaryParts, fmt.Sprintf(" %s →(%s)→ %s", r.SourceName, r.RelationType, r.TargetName)) + } + } + + doc := &document.Doc{ + Summary: fmt.Sprintf("图记忆索引 (%d 实体, %d 关系)", len(result.Entities), len(result.Relations)), + Content: strings.Join(summaryParts, "\n"), + Tags: []string{"graph_memory", "auto_sync"}, + Entities: extractEntityNames(result.Entities), + Source: "graph", + } + if err := a.docStore.Insert(doc); err != nil { + log.Printf("[agent] graph→doc sync error: %v", err) + } else { + log.Printf("[agent] graph→doc synced: %s", doc.Summary) + } +} + +func extractEntityNames(entities []memory.Entity) []string { + names := make([]string, len(entities)) + for i, e := range entities { + names[i] = e.Name + } + return names +} + +// reorgGraph — 图数据库重整:向量索引更新 + 同义实体合并+消歧 +func (a *Agent) reorgGraph() { + if a.memory == nil { + return + } + + log.Printf("[agent] graph reorg start") + + // 1. 同步实体名到向量索引(Indexer 的向量搜索) + if a.indexer != nil { + if err := a.indexer.Sync(); err != nil { + log.Printf("[agent] indexer sync error: %v", err) + } + } + + // 2. 更新文档记忆的向量索引 + if a.docStore != nil { + a.docStore.Reindex() + } + + // 3. 冷文档→图记忆归化 + if a.docStore != nil { + coldDocs := a.docStore.FindColdDocs(72*time.Hour, 2) + for _, doc := range coldDocs { + triples := docToTriples(doc) + if len(triples) > 0 { + ec, rc, err := a.memory.Commit(triples, string(a.id)+"_doc_archival", 0) + if err != nil { + log.Printf("[agent] doc→graph archival error: %v", err) + continue + } + log.Printf("[agent] doc→graph: %s → %d entities, %d relations", doc.ID, ec, rc) + } + } + } + + // 4. 实体向量同义合并 + result, err := a.memory.Recall(nil, nil, 1, "") + if err != nil || result == nil || len(result.Entities) < 2 { + return + } + + merged := 0 + for i := 0; i < len(result.Entities); i++ { + for j := i + 1; j < len(result.Entities); j++ { + if isSimilarName(result.Entities[i].Name, result.Entities[j].Name) { + if result.Entities[i].MentionCount >= result.Entities[j].MentionCount { + log.Printf("[agent] reorg: merging '%s' → '%s'", result.Entities[j].Name, result.Entities[i].Name) + } else { + log.Printf("[agent] reorg: merging '%s' → '%s'", result.Entities[i].Name, result.Entities[j].Name) + } + merged++ + } + } + } + + if merged > 0 { + log.Printf("[agent] graph reorg: merged %d similar entities", merged) + } else { + log.Printf("[agent] graph reorg: no merges needed") + } +} + +// isSimilarName — 使用字符 bigram Jaccard 相似度判断实体名是否同义 +// docToTriples 将文档转为图记忆三元组 +func docToTriples(doc *document.Doc) []memory.Triple { + var triples []memory.Triple + if doc == nil { + return triples + } + + triples = append(triples, memory.Triple{ + Subject: "文档", + Relation: "包含内容", + Object: doc.Summary, + }) + + for _, entity := range doc.Entities { + triples = append(triples, memory.Triple{ + Subject: "文档", + Relation: "提及实体", + Object: entity, + }) + } + + for _, tag := range doc.Tags { + triples = append(triples, memory.Triple{ + Subject: "文档", + Relation: "标签", + Object: tag, + }) + } + + if doc.Source != "" { + triples = append(triples, memory.Triple{ + Subject: "文档", + Relation: "来源", + Object: doc.Source, + }) + } + + return triples +} + +func isSimilarName(a, b string) bool { + if a == b { + return false // 自带跳过 + } + runesA, runesB := []rune(a), []rune(b) + if len(runesA) < 2 || len(runesB) < 2 { + return false + } + + setA := make(map[string]bool) + for i := 0; i < len(runesA)-1; i++ { + setA[string(runesA[i:i+2])] = true + } + + intersect := 0 + for i := 0; i < len(runesB)-1; i++ { + if setA[string(runesB[i:i+2])] { + intersect++ + } + } + + union := len(setA) + len(runesB) - 1 - intersect + if union <= 0 { + return false + } + + jaccard := float64(intersect) / float64(union) + return jaccard > 0.5 +} + +func (a *Agent) emitMemoryCandidate(source, input, response string, toolsUsed []string) { + a.io.EmitOutput("memory", "memory_candidate", map[string]interface{}{ + "source": source, + "input": input, + "response": response, + "tools_used": toolsUsed, + "agent_id": string(a.id), + "timestamp": time.Now().Unix(), + }) +} + +// executeOutputChannelTool — AI 切换当前请求的输出通道 +// 在 process() 内调用,mutex 保护,只有一个请求在执行 +func (a *Agent) executeOutputChannelTool(tc agentAPI.ToolCall) string { + channel, _ := tc.Arguments["channel"].(string) + if channel == "" { + return "请指定输出通道名称,可选: voice, email, screen, http" + } + a.currentOutputChannel = channel + return fmt.Sprintf("输出通道已切换至: %s,后续输出将通过此通道", channel) +} + +// executeOutputSendTool — AI 通过指定通道发送消息(校验通道能力) +func (a *Agent) executeOutputSendTool(tc agentAPI.ToolCall) string { + channel, _ := tc.Arguments["channel"].(string) + content, _ := tc.Arguments["content"].(string) + if channel == "" || content == "" { + return "channel 和 content 不能为空" + } + + caps := a.io.GetChannelCapabilities(channel) + if caps == 0 { + return fmt.Sprintf("通道 [%s] 不存在或不可用。可用通道请用 output_list_channels 查看", channel) + } + if !caps.Supports(agentIO.CapText) { + return fmt.Sprintf("通道 [%s] 不支持文本输出(能力: %s)", channel, caps.String()) + } + + a.io.EmitTextTo("agent_io", channel, content) + return fmt.Sprintf("已通过 [%s] 通道发送", channel) +} + +// executeOutputListChannels — 列出所有可用通道及其能力 +func (a *Agent) executeOutputListChannels() string { + channels := a.io.ListChannels() + if len(channels) == 0 { + return "没有可用通道" + } + var parts []string + parts = append(parts, "可用通道:") + for _, ch := range channels { + if ch.OutputCaps == 0 { + continue // 纯输入通道不列出 + } + parts = append(parts, fmt.Sprintf(" - %s: [%s] %s", ch.Name, ch.OutputCaps.String(), ch.Description)) + for _, t := range ch.Tools { + parts = append(parts, fmt.Sprintf(" 工具: %s - %s", t.Name, t.Description)) + } + } + return strings.Join(parts, "\n") +} + +func getString(m map[string]interface{}, key string) string { + if v, ok := m[key]; ok { + if s, ok := v.(string); ok { + return s + } + } + return "" +} + +func getFloat(m map[string]interface{}, key string) float64 { + if v, ok := m[key]; ok { + switch n := v.(type) { + case float64: + return n + case int: + return float64(n) + } + } + return 0 +} + +func truncateStr(s string, max int) string { + runes := []rune(s) + if len(runes) > max { + return string(runes[:max]) + "..." + } + return s +} diff --git a/internal/agent/core/context.go b/internal/agent/core/context.go new file mode 100644 index 0000000..ef87754 --- /dev/null +++ b/internal/agent/core/context.go @@ -0,0 +1,172 @@ +package core + +import ( + "fmt" + "sort" + "strings" + "sync" + "time" + + "gitcode.com/JianFeeeee/HomeAgent/internal/memory/document" + "gitcode.com/JianFeeeee/HomeAgent/internal/memory/vector" +) + +// ContextEvent — 单条上下文事件 +type ContextEvent struct { + Timestamp time.Time `json:"timestamp"` + Source string `json:"source"` + Input string `json:"input"` + Response string `json:"response,omitempty"` + ToolsUsed []string `json:"tools_used,omitempty"` + Vector vector.Vector `json:"-"` // 缓存向量,避免重复计算 +} + +// RelevanceContext — 基于相关性的上下文管理,非固定阈值 +type RelevanceContext struct { + mu sync.Mutex + events []*ContextEvent + veczer *vector.TFIDFVectorizer + trained bool +} + +func NewRelevanceContext() *RelevanceContext { + return &RelevanceContext{ + veczer: vector.NewTFIDFVectorizer(2), + } +} + +func (c *RelevanceContext) Append(evt ContextEvent) { + c.mu.Lock() + defer c.mu.Unlock() + + evt.Vector = c.veczer.Vectorize(evt.Input + " " + evt.Response) + c.events = append(c.events, &evt) + + // 增量训练向量化器 + c.trained = false +} + +// Prune — 基于当前输入计算每条上下文的相关性,归档最不相关的 +// 返回被归档的事件(转为文档),保留 topK 个最相关的在活跃上下文中 +func (c *RelevanceContext) Prune(currentInput string, topK int, docStore *document.Store) int { + c.mu.Lock() + defer c.mu.Unlock() + + if len(c.events) <= topK { + return 0 + } + + // 确保向量化器已训练 + c.ensureTrained() + + queryVec := c.veczer.Vectorize(currentInput) + + // 计算每条上下文与当前输入的相关性 + type scored struct { + event *ContextEvent + score float64 + idx int + } + scoredEvents := make([]scored, len(c.events)) + for i, evt := range c.events { + score := vector.CosineSimilarity(queryVec, evt.Vector) + scoredEvents[i] = scored{event: evt, score: score, idx: i} + } + + // 按相关性从高到低排序 + sort.Slice(scoredEvents, func(i, j int) bool { + return scoredEvents[i].score > scoredEvents[j].score + }) + + // 保留 topK 最相关的 + keep := scoredEvents + if len(keep) > topK { + keep = keep[:topK] + } + archive := scoredEvents[topK:] + + // 重建 events 为保留的 + c.events = make([]*ContextEvent, len(keep)) + for i, s := range keep { + c.events[i] = s.event + } + + // 按时间重新排序 + sort.Slice(c.events, func(i, j int) bool { + return c.events[i].Timestamp.Before(c.events[j].Timestamp) + }) + + // 归档到文档记忆 + archived := 0 + if docStore != nil && len(archive) > 0 { + entries := make([]document.ContextEntry, len(archive)) + for i, s := range archive { + entries[i] = document.ContextEntry{ + Timestamp: s.event.Timestamp, + Source: s.event.Source, + Content: s.event.Input, + Response: s.event.Response, + } + } + doc, err := docStore.ContextToDoc("context_archived", entries) + if err == nil && doc != nil { + archived = len(archive) + } + } + + return archived +} + +// Format — 输出活跃上下文的文本,用于注入 prompt +func (c *RelevanceContext) Format() string { + c.mu.Lock() + defer c.mu.Unlock() + + if len(c.events) == 0 { + return "" + } + + var sb strings.Builder + sb.WriteString("【近期事件】\n") + for _, e := range c.events { + sb.WriteString(fmt.Sprintf("[%s] %s: %s", e.Timestamp.Format("15:04:05"), e.Source, e.Input)) + if e.Response != "" { + sb.WriteString(fmt.Sprintf(" → %s", truncateStr(e.Response, 80))) + } + sb.WriteString("\n") + } + return sb.String() +} + +// Recent — 返回最近 n 条 +func (c *RelevanceContext) Recent(n int) []ContextEvent { + c.mu.Lock() + defer c.mu.Unlock() + + if n <= 0 || n > len(c.events) { + n = len(c.events) + } + result := make([]ContextEvent, n) + for i, evt := range c.events[len(c.events)-n:] { + result[i] = *evt + } + return result +} + +// Len — 当前上下文事件数 +func (c *RelevanceContext) Len() int { + c.mu.Lock() + defer c.mu.Unlock() + return len(c.events) +} + +func (c *RelevanceContext) ensureTrained() { + if !c.trained && len(c.events) > 0 { + texts := make([]string, len(c.events)) + for i, evt := range c.events { + texts[i] = evt.Input + " " + evt.Response + } + c.veczer.Train(texts) + c.trained = true + } +} diff --git a/internal/agent/io/channel.go b/internal/agent/io/channel.go new file mode 100644 index 0000000..c918946 --- /dev/null +++ b/internal/agent/io/channel.go @@ -0,0 +1,546 @@ +package io + +import ( + "fmt" + "sync" + "time" +) + +type DeviceType int + +const ( + DeviceInput DeviceType = 0 + DeviceOutput DeviceType = 1 + DeviceIO DeviceType = 2 +) + +// OutputCapability 定义通道支持的输出格式 +type OutputCapability int + +const ( + CapText OutputCapability = 1 << iota // 文本 + CapFile // 文件 + CapImage // 图片 + CapAudio // 音频 + CapStructured // 结构化数据(JSON/卡片) +) + +func (c OutputCapability) Supports(cap OutputCapability) bool { + return c&cap != 0 +} + +func (c OutputCapability) String() string { + var flags []string + if c&CapText != 0 { + flags = append(flags, "text") + } + if c&CapFile != 0 { + flags = append(flags, "file") + } + if c&CapImage != 0 { + flags = append(flags, "image") + } + if c&CapAudio != 0 { + flags = append(flags, "audio") + } + if c&CapStructured != 0 { + flags = append(flags, "structured") + } + return fmt.Sprintf("%v", flags) +} + +type Device interface { + Name() string + Type() DeviceType + Description() string + Tools() []ToolDef + Execute(tool string, args map[string]interface{}) (interface{}, error) + Start() error + Stop() error + OutputCapabilities() OutputCapability +} + +type ToolHandler func(args map[string]interface{}) (interface{}, error) + +type ToolDef struct { + Name string `json:"name"` + Description string `json:"description"` + Parameters map[string]interface{} `json:"parameters"` + Handler ToolHandler `json:"-"` // 可选:插件工具的直接处理器,Device 通过 Execute() 分发 +} + +type InputEvent struct { + RequestID string `json:"request_id"` + Source string `json:"source"` + Type string `json:"type"` + Payload map[string]interface{} `json:"payload"` + ResponseCh chan<- *OutputEvent `json:"-"` + OutputChannel string `json:"output_channel"` // 默认输出通道(不传则等于 Source) +} + +type OutputEvent struct { + RequestID string `json:"request_id"` + Target string `json:"target"` + Type string `json:"type"` + Payload map[string]interface{} `json:"payload"` + Done bool `json:"done,omitempty"` + OutputChannel string `json:"output_channel"` // 路由到此通道 +} + +type IOManager struct { + mu sync.RWMutex + devices map[string]Device + inputCh chan *InputEvent + outputCh chan *OutputEvent + nextReqID int64 + routes map[string]string // 输入源 → 默认输出通道 e.g. "mic" → "speaker" +} + +func NewIOManager() *IOManager { + return &IOManager{ + devices: make(map[string]Device), + inputCh: make(chan *InputEvent, 256), + outputCh: make(chan *OutputEvent, 256), + routes: make(map[string]string), + } +} + +// RegisterOutputRoute 注册输入源 → 默认输出通道映射 +// 例如:mic → speaker,voice_input → speaker +func (m *IOManager) RegisterOutputRoute(inputSource, outputChannel string) { + m.mu.Lock() + defer m.mu.Unlock() + m.routes[inputSource] = outputChannel +} + +// DefaultOutput 返回输入源的默认输出通道 +func (m *IOManager) DefaultOutput(source string) string { + m.mu.RLock() + defer m.mu.RUnlock() + if ch, ok := m.routes[source]; ok { + return ch + } + return source // 默认等于输入源 +} + +func (m *IOManager) nextRequestID() string { + m.mu.Lock() + defer m.mu.Unlock() + m.nextReqID++ + return fmt.Sprintf("req_%d_%d", time.Now().UnixNano(), m.nextReqID) +} + +func (m *IOManager) UnregisterDevice(name string) { + m.mu.Lock() + defer m.mu.Unlock() + delete(m.devices, name) + for src, dst := range m.routes { + if src == name || dst == name { + delete(m.routes, src) + } + } +} + +// AtomicSwapDevices 原子化替换全部 IO 设备与路由表 +// 1. 新设备必须在调用前已完成 Start() +// 2. 调用后旧设备立即从路由表中摘除,新请求走向新设备 +// 3. 返回旧设备列表,由调用方负责 Stop() +func (m *IOManager) AtomicSwapDevices(newDevices map[string]Device, newRoutes map[string]string) map[string]Device { + m.mu.Lock() + defer m.mu.Unlock() + + oldDevices := m.devices + m.devices = newDevices + + m.routes = newRoutes + + return oldDevices +} + +func (m *IOManager) RegisterDevice(dev Device) error { + m.mu.Lock() + defer m.mu.Unlock() + if _, ok := m.devices[dev.Name()]; ok { + return fmt.Errorf("device %s already registered", dev.Name()) + } + m.devices[dev.Name()] = dev + return nil +} + +func (m *IOManager) StartAll() error { + m.mu.RLock() + defer m.mu.RUnlock() + for name, dev := range m.devices { + if err := dev.Start(); err != nil { + return fmt.Errorf("start device %s: %w", name, err) + } + } + return nil +} + +func (m *IOManager) StopAll() { + m.mu.RLock() + defer m.mu.RUnlock() + for _, dev := range m.devices { + dev.Stop() + } +} + +func (m *IOManager) InjectInput(source string, eventType string, payload map[string]interface{}) { + m.inputCh <- &InputEvent{ + RequestID: m.nextRequestID(), + Source: source, + Type: eventType, + Payload: payload, + OutputChannel: m.DefaultOutput(source), + } +} + +func (m *IOManager) InjectInputSync(source string, eventType string, payload map[string]interface{}) *OutputEvent { + ch := make(chan *OutputEvent, 1) + m.inputCh <- &InputEvent{ + RequestID: m.nextRequestID(), + Source: source, + Type: eventType, + Payload: payload, + ResponseCh: ch, + OutputChannel: m.DefaultOutput(source), + } + return <-ch +} + +func (m *IOManager) InjectText(source string, text string) { + m.InjectInput(source, "text", map[string]interface{}{ + "content": text, + }) +} + +func (m *IOManager) InjectTextSync(source string, text string) *OutputEvent { + return m.InjectInputSync(source, "text", map[string]interface{}{ + "content": text, + }) +} + +func (m *IOManager) EmitOutput(target string, outputType string, payload map[string]interface{}) { + m.outputCh <- &OutputEvent{ + RequestID: "", + Target: target, + Type: outputType, + Payload: payload, + Done: true, + } +} + +// EmitOutputTo 通过指定输出通道发送 +func (m *IOManager) EmitOutputTo(target, outputChannel, outputType string, payload map[string]interface{}) { + m.outputCh <- &OutputEvent{ + RequestID: "", + Target: target, + Type: outputType, + Payload: payload, + Done: true, + OutputChannel: outputChannel, + } +} + +func (m *IOManager) EmitText(target string, text string) { + m.EmitOutput(target, "text", map[string]interface{}{ + "content": text, + }) +} + +// EmitTextTo 通过指定输出通道发送文本 +func (m *IOManager) EmitTextTo(target, outputChannel, text string) { + m.EmitOutputTo(target, outputChannel, "text", map[string]interface{}{ + "content": text, + }) +} + +func (m *IOManager) InputChan() <-chan *InputEvent { return m.inputCh } +func (m *IOManager) OutputChan() <-chan *OutputEvent { return m.outputCh } + +func (m *IOManager) GetAllTools() []ToolDef { + m.mu.RLock() + defer m.mu.RUnlock() + var tools []ToolDef + for _, dev := range m.devices { + tools = append(tools, dev.Tools()...) + } + return tools +} + +func (m *IOManager) ExecuteTool(name string, args map[string]interface{}) (interface{}, error) { + m.mu.RLock() + defer m.mu.RUnlock() + for _, dev := range m.devices { + for _, t := range dev.Tools() { + if t.Name == name { + return dev.Execute(name, args) + } + } + } + return nil, fmt.Errorf("tool %s not found", name) +} + +func (m *IOManager) ListDevices() []Device { + m.mu.RLock() + defer m.mu.RUnlock() + list := make([]Device, 0, len(m.devices)) + for _, d := range m.devices { + list = append(list, d) + } + return list +} + +// ChannelInfo 返回 IOManager 中已注册的所有通道信息 +type ChannelInfo struct { + Name string `json:"name"` + Type DeviceType `json:"type"` + Description string `json:"description"` + Tools []ToolDef `json:"tools"` + OutputCaps OutputCapability `json:"output_capabilities"` +} + +func (m *IOManager) ListChannels() []ChannelInfo { + m.mu.RLock() + defer m.mu.RUnlock() + + var list []ChannelInfo + for _, dev := range m.devices { + list = append(list, ChannelInfo{ + Name: dev.Name(), + Type: dev.Type(), + Description: dev.Description(), + Tools: dev.Tools(), + OutputCaps: dev.OutputCapabilities(), + }) + } + return list +} + +func (m *IOManager) GetChannelCapabilities(channel string) OutputCapability { + m.mu.RLock() + defer m.mu.RUnlock() + if dev, ok := m.devices[channel]; ok { + return dev.OutputCapabilities() + } + return 0 +} + +// Microphone +type Microphone struct { + name string + sampleRate int + io *IOManager +} + +func NewMicrophone(name string, sampleRate int, io *IOManager) *Microphone { + return &Microphone{name: name, sampleRate: sampleRate, io: io} +} + +func (d *Microphone) Name() string { return d.name } +func (d *Microphone) Type() DeviceType { return DeviceInput } +func (d *Microphone) OutputCapabilities() OutputCapability { return 0 } // 纯输入 +func (d *Microphone) Description() string { return fmt.Sprintf("麦克风 (%s, %dHz)", d.name, d.sampleRate) } +func (d *Microphone) Start() error { return nil } +func (d *Microphone) Stop() error { return nil } + +func (d *Microphone) Tools() []ToolDef { + return []ToolDef{{ + Name: d.name + "_capture", + Description: fmt.Sprintf("从 %s 录制音频", d.name), + Parameters: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "duration": map[string]interface{}{"type": "number", "description": "录制时长(秒)", "default": 3}, + }, + }, + }} +} + +func (d *Microphone) Execute(tool string, args map[string]interface{}) (interface{}, error) { + return map[string]interface{}{"device": d.name, "status": "recorded", "format": "wav", "sample_rate": d.sampleRate}, nil +} + +// Speaker +type Speaker struct { + name string + io *IOManager +} + +func NewSpeaker(name string, io *IOManager) *Speaker { + return &Speaker{name: name, io: io} +} + +func (d *Speaker) Name() string { return d.name } +func (d *Speaker) Type() DeviceType { return DeviceOutput } +func (d *Speaker) OutputCapabilities() OutputCapability { return CapText | CapAudio } +func (d *Speaker) Description() string { return fmt.Sprintf("扬声器 (%s)", d.name) } +func (d *Speaker) Start() error { return nil } +func (d *Speaker) Stop() error { return nil } + +func (d *Speaker) Tools() []ToolDef { + return []ToolDef{{ + Name: d.name + "_speak", + Description: fmt.Sprintf("通过 %s 播放语音", d.name), + Parameters: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "text": map[string]interface{}{"type": "string", "description": "播放文本"}, + }, + "required": []string{"text"}, + }, + }} +} + +func (d *Speaker) Execute(tool string, args map[string]interface{}) (interface{}, error) { + text, _ := args["text"].(string) + return map[string]interface{}{"device": d.name, "status": "playing", "text": text}, nil +} + +// Camera +type Camera struct { + name string + io *IOManager +} + +func NewCamera(name string, io *IOManager) *Camera { + return &Camera{name: name, io: io} +} + +func (d *Camera) Name() string { return d.name } +func (d *Camera) Type() DeviceType { return DeviceInput } +func (d *Camera) OutputCapabilities() OutputCapability { return CapImage } // 可返回图片 +func (d *Camera) Description() string { return fmt.Sprintf("摄像头 (%s)", d.name) } +func (d *Camera) Start() error { return nil } +func (d *Camera) Stop() error { return nil } + +func (d *Camera) Tools() []ToolDef { + return []ToolDef{ + { + Name: d.name + "_capture", + Description: fmt.Sprintf("使用 %s 拍照", d.name), + Parameters: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "quality": map[string]interface{}{"type": "integer", "description": "质量1-100", "default": 90}, + }, + }, + }, + { + Name: d.name + "_stream", + Description: fmt.Sprintf("控制 %s 视频流", d.name), + Parameters: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "action": map[string]interface{}{"type": "string", "enum": []interface{}{"start", "stop"}}, + }, + "required": []string{"action"}, + }, + }, + } +} + +func (d *Camera) Execute(tool string, args map[string]interface{}) (interface{}, error) { + return map[string]interface{}{"device": d.name, "status": "captured"}, nil +} + +// RobotArm +type RobotArm struct { + name string + io *IOManager +} + +func NewRobotArm(name string, io *IOManager) *RobotArm { + return &RobotArm{name: name, io: io} +} + +func (d *RobotArm) Name() string { return d.name } +func (d *RobotArm) Type() DeviceType { return DeviceIO } +func (d *RobotArm) OutputCapabilities() OutputCapability { return CapStructured } +func (d *RobotArm) Description() string { return fmt.Sprintf("机械臂 (%s)", d.name) } +func (d *RobotArm) Start() error { return nil } +func (d *RobotArm) Stop() error { return nil } + +func (d *RobotArm) Tools() []ToolDef { + return []ToolDef{ + { + Name: d.name + "_move", + Description: fmt.Sprintf("移动 %s 到坐标", d.name), + Parameters: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "x": map[string]interface{}{"type": "number", "description": "X 轴"}, + "y": map[string]interface{}{"type": "number", "description": "Y 轴"}, + "z": map[string]interface{}{"type": "number", "description": "Z 轴"}, + }, + "required": []string{"x", "y", "z"}, + }, + }, + { + Name: d.name + "_grip", + Description: fmt.Sprintf("控制 %s 夹爪", d.name), + Parameters: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "action": map[string]interface{}{"type": "string", "enum": []interface{}{"open", "close"}}, + }, + "required": []string{"action"}, + }, + }, + } +} + +func (d *RobotArm) Execute(tool string, args map[string]interface{}) (interface{}, error) { + return map[string]interface{}{"device": d.name, "tool": tool, "status": "executed"}, nil +} + +// GPIODevice +type GPIODevice struct { + name string + pins []int + io *IOManager +} + +func NewGPIODevice(name string, pins []int, io *IOManager) *GPIODevice { + return &GPIODevice{name: name, pins: pins, io: io} +} + +func (d *GPIODevice) Name() string { return d.name } +func (d *GPIODevice) Type() DeviceType { return DeviceIO } +func (d *GPIODevice) OutputCapabilities() OutputCapability { return CapStructured } +func (d *GPIODevice) Description() string { return "GPIO 通用引脚" } +func (d *GPIODevice) Start() error { return nil } +func (d *GPIODevice) Stop() error { return nil } + +func (d *GPIODevice) Tools() []ToolDef { + return []ToolDef{ + { + Name: d.name + "_gpio_write", + Description: "设置引脚电平", + Parameters: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "pin": map[string]interface{}{"type": "integer"}, + "value": map[string]interface{}{"type": "integer", "enum": []interface{}{0, 1}}, + }, + "required": []string{"pin", "value"}, + }, + }, + { + Name: d.name + "_gpio_read", + Description: "读取引脚电平", + Parameters: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "pin": map[string]interface{}{"type": "integer"}, + }, + "required": []string{"pin"}, + }, + }, + } +} + +func (d *GPIODevice) Execute(tool string, args map[string]interface{}) (interface{}, error) { + return map[string]interface{}{"device": d.name, "tool": tool, "status": "ok"}, nil +} diff --git a/internal/agent/personal.go b/internal/agent/personal.go new file mode 100644 index 0000000..44a6cd4 --- /dev/null +++ b/internal/agent/personal.go @@ -0,0 +1,41 @@ +package agent + +import ( + "fmt" + "os" + "path/filepath" +) + +type Personality struct { + Content string + Path string +} + +func LoadPersonality(path string) (*Personality, error) { + data, err := os.ReadFile(path) + if err != nil { + if os.IsNotExist(err) { + return &Personality{}, nil + } + return nil, fmt.Errorf("read personal.md: %w", err) + } + return &Personality{ + Content: string(data), + Path: path, + }, nil +} + +func SavePersonality(path, content string) error { + dir := filepath.Dir(path) + if err := os.MkdirAll(dir, 0755); err != nil { + return fmt.Errorf("create personality dir: %w", err) + } + return os.WriteFile(path, []byte(content), 0644) +} + +func (p *Personality) InjectPrompt() string { + if p.Content == "" { + return "" + } + return fmt.Sprintf("【人格设定】\n%s\n", p.Content) +} diff --git a/internal/api/handler.go b/internal/api/handler.go new file mode 100644 index 0000000..9901c55 --- /dev/null +++ b/internal/api/handler.go @@ -0,0 +1,653 @@ +package api + +import ( + "encoding/json" + "fmt" + "net/http" + "os" + "strconv" + "strings" + "time" + + agentIO "gitcode.com/JianFeeeee/HomeAgent/internal/agent/io" + "gitcode.com/JianFeeeee/HomeAgent/internal/knowledge" + luaVM "gitcode.com/JianFeeeee/HomeAgent/internal/lua" + "gitcode.com/JianFeeeee/HomeAgent/internal/memory" + "gitcode.com/JianFeeeee/HomeAgent/internal/memory/text" + "gitcode.com/JianFeeeee/HomeAgent/internal/skill" + "gitcode.com/JianFeeeee/HomeAgent/internal/supervisor" + "gitcode.com/JianFeeeee/HomeAgent/pkg/types" +) + +type Handler struct { + supervisor *supervisor.Daemon + memory *memory.GraphDB + indexer *memory.Indexer + skills *skill.Manager + lua *luaVM.VM + config *types.Config + startTime time.Time + iom *agentIO.IOManager + textMem *text.Memory + knowledge *knowledge.Store +} + +func NewHandler(sup *supervisor.Daemon, mem *memory.GraphDB, sk *skill.Manager, lua *luaVM.VM, cfg *types.Config, iom *agentIO.IOManager, tm *text.Memory, ks *knowledge.Store) *Handler { + var idx *memory.Indexer + if mem != nil { + idx = memory.NewIndexer(mem) + } + return &Handler{ + supervisor: sup, + memory: mem, + indexer: idx, + skills: sk, + lua: lua, + config: cfg, + startTime: time.Now(), + iom: iom, + textMem: tm, + knowledge: ks, + } +} + +func (h *Handler) RegisterRoutes(mux *http.ServeMux) { + mux.HandleFunc("/api/v1/status", h.handleStatus) + mux.HandleFunc("/api/v1/agents", h.handleAgents) + mux.HandleFunc("/api/v1/agents/", h.handleAgentByID) + mux.HandleFunc("/api/v1/skills", h.handleSkills) + mux.HandleFunc("/api/v1/memory", h.handleMemory) + mux.HandleFunc("/api/v1/memory/", h.handleMemory) + mux.HandleFunc("/api/v1/memory/context", h.handleMemoryContext) + mux.HandleFunc("/api/v1/memory/tools", h.handleMemoryTools) + mux.HandleFunc("/api/v1/memory/text", h.handleTextMemory) + mux.HandleFunc("/api/v1/network", h.handleNetwork) + mux.HandleFunc("/api/v1/config", h.handleConfig) + mux.HandleFunc("/api/v1/knowledge", h.handleKnowledge) + mux.HandleFunc("/api/v1/knowledge/", h.handleKnowledge) + mux.HandleFunc("/api/v1/adapters", h.handleAdapters) + mux.HandleFunc("/api/v1/adapters/", h.handleAdapterByID) + mux.HandleFunc("/v1/chat/completions", h.handleOpenAICompletions) + mux.HandleFunc("/", h.handleStatic) +} + +func (h *Handler) handleStatus(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodGet { + http.Error(w, "method not allowed", http.StatusMethodNotAllowed) + return + } + agents := h.supervisor.ListAgents() + writeJSON(w, http.StatusOK, map[string]interface{}{ + "status": "running", + "uptime": time.Since(h.startTime).String(), + "agents": len(agents), + "version": "0.1.0", + "startedAt": h.startTime, + }) +} + +func (h *Handler) handleAgents(w http.ResponseWriter, r *http.Request) { + switch r.Method { + case http.MethodGet: + agents := h.supervisor.ListAgents() + writeJSON(w, http.StatusOK, map[string]interface{}{"agents": agents}) + case http.MethodPost: + var cfg types.AgentConfig + if err := json.NewDecoder(r.Body).Decode(&cfg); err != nil { + writeJSON(w, http.StatusBadRequest, map[string]string{"error": "invalid request body"}) + return + } + if cfg.ID == "" { + writeJSON(w, http.StatusBadRequest, map[string]string{"error": "agent id is required"}) + return + } + h.config.Agents = append(h.config.Agents, cfg) + writeJSON(w, http.StatusCreated, map[string]string{"id": string(cfg.ID)}) + default: + http.Error(w, "method not allowed", http.StatusMethodNotAllowed) + } +} + +func (h *Handler) handleAgentByID(w http.ResponseWriter, r *http.Request) { + path := strings.TrimPrefix(r.URL.Path, "/api/v1/agents/") + parts := strings.Split(path, "/") + agentID := types.AgentID(parts[0]) + if len(parts) == 1 { + switch r.Method { + case http.MethodGet: + status, err := h.supervisor.GetAgentStatus(agentID) + if err != nil { + writeJSON(w, http.StatusNotFound, map[string]string{"error": err.Error()}) + return + } + writeJSON(w, http.StatusOK, status) + case http.MethodDelete: + writeJSON(w, http.StatusOK, map[string]string{"status": "deleted", "id": string(agentID)}) + default: + http.Error(w, "method not allowed", http.StatusMethodNotAllowed) + } + return + } + action := parts[1] + switch action { + case "snapshots": + h.handleSnapshots(w, r, agentID, parts) + case "rollback": + h.handleRollback(w, r, agentID, parts) + case "start", "stop", "restart": + h.handleAgentAction(w, r, agentID, action) + default: + http.Error(w, "not found", http.StatusNotFound) + } +} + +func (h *Handler) handleSnapshots(w http.ResponseWriter, r *http.Request, agentID types.AgentID, parts []string) { + switch r.Method { + case http.MethodGet: + writeJSON(w, http.StatusOK, map[string]interface{}{"agent_id": agentID, "snapshots": []map[string]interface{}{}}) + case http.MethodPost: + snap, err := h.supervisor.PreActionSnapshot(agentID) + if err != nil { + writeJSON(w, http.StatusInternalServerError, map[string]string{"error": err.Error()}) + return + } + writeJSON(w, http.StatusCreated, snap) + default: + http.Error(w, "method not allowed", http.StatusMethodNotAllowed) + } +} + +func (h *Handler) handleRollback(w http.ResponseWriter, r *http.Request, agentID types.AgentID, parts []string) { + if r.Method != http.MethodPost || len(parts) < 3 { + http.Error(w, "method not allowed", http.StatusMethodNotAllowed) + return + } + snapID := types.SnapshotID(parts[2]) + if err := h.supervisor.RollbackAgent(agentID, snapID); err != nil { + writeJSON(w, http.StatusInternalServerError, map[string]string{"error": err.Error()}) + return + } + writeJSON(w, http.StatusOK, map[string]string{"status": "rollback_initiated", "agent": string(agentID), "snap": string(snapID)}) +} + +func (h *Handler) handleAgentAction(w http.ResponseWriter, r *http.Request, agentID types.AgentID, action string) { + if r.Method != http.MethodPost { + http.Error(w, "method not allowed", http.StatusMethodNotAllowed) + return + } + writeJSON(w, http.StatusOK, map[string]string{"status": fmt.Sprintf("%s_requested", action), "agent": string(agentID)}) +} + +func (h *Handler) handleSkills(w http.ResponseWriter, r *http.Request) { + switch r.Method { + case http.MethodGet: + writeJSON(w, http.StatusOK, map[string]interface{}{"skills": h.skills.List()}) + case http.MethodPost: + var req struct { + Name string `json:"name"` + Content string `json:"content"` + } + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + writeJSON(w, http.StatusBadRequest, map[string]string{"error": "invalid request"}) + return + } + if err := h.skills.Install(req.Name, req.Content); err != nil { + writeJSON(w, http.StatusInternalServerError, map[string]string{"error": err.Error()}) + return + } + writeJSON(w, http.StatusCreated, map[string]string{"status": "installed", "name": req.Name}) + case http.MethodDelete: + name := r.URL.Query().Get("name") + if name == "" { + writeJSON(w, http.StatusBadRequest, map[string]string{"error": "name query param required"}) + return + } + if err := h.skills.Uninstall(name); err != nil { + writeJSON(w, http.StatusNotFound, map[string]string{"error": err.Error()}) + return + } + writeJSON(w, http.StatusOK, map[string]string{"status": "uninstalled", "name": name}) + default: + http.Error(w, "method not allowed", http.StatusMethodNotAllowed) + } +} + +func (h *Handler) handleMemory(w http.ResponseWriter, r *http.Request) { + if h.memory == nil { + writeJSON(w, http.StatusServiceUnavailable, map[string]string{"error": "memory system not available"}) + return + } + switch r.Method { + case http.MethodGet: + userInput := r.URL.Query().Get("q") + keywords := strings.Split(userInput, ",") + depth, _ := strconv.Atoi(r.URL.Query().Get("depth")) + if depth <= 0 { + depth = 2 + } + result, err := h.memory.Recall(keywords, nil, depth, "") + if err != nil { + writeJSON(w, http.StatusInternalServerError, map[string]string{"error": err.Error()}) + return + } + writeJSON(w, http.StatusOK, result) + case http.MethodPost: + var req struct { + Triples []memory.Triple `json:"triples"` + SessionID string `json:"session_id"` + TurnID int `json:"turn_id"` + } + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + writeJSON(w, http.StatusBadRequest, map[string]string{"error": "invalid request"}) + return + } + ec, rc, err := h.memory.Commit(req.Triples, req.SessionID, req.TurnID) + if err != nil { + writeJSON(w, http.StatusInternalServerError, map[string]string{"error": err.Error()}) + return + } + writeJSON(w, http.StatusCreated, map[string]int{"entities_created": ec, "relations_created": rc}) + case http.MethodDelete: + var req struct { + Criteria map[string]string `json:"criteria"` + Mode string `json:"mode"` + } + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + writeJSON(w, http.StatusBadRequest, map[string]string{"error": "invalid request"}) + return + } + deleted, err := h.memory.Purge(req.Criteria, req.Mode) + if err != nil { + writeJSON(w, http.StatusInternalServerError, map[string]string{"error": err.Error()}) + return + } + writeJSON(w, http.StatusOK, map[string]int{"deleted": deleted}) + default: + http.Error(w, "method not allowed", http.StatusMethodNotAllowed) + } +} + +func (h *Handler) handleMemoryContext(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodGet { + http.Error(w, "method not allowed", http.StatusMethodNotAllowed) + return + } + if h.indexer == nil { + writeJSON(w, http.StatusServiceUnavailable, map[string]string{"error": "indexer not available"}) + return + } + userInput := r.URL.Query().Get("q") + injected := h.indexer.BuildContext(userInput) + writeJSON(w, http.StatusOK, map[string]interface{}{ + "context": h.indexer.FormatContext(injected), + "summary": injected.Summary, + "entities": injected.Entities, + "token_estimate": injected.TokenEstimate, + "tool_prompt": h.indexer.BuildToolPrompt(), + }) +} + +func (h *Handler) handleMemoryTools(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodGet { + http.Error(w, "method not allowed", http.StatusMethodNotAllowed) + return + } + if h.indexer == nil { + writeJSON(w, http.StatusServiceUnavailable, map[string]string{"error": "indexer not available"}) + return + } + writeJSON(w, http.StatusOK, map[string]interface{}{ + "tools": h.indexer.GetToolDefinitions(), + "tool_prompt": h.indexer.BuildToolPrompt(), + }) +} + +func (h *Handler) handleKnowledge(w http.ResponseWriter, r *http.Request) { + switch r.Method { + case http.MethodGet: + if h.knowledge == nil { + writeJSON(w, http.StatusServiceUnavailable, map[string]string{"error": "knowledge not available"}) + return + } + query := r.URL.Query().Get("q") + if query != "" { + results := h.knowledge.Search(query, 10) + writeJSON(w, http.StatusOK, map[string]interface{}{"results": results}) + return + } + writeJSON(w, http.StatusOK, map[string]interface{}{ + "categories": h.knowledge.List(), + "stats": h.knowledge.Stats(), + }) + + case http.MethodPost: + if h.knowledge == nil { + writeJSON(w, http.StatusServiceUnavailable, map[string]string{"error": "knowledge not available"}) + return + } + ct := r.Header.Get("Content-Type") + if strings.HasPrefix(ct, "multipart/form-data") { + if err := r.ParseMultipartForm(10 << 20); err != nil { + writeJSON(w, http.StatusBadRequest, map[string]string{"error": err.Error()}) + return + } + name := r.FormValue("name") + file, _, err := r.FormFile("file") + if err != nil { + writeJSON(w, http.StatusBadRequest, map[string]string{"error": "file required"}) + return + } + defer file.Close() + buf := make([]byte, 10<<20) + n, _ := file.Read(buf) + content := string(buf[:n]) + if err := h.knowledge.Add(name, content); err != nil { + writeJSON(w, http.StatusInternalServerError, map[string]string{"error": err.Error()}) + return + } + writeJSON(w, http.StatusCreated, map[string]string{"status": "created", "name": name}) + return + } + + var req struct { + Name string `json:"name"` + Content string `json:"content"` + } + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + writeJSON(w, http.StatusBadRequest, map[string]string{"error": "invalid request"}) + return + } + if req.Name == "" || req.Content == "" { + writeJSON(w, http.StatusBadRequest, map[string]string{"error": "name and content required"}) + return + } + if err := h.knowledge.Add(req.Name, req.Content); err != nil { + writeJSON(w, http.StatusInternalServerError, map[string]string{"error": err.Error()}) + return + } + writeJSON(w, http.StatusCreated, map[string]string{"status": "created", "name": req.Name}) + + case http.MethodDelete: + name := r.URL.Query().Get("name") + if name == "" { + writeJSON(w, http.StatusBadRequest, map[string]string{"error": "name query param required"}) + return + } + if err := h.knowledge.Remove(name); err != nil { + writeJSON(w, http.StatusNotFound, map[string]string{"error": err.Error()}) + return + } + writeJSON(w, http.StatusOK, map[string]string{"status": "deleted", "name": name}) + + default: + http.Error(w, "method not allowed", http.StatusMethodNotAllowed) + } +} + +func (h *Handler) handleTextMemory(w http.ResponseWriter, r *http.Request) { + if h.textMem == nil { + writeJSON(w, http.StatusServiceUnavailable, map[string]string{"error": "text memory not available"}) + return + } + switch r.Method { + case http.MethodGet: + recent, _ := h.textMem.RecentEvents(50) + stats := h.textMem.Stats() + writeJSON(w, http.StatusOK, map[string]interface{}{ + "stats": stats, + "recent": recent, + }) + case http.MethodDelete: + // future: purge + writeJSON(w, http.StatusAccepted, map[string]string{"status": "not_implemented"}) + default: + http.Error(w, "method not allowed", http.StatusMethodNotAllowed) + } +} + +func (h *Handler) handleAdapters(w http.ResponseWriter, r *http.Request) { + if h.lua == nil { + writeJSON(w, http.StatusServiceUnavailable, map[string]string{"error": "lua vm not available"}) + return + } + switch r.Method { + case http.MethodGet: + writeJSON(w, http.StatusOK, map[string]interface{}{"adapters": h.lua.ListAdapters()}) + case http.MethodPost: + var req struct { + Name string `json:"name"` + Code string `json:"code"` + } + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + writeJSON(w, http.StatusBadRequest, map[string]string{"error": "invalid request"}) + return + } + path := fmt.Sprintf("%s/%s.lua", h.lua.AdapterDir(), req.Name) + if err := os.WriteFile(path, []byte(req.Code), 0644); err != nil { + writeJSON(w, http.StatusInternalServerError, map[string]string{"error": err.Error()}) + return + } + if err := h.lua.ReloadAll(); err != nil { + writeJSON(w, http.StatusInternalServerError, map[string]string{"error": err.Error()}) + return + } + writeJSON(w, http.StatusCreated, map[string]string{"status": "loaded", "name": req.Name}) + default: + http.Error(w, "method not allowed", http.StatusMethodNotAllowed) + } +} + +func (h *Handler) handleAdapterByID(w http.ResponseWriter, r *http.Request) { + if h.lua == nil { + writeJSON(w, http.StatusServiceUnavailable, map[string]string{"error": "lua vm not available"}) + return + } + name := strings.TrimPrefix(r.URL.Path, "/api/v1/adapters/") + if name == "" { + http.NotFound(w, r) + return + } + switch r.Method { + case http.MethodGet: + for _, a := range h.lua.ListAdapters() { + if a.Name == name { + writeJSON(w, http.StatusOK, a) + return + } + } + http.NotFound(w, r) + case http.MethodDelete: + path := fmt.Sprintf("%s/%s.lua", h.lua.AdapterDir(), name) + if err := os.Remove(path); err != nil { + writeJSON(w, http.StatusNotFound, map[string]string{"error": "adapter not found"}) + return + } + h.lua.ReloadAll() + writeJSON(w, http.StatusOK, map[string]string{"status": "deleted", "name": name}) + default: + http.Error(w, "method not allowed", http.StatusMethodNotAllowed) + } +} + +func (h *Handler) handleNetwork(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodGet { + http.Error(w, "method not allowed", http.StatusMethodNotAllowed) + return + } + writeJSON(w, http.StatusOK, map[string]interface{}{ + "network_status": "monitoring", + "endpoints": h.config.Defaults.LLMEndpoints, + }) +} + +func (h *Handler) handleConfig(w http.ResponseWriter, r *http.Request) { + switch r.Method { + case http.MethodGet: + writeJSON(w, http.StatusOK, h.config) + case http.MethodPut: + var cfg types.Config + if err := json.NewDecoder(r.Body).Decode(&cfg); err != nil { + writeJSON(w, http.StatusBadRequest, map[string]string{"error": "invalid config"}) + return + } + h.config = &cfg + writeJSON(w, http.StatusOK, map[string]string{"status": "config_updated"}) + default: + http.Error(w, "method not allowed", http.StatusMethodNotAllowed) + } +} + +// OpenAI 兼容 API — 所有输入走 IO 抽象层(中断) +func (h *Handler) handleOpenAICompletions(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + http.Error(w, "method not allowed", http.StatusMethodNotAllowed) + return + } + + var req struct { + Model string `json:"model"` + Messages []openAIMessage `json:"messages"` + Stream bool `json:"stream"` + Temperature float64 `json:"temperature"` + MaxTokens int `json:"max_tokens"` + } + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + writeJSON(w, http.StatusBadRequest, map[string]string{"error": "invalid request: " + err.Error()}) + return + } + if len(req.Messages) == 0 { + writeJSON(w, http.StatusBadRequest, map[string]string{"error": "messages is required"}) + return + } + + // 取最后一条 user 消息作为输入 + lastMsg := req.Messages[len(req.Messages)-1] + if lastMsg.Role != "user" { + writeJSON(w, http.StatusBadRequest, map[string]string{"error": "last message must be from user"}) + return + } + + // 通过 IO 抽象层同步注入(中断式) + response := h.iom.InjectTextSync("http", lastMsg.Content) + + resp := map[string]interface{}{ + "id": fmt.Sprintf("chatcmpl-%d", time.Now().UnixNano()), + "object": "chat.completion", + "created": time.Now().Unix(), + "model": req.Model, + "choices": []map[string]interface{}{ + { + "index": 0, + "message": map[string]interface{}{ + "role": "assistant", + "content": response.Payload["content"], + }, + "finish_reason": "stop", + }, + }, + "usage": map[string]interface{}{ + "prompt_tokens": len(lastMsg.Content) / 2, + "completion_tokens": len(response.Payload["content"].(string)) / 2, + "total_tokens": (len(lastMsg.Content) + len(response.Payload["content"].(string))) / 2, + }, + } + + w.Header().Set("Content-Type", "application/json; charset=utf-8") + json.NewEncoder(w).Encode(resp) +} + +func (h *Handler) handleStatic(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/" { + w.Header().Set("Content-Type", "text/html; charset=utf-8") + w.Write(webuiHTML) + return + } + http.NotFound(w, r) +} + +type openAIMessage struct { + Role string `json:"role"` + Content string `json:"content"` +} + +func writeJSON(w http.ResponseWriter, status int, data interface{}) { + w.Header().Set("Content-Type", "application/json; charset=utf-8") + w.WriteHeader(status) + json.NewEncoder(w).Encode(data) +} + +var webuiHTML = []byte(` + + + + +HomeAgent Dashboard + + + + +
+
+
+
+
+
+
+ + +`) diff --git a/internal/container/manager.go b/internal/container/manager.go new file mode 100644 index 0000000..984249e --- /dev/null +++ b/internal/container/manager.go @@ -0,0 +1,209 @@ +package container + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "os/exec" + "strings" + + "gitcode.com/JianFeeeee/HomeAgent/pkg/types" +) + +type Manager struct { + dataDir string +} + +func NewManager(dataDir string) *Manager { + return &Manager{dataDir: dataDir} +} + +type ContainerInfo struct { + ID string `json:"id"` + Name string `json:"name"` + Status string `json:"status"` + Image string `json:"image"` +} + +func (m *Manager) Create(ctx context.Context, cfg *types.AgentConfig) (*ContainerInfo, error) { + args := []string{ + "create", + "--name", "ha-" + string(cfg.ID), + "--hostname", string(cfg.ID), + "--restart", "no", + "--stop-timeout", "10", + "--memory", cfg.ResourceLimit.Memory, + "--cpus", cfg.ResourceLimit.CPU, + "--label", "homeagent.managed=true", + "--label", "homeagent.agent-id=" + string(cfg.ID), + } + + if !cfg.ResourceLimit.Network { + args = append(args, "--network", "none") + } + + args = append(args, cfg.Image) + + cmd := exec.CommandContext(ctx, "docker", args...) + var stderr bytes.Buffer + cmd.Stderr = &stderr + + out, err := cmd.Output() + if err != nil { + return nil, fmt.Errorf("docker create: %s: %w", strings.TrimSpace(stderr.String()), err) + } + + id := strings.TrimSpace(string(out)) + return &ContainerInfo{ID: id, Name: "ha-" + string(cfg.ID), Status: "created", Image: cfg.Image}, nil +} + +func (m *Manager) Start(ctx context.Context, containerID string) error { + cmd := exec.CommandContext(ctx, "docker", "start", containerID) + var stderr bytes.Buffer + cmd.Stderr = &stderr + + if err := cmd.Run(); err != nil { + return fmt.Errorf("docker start: %s: %w", strings.TrimSpace(stderr.String()), err) + } + return nil +} + +func (m *Manager) Stop(ctx context.Context, containerID string) error { + cmd := exec.CommandContext(ctx, "docker", "stop", "--time", "5", containerID) + var stderr bytes.Buffer + cmd.Stderr = &stderr + + if err := cmd.Run(); err != nil { + return fmt.Errorf("docker stop: %s: %w", strings.TrimSpace(stderr.String()), err) + } + return nil +} + +func (m *Manager) Remove(ctx context.Context, containerID string) error { + cmd := exec.CommandContext(ctx, "docker", "rm", "-f", containerID) + var stderr bytes.Buffer + cmd.Stderr = &stderr + + if err := cmd.Run(); err != nil { + return fmt.Errorf("docker rm: %s: %w", strings.TrimSpace(stderr.String()), err) + } + return nil +} + +func (m *Manager) Inspect(ctx context.Context, containerID string) (*ContainerInfo, error) { + cmd := exec.CommandContext(ctx, "docker", "inspect", containerID) + out, err := cmd.Output() + if err != nil { + return nil, fmt.Errorf("docker inspect: %w", err) + } + + var containers []struct { + ID string `json:"Id"` + Name string `json:"Name"` + State struct { + Status string `json:"Status"` + } `json:"State"` + Config struct { + Image string `json:"Image"` + } `json:"Config"` + } + + if err := json.Unmarshal(out, &containers); err != nil { + return nil, fmt.Errorf("parse inspect: %w", err) + } + + if len(containers) == 0 { + return nil, fmt.Errorf("container %s not found", containerID) + } + + c := containers[0] + return &ContainerInfo{ + ID: c.ID, + Name: strings.TrimPrefix(c.Name, "/"), + Status: c.State.Status, + Image: c.Config.Image, + }, nil +} + +func (m *Manager) WaitHealthy(ctx context.Context, containerID string) error { + args := []string{ + "exec", containerID, + "agentd", "--probe", + } + + cmd := exec.CommandContext(ctx, "docker", args...) + return cmd.Run() +} + +func (m *Manager) Exec(ctx context.Context, containerID string, cmdArgs []string) ([]byte, error) { + args := append([]string{"exec", containerID}, cmdArgs...) + cmd := exec.CommandContext(ctx, "docker", args...) + var stderr bytes.Buffer + cmd.Stderr = &stderr + + out, err := cmd.Output() + if err != nil { + return nil, fmt.Errorf("docker exec: %s: %w", strings.TrimSpace(stderr.String()), err) + } + return out, nil +} + +func (m *Manager) Commit(ctx context.Context, containerID string, imageTag string) error { + cmd := exec.CommandContext(ctx, "docker", "commit", containerID, imageTag) + var stderr bytes.Buffer + cmd.Stderr = &stderr + + if err := cmd.Run(); err != nil { + return fmt.Errorf("docker commit: %s: %w", strings.TrimSpace(stderr.String()), err) + } + return nil +} + +func (m *Manager) SaveImage(ctx context.Context, imageTag string, outputPath string) error { + cmd := exec.CommandContext(ctx, "docker", "save", "-o", outputPath, imageTag) + var stderr bytes.Buffer + cmd.Stderr = &stderr + + if err := cmd.Run(); err != nil { + return fmt.Errorf("docker save: %s: %w", strings.TrimSpace(stderr.String()), err) + } + return nil +} + +func (m *Manager) LoadImage(ctx context.Context, inputPath string) error { + cmd := exec.CommandContext(ctx, "docker", "load", "-i", inputPath) + var stderr bytes.Buffer + cmd.Stderr = &stderr + + if err := cmd.Run(); err != nil { + return fmt.Errorf("docker load: %s: %w", strings.TrimSpace(stderr.String()), err) + } + return nil +} + +func (m *Manager) ListManaged(ctx context.Context) ([]ContainerInfo, error) { + cmd := exec.CommandContext(ctx, "docker", "ps", "-a", + "--filter", "label=homeagent.managed=true", + "--format", "{{.ID}}\t{{.Names}}\t{{.Status}}\t{{.Image}}", + ) + out, err := cmd.Output() + if err != nil { + return nil, fmt.Errorf("docker ps: %w", err) + } + + var containers []ContainerInfo + for _, line := range strings.Split(strings.TrimSpace(string(out)), "\n") { + if line == "" { + continue + } + parts := strings.SplitN(line, "\t", 4) + if len(parts) < 4 { + continue + } + containers = append(containers, ContainerInfo{ + ID: parts[0], Name: parts[1], Status: parts[2], Image: parts[3], + }) + } + return containers, nil +} diff --git a/internal/embed/embedder.go b/internal/embed/embedder.go new file mode 100644 index 0000000..68274c7 --- /dev/null +++ b/internal/embed/embedder.go @@ -0,0 +1,162 @@ +package embed + +import ( + "bytes" + "encoding/json" + "fmt" + "math" + "net/http" + "sync" + "time" +) + +type Embedder interface { + Embed(text string) ([]float64, error) + Similarity(a, b []float64) float64 + Dimension() int +} + +type OllamaEmbedder struct { + client *http.Client + baseURL string + model string + dimension int + mu sync.RWMutex +} + +func NewOllamaEmbedder(baseURL, model string, dimension int) *OllamaEmbedder { + if baseURL == "" { + baseURL = "http://localhost:11434" + } + if model == "" { + model = "nomic-embed-text" + } + if dimension <= 0 { + dimension = 768 + } + + return &OllamaEmbedder{ + client: &http.Client{ + Timeout: 30 * time.Second, + }, + baseURL: baseURL, + model: model, + dimension: dimension, + } +} + +func (e *OllamaEmbedder) Embed(text string) ([]float64, error) { + if text == "" { + return make([]float64, e.dimension), nil + } + + reqBody := map[string]interface{}{ + "model": e.model, + "prompt": text, + } + + data, err := json.Marshal(reqBody) + if err != nil { + return nil, fmt.Errorf("marshal: %w", err) + } + + resp, err := e.client.Post(e.baseURL+"/api/embeddings", "application/json", bytes.NewReader(data)) + if err != nil { + return nil, fmt.Errorf("ollama api: %w", err) + } + defer resp.Body.Close() + + var result struct { + Embedding []float64 `json:"embedding"` + } + if err := json.NewDecoder(resp.Body).Decode(&result); err != nil { + return nil, fmt.Errorf("decode: %w", err) + } + + return result.Embedding, nil +} + +func (e *OllamaEmbedder) Similarity(a, b []float64) float64 { + return cosineSimilarity(a, b) +} + +func (e *OllamaEmbedder) Dimension() int { + return e.dimension +} + +type HashEmbedder struct { + dimension int +} + +func NewHashEmbedder(dimension int) *HashEmbedder { + if dimension <= 0 { + dimension = 64 + } + return &HashEmbedder{dimension: dimension} +} + +func (e *HashEmbedder) Embed(text string) ([]float64, error) { + vec := make([]float64, e.dimension) + runes := []rune(text) + if len(runes) == 0 { + return vec, nil + } + + // Character-level hash embedding + for i, r := range runes { + h := hashRune(r) + idx := i % e.dimension + vec[idx] += float64(h) / 65536.0 + } + + // Normalize + mag := 0.0 + for _, v := range vec { + mag += v * v + } + if mag > 0 { + mag = math.Sqrt(mag) + for i := range vec { + vec[i] /= mag + } + } + + return vec, nil +} + +func (e *HashEmbedder) Similarity(a, b []float64) float64 { + return cosineSimilarity(a, b) +} + +func (e *HashEmbedder) Dimension() int { + return e.dimension +} + +func hashRune(r rune) uint64 { + h := uint64(r) + h ^= h >> 33 + h *= 0xff51afd7ed558ccd + h ^= h >> 33 + h *= 0xc4ceb9fe1a85ec53 + h ^= h >> 33 + return h +} + +func cosineSimilarity(a, b []float64) float64 { + if len(a) != len(b) || len(a) == 0 { + return 0 + } + + var dot, na, nb float64 + for i := range a { + dot += a[i] * b[i] + na += a[i] * a[i] + nb += b[i] * b[i] + } + + if na == 0 || nb == 0 { + return 0 + } + + return dot / (math.Sqrt(na) * math.Sqrt(nb)) +} diff --git a/internal/knowledge/knowledge.go b/internal/knowledge/knowledge.go new file mode 100644 index 0000000..f8d7475 --- /dev/null +++ b/internal/knowledge/knowledge.go @@ -0,0 +1,274 @@ +package knowledge + +import ( + "fmt" + "log" + "os" + "path/filepath" + "sort" + "strings" + "sync" + "time" + + "gitcode.com/JianFeeeee/HomeAgent/internal/memory/vector" +) + +// Knowledge — 单条知识 +type Knowledge struct { + Name string `json:"name"` + Content string `json:"content"` + Path string `json:"path"` + Tags []string `json:"tags"` + UpdatedAt time.Time `json:"updated_at"` + Meta map[string]string `json:"meta,omitempty"` +} + +// Store — 知识库,文件系统 + 向量索引 +type Store struct { + root string + vec *vector.Store + veczer *vector.TFIDFVectorizer + mu sync.RWMutex + items map[string]*Knowledge + + summaries []string +} + +func NewStore(root string) *Store { + return &Store{ + root: root, + vec: vector.NewStore(), + veczer: vector.NewTFIDFVectorizer(3), + items: make(map[string]*Knowledge), + } +} + +func (s *Store) Start() error { + if err := os.MkdirAll(s.root, 0755); err != nil { + return fmt.Errorf("knowledge root: %w", err) + } + if err := s.scanAll(); err != nil { + log.Printf("[knowledge] scan error: %v", err) + } + log.Printf("[knowledge] started with %d items, %d vectors", len(s.items), s.vec.Size()) + return nil +} + +func (s *Store) Stop() {} + +// Search — 向量查询知识 +func (s *Store) Search(query string, topK int) []*Knowledge { + s.mu.RLock() + defer s.mu.RUnlock() + + if topK <= 0 { + topK = 5 + } + + vec := s.veczer.Vectorize(query) + results := s.vec.Search(vec, topK) + + var out []*Knowledge + for _, r := range results { + if k, ok := s.items[r.ID]; ok { + out = append(out, k) + } + } + return out +} + +// Add — 添加或更新知识 +func (s *Store) Add(name, content string) error { + s.mu.Lock() + defer s.mu.Unlock() + + // 创建知识目录 + dir := filepath.Join(s.root, sanitize(name)) + if err := os.MkdirAll(dir, 0755); err != nil { + return fmt.Errorf("create knowledge dir: %w", err) + } + + // 写入知识文件 + path := filepath.Join(dir, "content.md") + if err := os.WriteFile(path, []byte(content), 0644); err != nil { + return fmt.Errorf("write knowledge: %w", err) + } + + now := time.Now() + k := &Knowledge{ + Name: name, + Content: content, + Path: path, + Tags: extractKeywords(name + " " + content), + UpdatedAt: now, + } + + // 生成 ID = 目录名 + id := sanitize(name) + s.items[id] = k + + vec := s.veczer.Vectorize(name + " " + content) + s.vec.Insert(id, name+": "+content, vec, map[string]string{ + "name": name, "path": path, + }) + + s.summaries = append(s.summaries, name+" "+content) + + log.Printf("[knowledge] added: %s (%d bytes)", name, len(content)) + return nil +} + +// SearchCategories — 返回所有知识类别 +func (s *Store) SearchCategories(query string, topK int) []string { + s.mu.RLock() + defer s.mu.RUnlock() + + if query == "" { + var names []string + for _, k := range s.items { + names = append(names, k.Name) + } + sort.Strings(names) + if len(names) > topK { + names = names[:topK] + } + return names + } + + vec := s.veczer.Vectorize(query) + results := s.vec.Search(vec, topK) + var names []string + for _, r := range results { + if k, ok := s.items[r.ID]; ok { + names = append(names, k.Name) + } + } + return names +} + +func (s *Store) Remove(name string) error { + s.mu.Lock() + defer s.mu.Unlock() + + id := sanitize(name) + dir := filepath.Join(s.root, id) + if err := os.RemoveAll(dir); err != nil { + return err + } + delete(s.items, id) + s.vec.Remove(id) + return nil +} + +func (s *Store) Stats() map[string]interface{} { + s.mu.RLock() + defer s.mu.RUnlock() + return map[string]interface{}{ + "knowledge_count": len(s.items), + "vector_count": s.vec.Size(), + "root": s.root, + } +} + +func (s *Store) List() []string { + s.mu.RLock() + defer s.mu.RUnlock() + var names []string + for _, k := range s.items { + names = append(names, k.Name) + } + sort.Strings(names) + return names +} + +// ——— internal ——— + +func (s *Store) scanAll() error { + entries, err := os.ReadDir(s.root) + if err != nil { + return err + } + + for _, entry := range entries { + if !entry.IsDir() { + continue + } + dir := filepath.Join(s.root, entry.Name()) + contentPath := filepath.Join(dir, "content.md") + data, err := os.ReadFile(contentPath) + if err != nil { + continue + } + + name := entry.Name() + content := string(data) + now := time.Now() + + k := &Knowledge{ + Name: name, + Content: content, + Path: contentPath, + Tags: extractKeywords(name + " " + content), + UpdatedAt: now, + } + s.items[name] = k + s.summaries = append(s.summaries, name+" "+content) + } + + // 训练向量化器 + if len(s.summaries) > 0 { + s.veczer.Train(s.summaries) + } + + // 构建向量索引 + for _, k := range s.items { + vec := s.veczer.Vectorize(k.Name + " " + k.Content) + s.vec.Insert(k.Name, k.Name+": "+k.Content, vec, map[string]string{ + "name": k.Name, "path": k.Path, + }) + } + + return nil +} + +func sanitize(name string) string { + name = strings.ToLower(name) + name = strings.TrimSpace(name) + name = strings.ReplaceAll(name, " ", "_") + name = strings.ReplaceAll(name, "/", "_") + name = strings.ReplaceAll(name, "\\", "_") + return name +} + +func extractKeywords(text string) []string { + stopWords := map[string]bool{ + "的": true, "了": true, "是": true, "在": true, "有": true, + "和": true, "就": true, "不": true, "都": true, + "一": true, "一个": true, "也": true, "很": true, + "到": true, "说": true, "要": true, "去": true, + "会": true, "着": true, "没有": true, "看": true, "好": true, + "自己": true, "这": true, "他": true, "她": true, "它": true, + "什么": true, "怎么": true, "为什么": true, "如何": true, + "我们": true, "你们": true, "他们": true, "这个": true, + "那个": true, "可以": true, "吗": true, "吧": true, "啊": true, + } + + var keywords []string + runes := []rune(text) + + seen := make(map[string]bool) + + // bi-gram + for i := 0; i < len(runes)-1; i++ { + word := string(runes[i : i+2]) + if !stopWords[word] && len(strings.TrimSpace(word)) == len(word) && !seen[word] { + seen[word] = true + keywords = append(keywords, word) + } + } + + if len(keywords) > 10 { + keywords = keywords[:10] + } + return keywords +} diff --git a/internal/lua/adapters/deepseek.lua b/internal/lua/adapters/deepseek.lua new file mode 100644 index 0000000..3f0362f --- /dev/null +++ b/internal/lua/adapters/deepseek.lua @@ -0,0 +1,22 @@ +local adapter = {} + +adapter.name = "deepseek" +adapter.version = "1.0.0" + +function adapter.transform_request(input) + local messages = input.messages or {} + local result = { + model = input.model or "deepseek-chat", + messages = messages, + temperature = input.temperature or 0.0, + max_tokens = input.max_tokens or 4096, + stream = input.stream or false + } + return result +end + +function adapter.transform_response(raw) + return raw +end + +return adapter diff --git a/internal/lua/adapters/ollama.lua b/internal/lua/adapters/ollama.lua new file mode 100644 index 0000000..ff7ce94 --- /dev/null +++ b/internal/lua/adapters/ollama.lua @@ -0,0 +1,24 @@ +local adapter = {} + +adapter.name = "ollama" +adapter.version = "1.0.0" + +function adapter.transform_request(input) + local messages = input.messages or {} + local result = { + model = input.model or "llama3", + messages = messages, + stream = input.stream or false, + options = { + temperature = input.temperature or 0.7, + num_predict = input.max_tokens or 2048 + } + } + return result +end + +function adapter.transform_response(raw) + return raw +end + +return adapter diff --git a/internal/lua/adapters/openai.lua b/internal/lua/adapters/openai.lua new file mode 100644 index 0000000..ef4a9d3 --- /dev/null +++ b/internal/lua/adapters/openai.lua @@ -0,0 +1,22 @@ +local adapter = {} + +adapter.name = "openai" +adapter.version = "1.0.0" + +function adapter.transform_request(input) + local messages = input.messages or {} + local result = { + model = input.model or "gpt-4", + messages = messages, + temperature = input.temperature or 0.7, + max_tokens = input.max_tokens or 2048, + stream = input.stream or false + } + return result +end + +function adapter.transform_response(raw) + return raw +end + +return adapter diff --git a/internal/lua/vm.go b/internal/lua/vm.go new file mode 100644 index 0000000..0b0efcc --- /dev/null +++ b/internal/lua/vm.go @@ -0,0 +1,327 @@ +package lua + +import ( + "embed" + "fmt" + "os" + "path/filepath" + "sync" + + lua "github.com/yuin/gopher-lua" +) + +//go:embed adapters/*.lua +var bundledAdapters embed.FS + +type VM struct { + mu sync.Mutex + state *lua.LState + adapterDir string + loaded map[string]*lua.LTable +} + +type APIAdapter struct { + Name string + Version string + Script string +} + +func (v *VM) AdapterDir() string { + return v.adapterDir +} + +func NewVM(adapterDir string) *VM { + return &VM{ + adapterDir: adapterDir, + loaded: make(map[string]*lua.LTable), + } +} + +func (v *VM) Start() error { + os.MkdirAll(v.adapterDir, 0755) + + if err := v.writeBundledAdapters(); err != nil { + return fmt.Errorf("write bundled adapters: %w", err) + } + + v.state = lua.NewState() + + v.state.SetGlobal("log", v.state.NewFunction(func(L *lua.LState) int { + level := L.ToString(1) + msg := L.ToString(2) + fmt.Printf("[lua/%s] %s\n", level, msg) + return 0 + })) + + v.state.SetGlobal("json_encode", v.state.NewFunction(func(L *lua.LState) int { + val := L.CheckAny(1) + L.Push(lua.LString(fmt.Sprintf("%v", val))) + return 1 + })) + + v.state.SetGlobal("http_get", v.state.NewFunction(func(L *lua.LState) int { + url := L.ToString(1) + L.Push(lua.LString(fmt.Sprintf(`{"url":%q,"status":200,"body":"mock"}`, url))) + return 1 + })) + + v.state.SetGlobal("http_post", v.state.NewFunction(func(L *lua.LState) int { + url := L.ToString(1) + body := L.ToString(2) + L.Push(lua.LString(fmt.Sprintf(`{"url":%q,"body":%q,"status":200}`, url, body))) + return 1 + })) + + if err := v.loadAdapters(); err != nil { + return fmt.Errorf("load adapters: %w", err) + } + + return nil +} + +func (v *VM) Stop() { + if v.state != nil { + v.state.Close() + } +} + +func (v *VM) loadAdapters() error { + entries, err := os.ReadDir(v.adapterDir) + if err != nil { + if os.IsNotExist(err) { + return nil + } + return err + } + + for _, entry := range entries { + if filepath.Ext(entry.Name()) != ".lua" { + continue + } + path := filepath.Join(v.adapterDir, entry.Name()) + if err := v.LoadAdapter(path); err != nil { + fmt.Printf("[lua] load %s: %v\n", entry.Name(), err) + } + } + + return nil +} + +func (v *VM) LoadAdapter(path string) error { + v.mu.Lock() + defer v.mu.Unlock() + + data, err := os.ReadFile(path) + if err != nil { + return fmt.Errorf("read adapter: %w", err) + } + + script := string(data) + + if err := v.state.DoString(script); err != nil { + return fmt.Errorf("execute adapter script: %w", err) + } + + adapterTable := v.state.Get(-1) + v.state.Pop(1) + + tbl, ok := adapterTable.(*lua.LTable) + if !ok { + return fmt.Errorf("adapter script must return a table") + } + + name := "" + if nameVal := tbl.RawGetString("name"); nameVal != nil { + name = nameVal.String() + } + if name == "" { + name = filepath.Base(path) + } + + v.loaded[name] = tbl + fmt.Printf("[lua] loaded adapter: %s\n", name) + return nil +} + +func (v *VM) CallTransform(name string, input map[string]interface{}) (map[string]interface{}, error) { + v.mu.Lock() + adapter, ok := v.loaded[name] + v.mu.Unlock() + + if !ok { + return nil, fmt.Errorf("adapter %s not loaded", name) + } + + v.mu.Lock() + defer v.mu.Unlock() + + fn := adapter.RawGetString("transform_request") + if fn == nil { + return nil, fmt.Errorf("adapter %s missing transform_request", name) + } + + inputTable := mapToTable(v.state, input) + v.state.Push(fn) + v.state.Push(inputTable) + + if err := v.state.PCall(1, 1, nil); err != nil { + return nil, fmt.Errorf("transform_request: %w", err) + } + + result := v.state.Get(-1) + v.state.Pop(1) + + resultTable, ok := result.(*lua.LTable) + if !ok { + return nil, fmt.Errorf("transform_request must return a table") + } + + return tableToMap(resultTable), nil +} + +func (v *VM) CallResponseTransform(name string, raw []byte) ([]byte, error) { + v.mu.Lock() + adapter, ok := v.loaded[name] + v.mu.Unlock() + + if !ok { + return raw, nil + } + + v.mu.Lock() + defer v.mu.Unlock() + + fn := adapter.RawGetString("transform_response") + if fn == nil { + return raw, nil + } + + v.state.Push(fn) + v.state.Push(lua.LString(string(raw))) + + if err := v.state.PCall(1, 1, nil); err != nil { + return nil, fmt.Errorf("transform_response: %w", err) + } + + result := v.state.Get(-1) + v.state.Pop(1) + + return []byte(result.String()), nil +} + +func (v *VM) ListAdapters() []APIAdapter { + v.mu.Lock() + defer v.mu.Unlock() + + adapters := make([]APIAdapter, 0) + for name, tbl := range v.loaded { + adapter := APIAdapter{Name: name} + if v := tbl.RawGetString("version"); v != nil { + adapter.Version = v.String() + } + adapters = append(adapters, adapter) + } + return adapters +} + +func (v *VM) ReloadAll() error { + v.mu.Lock() + v.loaded = make(map[string]*lua.LTable) + v.mu.Unlock() + + if v.state != nil { + v.state.Close() + } + v.state = lua.NewState() + + return v.Start() +} + +func mapToTable(L *lua.LState, m map[string]interface{}) *lua.LTable { + tbl := L.NewTable() + for k, v := range m { + switch val := v.(type) { + case string: + tbl.RawSetString(k, lua.LString(val)) + case float64: + tbl.RawSetString(k, lua.LNumber(val)) + case int: + tbl.RawSetString(k, lua.LNumber(val)) + case bool: + tbl.RawSetString(k, lua.LBool(val)) + case map[string]interface{}: + tbl.RawSetString(k, mapToTable(L, val)) + case []interface{}: + arr := L.NewTable() + for i, item := range val { + if m, ok := item.(map[string]interface{}); ok { + arr.RawSetInt(i+1, mapToTable(L, m)) + } else { + arr.RawSetInt(i+1, lua.LString(fmt.Sprintf("%v", item))) + } + } + tbl.RawSetString(k, arr) + default: + tbl.RawSetString(k, lua.LString(fmt.Sprintf("%v", v))) + } + } + return tbl +} + +func (v *VM) writeBundledAdapters() error { + entries, err := bundledAdapters.ReadDir("adapters") + if err != nil { + return nil + } + + for _, entry := range entries { + if entry.IsDir() { + continue + } + dstPath := filepath.Join(v.adapterDir, entry.Name()) + if _, err := os.Stat(dstPath); err == nil { + continue + } + + data, err := bundledAdapters.ReadFile(filepath.Join("adapters", entry.Name())) + if err != nil { + continue + } + + if err := os.WriteFile(dstPath, data, 0644); err != nil { + return fmt.Errorf("write %s: %w", entry.Name(), err) + } + fmt.Printf("[lua] installed bundled adapter: %s\n", entry.Name()) + } + return nil +} + +func tableToMap(tbl *lua.LTable) map[string]interface{} { + result := make(map[string]interface{}) + tbl.ForEach(func(key lua.LValue, val lua.LValue) { + k := key.String() + switch v := val.(type) { + case lua.LString: + result[k] = string(v) + case lua.LNumber: + result[k] = float64(v) + case lua.LBool: + result[k] = bool(v) + case *lua.LTable: + if v.MaxN() == 0 { + result[k] = tableToMap(v) + } else { + var arr []interface{} + v.ForEach(func(_, item lua.LValue) { + if tbl, ok := item.(*lua.LTable); ok { + arr = append(arr, tableToMap(tbl)) + } else { + arr = append(arr, item.String()) + } + }) + result[k] = arr + } + } + }) + return result +} diff --git a/internal/memory/document/document.go b/internal/memory/document/document.go new file mode 100644 index 0000000..6f1117d --- /dev/null +++ b/internal/memory/document/document.go @@ -0,0 +1,388 @@ +package document + +import ( + "encoding/json" + "fmt" + "log" + "os" + "path/filepath" + "sort" + "strings" + "sync" + "time" + + "gitcode.com/JianFeeeee/HomeAgent/internal/memory/vector" +) + +// Doc — 记忆文档:由上下文提炼而来 +type Doc struct { + ID string `json:"id"` + Summary string `json:"summary"` + Content string `json:"content"` + Tags []string `json:"tags"` + Entities []string `json:"entities"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` + Source string `json:"source"` // context / graph / manual + Meta map[string]string `json:"meta,omitempty"` + AccessCount int `json:"access_count"` // 访问次数 + LastAccess time.Time `json:"last_access"` // 最后访问时间 +} + +// Store — 文档记忆存储,包含向量索引 +type Store struct { + dir string + vec *vector.Store + veczer *vector.TFIDFVectorizer + mu sync.RWMutex + + docs map[string]*Doc + summaries []string // 用于训练向量化器 + + dirty bool +} + +func NewStore(dir string) *Store { + return &Store{ + dir: dir, + vec: vector.NewStore(), + veczer: vector.NewTFIDFVectorizer(2), + docs: make(map[string]*Doc), + } +} + +func (s *Store) Start() error { + if err := os.MkdirAll(s.dir, 0755); err != nil { + return fmt.Errorf("document store dir: %w", err) + } + if err := s.loadAll(); err != nil { + log.Printf("[document memory] load error: %v", err) + } + log.Printf("[document memory] started with %d docs, %d vectors", len(s.docs), s.vec.Size()) + return nil +} + +func (s *Store) Stop() { + s.flush() +} + +// Insert 创建/更新文档 +func (s *Store) Insert(doc *Doc) error { + s.mu.Lock() + defer s.mu.Unlock() + + if doc.ID == "" { + doc.ID = fmt.Sprintf("doc_%d", time.Now().UnixNano()) + doc.CreatedAt = time.Now() + } + doc.UpdatedAt = time.Now() + doc.LastAccess = time.Now() + if doc.AccessCount == 0 { + doc.AccessCount = 1 + } + + s.docs[doc.ID] = doc + + vec := s.veczer.Vectorize(doc.Summary + " " + doc.Content) + s.vec.Insert(doc.ID, doc.Summary, vec, doc.Meta) + + // 更新训练集 + s.summaries = append(s.summaries, doc.Summary) + + s.dirty = true + return nil +} + +// ContextToDoc — 将一段上下文对话历史提炼为文档 +func (s *Store) ContextToDoc(source string, entries []ContextEntry) (*Doc, error) { + if len(entries) == 0 { + return nil, nil + } + + var parts []string + for _, e := range entries { + line := fmt.Sprintf("[%s] %s: %s", e.Timestamp.Format("15:04"), e.Source, e.Content) + if e.Response != "" { + line += fmt.Sprintf(" → %s", truncate(e.Response, 100)) + } + parts = append(parts, line) + } + content := strings.Join(parts, "\n") + + summary := summarizeEntries(entries) + tags := extractTags(entries) + entities := extractEntities(entries) + + doc := &Doc{ + ID: fmt.Sprintf("doc_%d", time.Now().UnixNano()), + Summary: summary, + Content: content, + Tags: tags, + Entities: entities, + CreatedAt: time.Now(), + UpdatedAt: time.Now(), + Source: source, + } + + if err := s.Insert(doc); err != nil { + return nil, err + } + + return doc, nil +} + +// Query — 向量相似度查询文档 +func (s *Store) Query(text string, topK int) []*Doc { + s.mu.RLock() + defer s.mu.RUnlock() + + if topK <= 0 { + topK = 5 + } + + vec := s.veczer.Vectorize(text) + results := s.vec.Search(vec, topK) + + var docs []*Doc + for _, r := range results { + if d, ok := s.docs[r.ID]; ok { + d.AccessCount++ + d.LastAccess = time.Now() + docs = append(docs, d) + } + } + return docs +} + +// Reindex — 重新训练并重建向量索引 +func (s *Store) Reindex() { + s.mu.Lock() + defer s.mu.Unlock() + + log.Printf("[document memory] reindexing %d docs", len(s.docs)) + + s.veczer.Train(s.summaries) + + s.vec = vector.NewStore() + for _, doc := range s.docs { + vec := s.veczer.Vectorize(doc.Summary + " " + doc.Content) + s.vec.Insert(doc.ID, doc.Summary, vec, doc.Meta) + } + + log.Printf("[document memory] reindex complete (%d vectors)", s.vec.Size()) +} + +func (s *Store) Stats() map[string]interface{} { + s.mu.RLock() + defer s.mu.RUnlock() + + return map[string]interface{}{ + "doc_count": len(s.docs), + "vector_count": s.vec.Size(), + "summary_count": len(s.summaries), + "dir": s.dir, + } +} + +// FindColdDocs — 查找冷文档:超过 maxAge 未访问且访问次数 <= minAccess +func (s *Store) FindColdDocs(maxAge time.Duration, minAccess int) []*Doc { + s.mu.RLock() + defer s.mu.RUnlock() + + cutoff := time.Now().Add(-maxAge) + var cold []*Doc + for _, d := range s.docs { + if d.AccessCount <= minAccess && d.LastAccess.Before(cutoff) { + cold = append(cold, d) + } + } + return cold +} + +func (s *Store) RecentDocs(n int) []*Doc { + s.mu.RLock() + defer s.mu.RUnlock() + + var list []*Doc + for _, d := range s.docs { + list = append(list, d) + } + sort.Slice(list, func(i, j int) bool { + return list[i].CreatedAt.After(list[j].CreatedAt) + }) + if len(list) > n { + list = list[:n] + } + return list +} + +// ——— internal ——— + +func (s *Store) loadAll() error { + entries, err := os.ReadDir(s.dir) + if err != nil { + return err + } + + for _, e := range entries { + if !strings.HasSuffix(e.Name(), ".json") || !strings.HasPrefix(e.Name(), "doc_") { + continue + } + path := filepath.Join(s.dir, e.Name()) + data, err := os.ReadFile(path) + if err != nil { + continue + } + var doc Doc + if err := json.Unmarshal(data, &doc); err != nil { + continue + } + s.docs[doc.ID] = &doc + s.summaries = append(s.summaries, doc.Summary) + } + + // 训练向量化器 + if len(s.summaries) > 0 { + s.veczer.Train(s.summaries) + } + + // 重建向量索引 + for _, doc := range s.docs { + vec := s.veczer.Vectorize(doc.Summary + " " + doc.Content) + s.vec.Insert(doc.ID, doc.Summary, vec, nil) + } + + return nil +} + +func (s *Store) flush() { + s.mu.Lock() + defer s.mu.Unlock() + + if !s.dirty { + return + } + + for _, doc := range s.docs { + path := filepath.Join(s.dir, doc.ID+".json") + data, err := json.MarshalIndent(doc, "", " ") + if err != nil { + continue + } + os.WriteFile(path, data, 0644) + } + s.dirty = false +} + +type ContextEntry struct { + Timestamp time.Time + Source string + Content string + Response string +} + +func summarizeEntries(entries []ContextEntry) string { + if len(entries) == 0 { + return "" + } + sources := make(map[string]int) + var topics []string + for _, e := range entries { + sources[e.Source]++ + words := extractKeywords(e.Content) + topics = append(topics, words...) + } + + summary := fmt.Sprintf("来自 %d 个来源的 %d 条对话", len(sources), len(entries)) + var srcList []string + for s := range sources { + srcList = append(srcList, s) + } + summary += " (" + strings.Join(srcList, ", ") + ")" + + if len(topics) > 0 { + seen := make(map[string]bool) + var uniq []string + for _, t := range topics { + if !seen[t] { + seen[t] = true + uniq = append(uniq, t) + } + } + if len(uniq) > 5 { + uniq = uniq[:5] + } + summary += " 涉及: " + strings.Join(uniq, ", ") + } + + return summary +} + +func extractTags(entries []ContextEntry) []string { + tagSet := make(map[string]bool) + for _, e := range entries { + for _, kw := range extractKeywords(e.Content) { + tagSet[kw] = true + } + } + var tags []string + for t := range tagSet { + if len(tags) >= 10 { + break + } + tags = append(tags, t) + } + return tags +} + +func extractEntities(entries []ContextEntry) []string { + // 简易实体提取:提取引号内的内容、粗体/标记词 + var entities []string + seen := make(map[string]bool) + for _, e := range entries { + for _, kw := range extractKeywords(e.Content) { + if len(kw) >= 2 && !seen[kw] { + seen[kw] = true + entities = append(entities, kw) + } + } + } + if len(entities) > 20 { + entities = entities[:20] + } + return entities +} + +func extractKeywords(text string) []string { + stopWords := map[string]bool{ + "的": true, "了": true, "是": true, "在": true, "有": true, + "和": true, "就": true, "不": true, "人": true, "都": true, + "一": true, "一个": true, "上": true, "也": true, "很": true, + "到": true, "说": true, "要": true, "去": true, "你": true, + "会": true, "着": true, "没有": true, "看": true, "好": true, + "自己": true, "这": true, "他": true, "她": true, "它": true, + "什么": true, "怎么": true, "为什么": true, "如何": true, + "我": true, "我们": true, "你们": true, "他们": true, "这个": true, + "那个": true, "可以": true, "吗": true, "吧": true, "啊": true, + } + + var keywords []string + runes := []rune(text) + + // bi-gram + for i := 0; i < len(runes)-1; i++ { + word := string(runes[i : i+2]) + if !stopWords[word] && len(strings.TrimSpace(word)) == len(word) { + keywords = append(keywords, word) + } + } + return keywords +} + +func truncate(s string, max int) string { + runes := []rune(s) + if len(runes) > max { + return string(runes[:max]) + "..." + } + return s +} diff --git a/internal/memory/graph.go b/internal/memory/graph.go new file mode 100644 index 0000000..bfc6aa9 --- /dev/null +++ b/internal/memory/graph.go @@ -0,0 +1,555 @@ +package memory + +import ( + "database/sql" + "fmt" + "sync" + "time" + + _ "github.com/mattn/go-sqlite3" +) + +type Entity struct { + ID int64 `json:"id"` + Name string `json:"name"` + Type string `json:"type"` + MentionCount int `json:"mention_count"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` +} + +type Relation struct { + ID int64 `json:"id"` + SourceID int64 `json:"source_id"` + TargetID int64 `json:"target_id"` + SourceName string `json:"source_name"` + TargetName string `json:"target_name"` + RelationType string `json:"relation_type"` + Confidence float64 `json:"confidence"` + Status string `json:"status"` + SessionID string `json:"session_id"` + TurnID int `json:"turn_id"` + CreatedAt time.Time `json:"created_at"` + DateBucket string `json:"date_bucket"` +} + +type Triple struct { + Subject string `json:"subject"` + Relation string `json:"relation"` + Object string `json:"object"` + Confidence float64 `json:"confidence,omitempty"` + SubjectType string `json:"subject_type,omitempty"` + ObjectType string `json:"object_type,omitempty"` +} + +type GraphDB struct { + db *sql.DB + mu sync.RWMutex + dbPath string +} + +func NewGraphDB(dbPath string) (*GraphDB, error) { + db, err := sql.Open("sqlite3", dbPath+"?_journal_mode=WAL&_foreign_keys=on") + if err != nil { + return nil, fmt.Errorf("open graph db: %w", err) + } + + g := &GraphDB{db: db, dbPath: dbPath} + if err := g.initSchema(); err != nil { + return nil, fmt.Errorf("init schema: %w", err) + } + + return g, nil +} + +func (g *GraphDB) initSchema() error { + g.mu.Lock() + defer g.mu.Unlock() + + tx, err := g.db.Begin() + if err != nil { + return err + } + defer tx.Rollback() + + schemas := []string{ + `CREATE TABLE IF NOT EXISTS entities ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + name TEXT UNIQUE NOT NULL, + type TEXT DEFAULT 'Concept', + mention_count INTEGER DEFAULT 1, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP + )`, + `CREATE TABLE IF NOT EXISTS relations ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + source_id INTEGER NOT NULL, + target_id INTEGER NOT NULL, + relation_type TEXT NOT NULL, + confidence REAL DEFAULT 1.0, + status TEXT DEFAULT 'active', + session_id TEXT, + turn_id INTEGER DEFAULT 0, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + date_bucket TEXT, + FOREIGN KEY (source_id) REFERENCES entities(id), + FOREIGN KEY (target_id) REFERENCES entities(id) + )`, + `CREATE INDEX IF NOT EXISTS idx_entity_name ON entities(name)`, + `CREATE INDEX IF NOT EXISTS idx_entity_type ON entities(type)`, + `CREATE INDEX IF NOT EXISTS idx_relation_source ON relations(source_id)`, + `CREATE INDEX IF NOT EXISTS idx_relation_target ON relations(target_id)`, + `CREATE INDEX IF NOT EXISTS idx_relation_type ON relations(relation_type)`, + `CREATE INDEX IF NOT EXISTS idx_relation_status ON relations(status)`, + `CREATE INDEX IF NOT EXISTS idx_relation_session ON relations(session_id)`, + } + + for _, s := range schemas { + if _, err := tx.Exec(s); err != nil { + return fmt.Errorf("schema exec: %w", err) + } + } + + return tx.Commit() +} + +func (g *GraphDB) Commit(triples []Triple, sessionID string, turnID int) (int, int, error) { + g.mu.Lock() + defer g.mu.Unlock() + + tx, err := g.db.Begin() + if err != nil { + return 0, 0, err + } + defer tx.Rollback() + + entitiesCreated := 0 + relationsCreated := 0 + dateBucket := time.Now().Format("2006-01-02") + + for _, t := range triples { + if t.Subject == "" || t.Relation == "" || t.Object == "" { + continue + } + + subjType := t.SubjectType + if subjType == "" { + subjType = "Concept" + } + objType := t.ObjectType + if objType == "" { + objType = "Concept" + } + confidence := t.Confidence + if confidence == 0 { + confidence = 1.0 + } + + ec, err := g.upsertEntity(tx, t.Subject, subjType) + if err != nil { + return 0, 0, err + } + entitiesCreated += ec + + ec, err = g.upsertEntity(tx, t.Object, objType) + if err != nil { + return 0, 0, err + } + entitiesCreated += ec + + var sourceID, targetID int64 + err = tx.QueryRow("SELECT id FROM entities WHERE name = ?", t.Subject).Scan(&sourceID) + if err != nil { + return 0, 0, err + } + err = tx.QueryRow("SELECT id FROM entities WHERE name = ?", t.Object).Scan(&targetID) + if err != nil { + return 0, 0, err + } + + _, err = tx.Exec( + `INSERT INTO relations (source_id, target_id, relation_type, confidence, session_id, turn_id, date_bucket) + VALUES (?, ?, ?, ?, ?, ?, ?)`, + sourceID, targetID, t.Relation, confidence, sessionID, turnID, dateBucket, + ) + if err != nil { + return 0, 0, err + } + relationsCreated++ + } + + if err := tx.Commit(); err != nil { + return 0, 0, err + } + + return entitiesCreated, relationsCreated, nil +} + +func (g *GraphDB) upsertEntity(tx *sql.Tx, name string, entityType string) (int, error) { + result, err := tx.Exec( + `INSERT INTO entities (name, type) VALUES (?, ?) + ON CONFLICT(name) DO UPDATE SET + mention_count = mention_count + 1, + updated_at = CURRENT_TIMESTAMP`, + name, entityType, + ) + if err != nil { + return 0, err + } + rows, _ := result.RowsAffected() + if rows > 0 { + return 1, nil + } + return 0, nil +} + +type RecallResult struct { + Entities []Entity `json:"entities"` + Relations []Relation `json:"relations"` +} + +func (g *GraphDB) Recall(keywords []string, seedEntities []string, depth int, sessionFilter string) (*RecallResult, error) { + g.mu.RLock() + defer g.mu.RUnlock() + + result := &RecallResult{} + + if len(keywords) == 0 && len(seedEntities) == 0 { + rows, err := g.db.Query( + `SELECT id, name, type, mention_count, created_at, updated_at + FROM entities ORDER BY mention_count DESC LIMIT 50`, + ) + if err != nil { + return nil, err + } + defer rows.Close() + + for rows.Next() { + var e Entity + if err := rows.Scan(&e.ID, &e.Name, &e.Type, &e.MentionCount, &e.CreatedAt, &e.UpdatedAt); err != nil { + return nil, err + } + result.Entities = append(result.Entities, e) + } + + relRows, err := g.db.Query( + `SELECT r.id, r.source_id, r.target_id, e1.name, e2.name, + r.relation_type, r.confidence, r.status, r.session_id, + r.turn_id, r.created_at, COALESCE(r.date_bucket, '') + FROM relations r + JOIN entities e1 ON r.source_id = e1.id + JOIN entities e2 ON r.target_id = e2.id + WHERE r.status = 'active' + ORDER BY r.created_at DESC LIMIT 30`, + ) + if err != nil { + return nil, err + } + defer relRows.Close() + for relRows.Next() { + var rel Relation + if err := relRows.Scan(&rel.ID, &rel.SourceID, &rel.TargetID, + &rel.SourceName, &rel.TargetName, &rel.RelationType, + &rel.Confidence, &rel.Status, &rel.SessionID, + &rel.TurnID, &rel.CreatedAt, &rel.DateBucket); err != nil { + return nil, err + } + result.Relations = append(result.Relations, rel) + } + + return result, nil + } + + entityIDs := make(map[int64]bool) + + for _, kw := range keywords { + rows, err := g.db.Query( + `SELECT id, name, type, mention_count, created_at, updated_at + FROM entities WHERE LOWER(name) LIKE ?`, + "%"+kw+"%", + ) + if err != nil { + return nil, err + } + defer rows.Close() + + for rows.Next() { + var e Entity + if err := rows.Scan(&e.ID, &e.Name, &e.Type, &e.MentionCount, &e.CreatedAt, &e.UpdatedAt); err != nil { + return nil, err + } + if !entityIDs[e.ID] { + entityIDs[e.ID] = true + result.Entities = append(result.Entities, e) + } + } + } + + for _, se := range seedEntities { + row := g.db.QueryRow( + `SELECT id, name, type, mention_count, created_at, updated_at + FROM entities WHERE name = ?`, se) + var e Entity + if err := row.Scan(&e.ID, &e.Name, &e.Type, &e.MentionCount, &e.CreatedAt, &e.UpdatedAt); err == nil { + if !entityIDs[e.ID] { + entityIDs[e.ID] = true + result.Entities = append(result.Entities, e) + } + } + } + + if len(entityIDs) == 0 { + return result, nil + } + + for depthLevel := 0; depthLevel < depth; depthLevel++ { + ids := make([]interface{}, 0, len(entityIDs)) + for id := range entityIDs { + ids = append(ids, id) + } + + if len(ids) == 0 { + break + } + + query := fmt.Sprintf( + `SELECT r.id, r.source_id, r.target_id, e1.name, e2.name, + r.relation_type, r.confidence, r.status, r.session_id, + r.turn_id, r.created_at, COALESCE(r.date_bucket, '') + FROM relations r + JOIN entities e1 ON r.source_id = e1.id + JOIN entities e2 ON r.target_id = e2.id + WHERE (r.source_id IN (%s) OR r.target_id IN (%s)) + AND r.status = 'active'`, + placeholders(len(ids)), + placeholders(len(ids)), + ) + allIDs := append(ids, ids...) + + if sessionFilter != "" { + query += " AND r.session_id = ?" + allIDs = append(allIDs, sessionFilter) + } + + relRows, err := g.db.Query(query, allIDs...) + if err != nil { + return nil, err + } + defer relRows.Close() + + newIDs := make(map[int64]bool) + for relRows.Next() { + var rel Relation + if err := relRows.Scan(&rel.ID, &rel.SourceID, &rel.TargetID, + &rel.SourceName, &rel.TargetName, &rel.RelationType, + &rel.Confidence, &rel.Status, &rel.SessionID, + &rel.TurnID, &rel.CreatedAt, &rel.DateBucket); err != nil { + return nil, err + } + result.Relations = append(result.Relations, rel) + + if !entityIDs[rel.SourceID] { + newIDs[rel.SourceID] = true + } + if !entityIDs[rel.TargetID] { + newIDs[rel.TargetID] = true + } + } + + if len(newIDs) == 0 { + break + } + + ids2 := make([]interface{}, 0, len(newIDs)) + for id := range newIDs { + ids2 = append(ids2, id) + } + + eRows, err := g.db.Query( + fmt.Sprintf( + `SELECT id, name, type, mention_count, created_at, updated_at + FROM entities WHERE id IN (%s)`, placeholders(len(ids2))), + ids2..., + ) + if err != nil { + return nil, err + } + defer eRows.Close() + + for eRows.Next() { + var e Entity + if err := eRows.Scan(&e.ID, &e.Name, &e.Type, &e.MentionCount, &e.CreatedAt, &e.UpdatedAt); err != nil { + return nil, err + } + if !entityIDs[e.ID] { + entityIDs[e.ID] = true + result.Entities = append(result.Entities, e) + } + } + + for id := range newIDs { + entityIDs[id] = true + } + } + + return result, nil +} + +func (g *GraphDB) Purge(criteria map[string]string, mode string) (int, error) { + g.mu.Lock() + defer g.mu.Unlock() + + conds := []string{"r.status = 'active'"} + args := []interface{}{} + + if v, ok := criteria["subject_contains"]; ok { + rows, err := g.db.Query("SELECT id FROM entities WHERE name LIKE ?", "%"+v+"%") + if err != nil { + return 0, err + } + defer rows.Close() + var ids []interface{} + for rows.Next() { + var id int64 + rows.Scan(&id) + ids = append(ids, id) + } + if len(ids) > 0 { + conds = append(conds, fmt.Sprintf("r.source_id IN (%s)", placeholders(len(ids)))) + args = append(args, ids...) + } + } + + if v, ok := criteria["target_contains"]; ok { + rows, err := g.db.Query("SELECT id FROM entities WHERE name LIKE ?", "%"+v+"%") + if err != nil { + return 0, err + } + defer rows.Close() + var ids []interface{} + for rows.Next() { + var id int64 + rows.Scan(&id) + ids = append(ids, id) + } + if len(ids) > 0 { + conds = append(conds, fmt.Sprintf("r.target_id IN (%s)", placeholders(len(ids)))) + args = append(args, ids...) + } + } + + if v, ok := criteria["relation_type"]; ok { + conds = append(conds, "r.relation_type = ?") + args = append(args, v) + } + + if v, ok := criteria["session_id"]; ok { + conds = append(conds, "r.session_id = ?") + args = append(args, v) + } + + if len(conds) == 1 { + return 0, fmt.Errorf("no criteria provided") + } + + where := "" + for i, c := range conds { + if i == 0 { + where = c + } else { + where += " AND " + c + } + } + + if mode == "hard" { + result, err := g.db.Exec( + fmt.Sprintf(`DELETE FROM relations WHERE %s`, where), args...) + if err != nil { + return 0, err + } + n, _ := result.RowsAffected() + + g.db.Exec(`DELETE FROM entities WHERE id NOT IN ( + SELECT DISTINCT source_id FROM relations + UNION SELECT DISTINCT target_id FROM relations)`) + + return int(n), nil + } + + result, err := g.db.Exec( + fmt.Sprintf(`UPDATE relations SET status = 'deleted', updated_at = CURRENT_TIMESTAMP WHERE %s`, where), + args..., + ) + if err != nil { + return 0, err + } + n, _ := result.RowsAffected() + return int(n), nil +} + +func (g *GraphDB) Introspect() (map[string]interface{}, error) { + g.mu.RLock() + defer g.mu.RUnlock() + + var entityCount, relationCount int + g.db.QueryRow("SELECT COUNT(*) FROM entities").Scan(&entityCount) + g.db.QueryRow("SELECT COUNT(*) FROM relations WHERE status = 'active'").Scan(&relationCount) + + hotspots := []map[string]interface{}{} + rows, err := g.db.Query( + `SELECT name, mention_count, type FROM entities ORDER BY mention_count DESC LIMIT 10`, + ) + if err == nil { + defer rows.Close() + for rows.Next() { + var name, etype string + var count int + if err := rows.Scan(&name, &count, &etype); err == nil { + hotspots = append(hotspots, map[string]interface{}{ + "name": name, "count": count, "type": etype, + }) + } + } + } + + return map[string]interface{}{ + "entity_count": entityCount, + "relation_count": relationCount, + "memory_hotspots": hotspots, + }, nil +} + +func (g *GraphDB) Archive(days int) (int, error) { + g.mu.Lock() + defer g.mu.Unlock() + + result, err := g.db.Exec( + `UPDATE relations SET status = 'archived', updated_at = CURRENT_TIMESTAMP + WHERE status = 'active' AND created_at < datetime('now', ?)`, + fmt.Sprintf("-%d days", days), + ) + if err != nil { + return 0, err + } + n, _ := result.RowsAffected() + return int(n), nil +} + +func (g *GraphDB) Close() error { + return g.db.Close() +} + +func placeholders(n int) string { + if n <= 0 { + return "NULL" + } + b := make([]byte, 0, n*2-1) + for i := 0; i < n; i++ { + if i > 0 { + b = append(b, ',') + } + b = append(b, '?') + } + return string(b) +} diff --git a/internal/memory/indexer.go b/internal/memory/indexer.go new file mode 100644 index 0000000..2ce0734 --- /dev/null +++ b/internal/memory/indexer.go @@ -0,0 +1,329 @@ +package memory + +import ( + "fmt" + "log" + "strings" + "sync" + + "gitcode.com/JianFeeeee/HomeAgent/internal/memory/vector" +) + +type Indexer struct { + db *GraphDB + vec *vector.Store + veczer *vector.TFIDFVectorizer + mu sync.RWMutex + trained bool +} + +func NewIndexer(db *GraphDB) *Indexer { + return &Indexer{ + db: db, + vec: vector.NewStore(), + veczer: vector.NewTFIDFVectorizer(2), + } +} + +// Sync 从图数据库中同步实体名到向量索引 +func (idx *Indexer) Sync() error { + idx.mu.Lock() + defer idx.mu.Unlock() + + if idx.db == nil { + return nil + } + + result, err := idx.db.Recall(nil, nil, 1, "") + if err != nil || result == nil { + return err + } + + // 收集实体名 + var names []string + for _, e := range result.Entities { + names = append(names, e.Name) + } + + if len(names) == 0 { + return nil + } + + // 训练向量化器 + idx.veczer.Train(names) + + // 重建向量索引 + idx.vec = vector.NewStore() + for _, e := range result.Entities { + vec := idx.veczer.Vectorize(e.Name) + idx.vec.Insert(fmt.Sprintf("entity_%d", e.ID), e.Name, vec, map[string]string{ + "type": "entity", + "name": e.Name, + }) + } + + idx.trained = true + log.Printf("[indexer] synced %d entities to vector index", len(names)) + return nil +} + +type InjectedContext struct { + Entities []Entity `json:"entities"` + Relations []Relation `json:"relations"` + Summary string `json:"summary"` + TokenEstimate int `json:"token_estimate"` +} + +func (idx *Indexer) BuildContext(userInput string) *InjectedContext { + if idx.db == nil { + return &InjectedContext{Summary: ""} + } + + // 1. 向量搜索:从实体名向量索引中找到相关实体 + vectorEntities := idx.vectorSearchEntities(userInput) + + // 2. 关键词搜索:已有逻辑 + keywords := extractKeywords(userInput) + if len(keywords) == 0 && len(vectorEntities) == 0 { + keywords = []string{userInput} + } + + // 合并关键词和向量找到的实体名 + seedNames := make([]string, 0, len(vectorEntities)) + for _, e := range vectorEntities { + seedNames = append(seedNames, e.Name) + } + allKeywords := append(keywords, seedNames...) + + result, err := idx.db.Recall(allKeywords, nil, 2, "") + if err != nil || result == nil { + return &InjectedContext{Summary: ""} + } + + ctx := &InjectedContext{ + Entities: result.Entities, + Relations: nil, + } + + if len(result.Entities) > 0 { + summary := buildIndexSummary(result.Entities) + ctx.Summary = summary + ctx.TokenEstimate = estimateTokens(summary) + len(result.Entities)*8 + } else { + ctx.Summary = "" + } + + return ctx +} + +// vectorSearchEntities 在实体名向量索引中搜索 +func (idx *Indexer) vectorSearchEntities(query string) []Entity { + idx.mu.RLock() + defer idx.mu.RUnlock() + + if !idx.trained || idx.vec.Size() == 0 { + return nil + } + + queryVec := idx.veczer.Vectorize(query) + results := idx.vec.Search(queryVec, 5) + + var entities []Entity + for _, r := range results { + if r.Meta != nil && r.Meta["type"] == "entity" { + entities = append(entities, Entity{Name: r.Meta["name"]}) + } + } + return entities +} + +func (idx *Indexer) BuildToolPrompt() string { + return `## 图记忆工具 + +你有以下工具可以操作长期图记忆系统: + +### memory_recall +检索与关键词相关的实体和关系。 +参数: +- query_intent: 查询关键词,逗号分隔 +- depth: 遍历深度(默认2) + +### memory_commit +将三元组写入图记忆。 +参数: +- triples: [{"subject": "实体名", "relation": "关系类型", "object": "目标实体"}] + +### memory_introspect +查看记忆统计信息。 + +### memory_purge +删除或修正记忆。 +参数: +- criteria: {"subject_contains": "...", "relation_type": "..."} +- mode: "soft" | "supersede" + +使用方法:在推理过程中调用对应的 tool,系统会自动执行并返回结果。` +} + +func (idx *Indexer) FormatContext(ctx *InjectedContext) string { + if ctx == nil || len(ctx.Entities) == 0 { + return "" + } + + var b strings.Builder + b.WriteString("【记忆索引】") + + if ctx.Summary != "" { + b.WriteString(" ") + b.WriteString(ctx.Summary) + } + + b.WriteString(fmt.Sprintf(" 索引: ")) + for i, e := range ctx.Entities { + if i >= 5 { + b.WriteString("…") + break + } + if i > 0 { + b.WriteString(", ") + } + b.WriteString(e.Name) + if e.Type != "Concept" { + b.WriteString("(" + e.Type + ")") + } + } + + b.WriteString(" | 需更多细节请用 memory_recall 查询") + return b.String() +} + +func (idx *Indexer) GetToolDefinitions() []map[string]interface{} { + return []map[string]interface{}{ + { + "type": "function", + "function": map[string]interface{}{ + "name": "memory_recall", + "description": "检索图记忆。输入查询意图关键词,返回相关实体和关系。", + "parameters": map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "query_intent": map[string]interface{}{ + "type": "string", + "description": "查询意图,支持逗号分隔多个关键词", + }, + "depth": map[string]interface{}{ + "type": "integer", + "description": "遍历深度,默认2", + "default": 2, + }, + }, + "required": []string{"query_intent"}, + }, + }, + }, + { + "type": "function", + "function": map[string]interface{}{ + "name": "memory_commit", + "description": "写入图记忆。将三元组列表写入长期记忆。", + "parameters": map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "triples": map[string]interface{}{ + "type": "array", + "description": "三元组列表", + "items": map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "subject": map[string]interface{}{"type": "string"}, + "relation": map[string]interface{}{"type": "string"}, + "object": map[string]interface{}{"type": "string"}, + }, + "required": []string{"subject", "relation", "object"}, + }, + }, + }, + "required": []string{"triples"}, + }, + }, + }, + { + "type": "function", + "function": map[string]interface{}{ + "name": "memory_introspect", + "description": "查看图记忆统计信息:实体数量、关系数量、热点实体。", + "parameters": map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{}, + }, + }, + }, + } +} + +func extractKeywords(input string) []string { + stopWords := map[string]bool{ + "的": true, "了": true, "是": true, "在": true, "有": true, + "和": true, "就": true, "不": true, "人": true, "都": true, + "一": true, "一个": true, "上": true, "也": true, "很": true, + "到": true, "说": true, "要": true, "去": true, "你": true, + "会": true, "着": true, "没有": true, "看": true, "好": true, + "自己": true, "这": true, "他": true, "她": true, "它": true, + "什么": true, "怎么": true, "为什么": true, "如何": true, + } + + var keywords []string + seen := make(map[string]bool) + + runes := []rune(input) + + bigram := []rune{} + for _, r := range runes { + bigram = append(bigram, r) + if len(bigram) >= 2 { + word := string(bigram) + if !stopWords[word] && !seen[word] { + seen[word] = true + keywords = append(keywords, word) + } + bigram = bigram[1:] + } + } + + if len(keywords) == 0 && len(runes) > 0 { + keywords = []string{string(runes)} + } + + if len(keywords) > 5 { + keywords = keywords[:5] + } + + return keywords +} + +func buildIndexSummary(entities []Entity) string { + if len(entities) == 0 { + return "" + } + + var b strings.Builder + b.WriteString(fmt.Sprintf("关联 %d 个记忆实体", len(entities))) + + topN := 3 + if len(entities) < topN { + topN = len(entities) + } + b.WriteString(",高频:") + for i := 0; i < topN; i++ { + if i > 0 { + b.WriteString("、") + } + b.WriteString(entities[i].Name) + } + + return b.String() +} + +func estimateTokens(s string) int { + return len(s) / 2 +} diff --git a/internal/memory/pipeline/pipeline.go b/internal/memory/pipeline/pipeline.go new file mode 100644 index 0000000..432ec9b --- /dev/null +++ b/internal/memory/pipeline/pipeline.go @@ -0,0 +1,291 @@ +package pipeline + +import ( + "context" + "fmt" + "log" + "os" + "path/filepath" + "sync" + "time" + + "gitcode.com/JianFeeeee/HomeAgent/internal/memory" +) + +type RawRecord struct { + ID int64 `json:"id"` + SessionID string `json:"session_id"` + Role string `json:"role"` + Content string `json:"content"` + CreatedAt time.Time `json:"created_at"` + Distilled bool `json:"distilled"` +} + +type DistillerConfig struct { + Interval time.Duration `json:"interval"` + RetentionDays int `json:"retention_days"` + BatchSize int `json:"batch_size"` +} + +type Distiller struct { + mu sync.Mutex + db *memory.GraphDB + rawPath string + records []RawRecord + nextID int64 + cfg DistillerConfig + ctx context.Context + cancel context.CancelFunc + onMemory func(input, response string) +} + +func NewDistiller(db *memory.GraphDB, dataDir string, cfg DistillerConfig) *Distiller { + ctx, cancel := context.WithCancel(context.Background()) + return &Distiller{ + db: db, + rawPath: filepath.Join(dataDir, "memory", "raw"), + cfg: cfg, + ctx: ctx, + cancel: cancel, + } +} + +func (d *Distiller) OnMemoryCandidate(fn func(input, response string)) { + d.onMemory = fn +} + +func (d *Distiller) Start() { + if err := os.MkdirAll(d.rawPath, 0755); err != nil { + log.Printf("[memory] create raw path: %v", err) + } + d.loadExisting() + log.Printf("[memory] distiller started (interval: %v, retention: %d days)", d.cfg.Interval, d.cfg.RetentionDays) + go d.distillLoop() +} + +func (d *Distiller) Stop() { + d.cancel() + d.flush() +} + +func (d *Distiller) Append(sessionID string, role string, content string) { + d.mu.Lock() + defer d.mu.Unlock() + d.nextID++ + d.records = append(d.records, RawRecord{ + ID: d.nextID, SessionID: sessionID, Role: role, + Content: content, CreatedAt: time.Now(), + }) +} + +func (d *Distiller) flush() { + d.mu.Lock() + defer d.mu.Unlock() + if len(d.records) == 0 { + return + } + path := filepath.Join(d.rawPath, fmt.Sprintf("raw_%d.jsonl", time.Now().UnixNano())) + f, err := os.Create(path) + if err != nil { + log.Printf("[memory] flush error: %v", err) + return + } + defer f.Close() + for _, r := range d.records { + line := fmt.Sprintf("%d\t%s\t%s\t%s\t%d\n", r.ID, r.SessionID, r.Role, r.Content, r.CreatedAt.Unix()) + f.WriteString(line) + } +} + +func (d *Distiller) loadExisting() { + entries, err := os.ReadDir(d.rawPath) + if err != nil { + return + } + for _, entry := range entries { + if filepath.Ext(entry.Name()) != ".jsonl" { + continue + } + path := filepath.Join(d.rawPath, entry.Name()) + data, err := os.ReadFile(path) + if err != nil { + continue + } + for _, line := range parseLines(string(data)) { + parts := splitLine(line) + if len(parts) >= 4 { + d.records = append(d.records, RawRecord{ + ID: d.nextID, SessionID: parts[1], Role: parts[2], Content: parts[3], + }) + d.nextID++ + } + } + } +} + +func (d *Distiller) distillLoop() { + ticker := time.NewTicker(d.cfg.Interval) + defer ticker.Stop() + for { + select { + case <-ticker.C: + d.distillOnce() + case <-d.ctx.Done(): + return + } + } +} + +func (d *Distiller) distillOnce() { + d.mu.Lock() + cutoff := time.Now().AddDate(0, 0, -d.cfg.RetentionDays) + var toDistill []RawRecord + var remaining []RawRecord + for _, r := range d.records { + if r.CreatedAt.Before(cutoff) && !r.Distilled { + toDistill = append(toDistill, r) + } else { + remaining = append(remaining, r) + } + } + d.records = remaining + d.mu.Unlock() + + if len(toDistill) == 0 { + return + } + + batchSize := d.cfg.BatchSize + if batchSize <= 0 { + batchSize = 50 + } + for i := 0; i < len(toDistill); i += batchSize { + end := i + batchSize + if end > len(toDistill) { + end = len(toDistill) + } + d.distillBatch(toDistill[i:end]) + } + d.cleanupRawFiles() + log.Printf("[memory] distilled %d records", len(toDistill)) +} + +func (d *Distiller) distillBatch(batch []RawRecord) { + var userContent, assistantContent string + sessionIDs := make(map[string]bool) + for _, r := range batch { + sessionIDs[r.SessionID] = true + if r.Role == "user" { + userContent += r.Content + " " + } else { + assistantContent += r.Content + " " + } + } + triples := extractKeyTriples(userContent, assistantContent) + if len(triples) > 0 { + sessionID := "" + for sid := range sessionIDs { + sessionID = sid + break + } + if _, _, err := d.db.Commit(triples, sessionID, 0); err != nil { + log.Printf("[memory] distill commit: %v", err) + } + } +} + +func (d *Distiller) cleanupRawFiles() { + entries, err := os.ReadDir(d.rawPath) + if err != nil { + return + } + cutoff := time.Now().AddDate(0, 0, -(d.cfg.RetentionDays + 1)) + for _, entry := range entries { + info, err := entry.Info() + if err != nil { + continue + } + if info.ModTime().Before(cutoff) { + os.Remove(filepath.Join(d.rawPath, entry.Name())) + } + } +} + +func extractKeyTriples(userContent, assistantContent string) []memory.Triple { + var triples []memory.Triple + if len(userContent) > 0 && len(userContent) < 500 { + triples = append(triples, memory.Triple{Subject: "用户", Relation: "提及", Object: truncate(userContent, 200)}) + } + if len(assistantContent) > 0 && len(assistantContent) < 500 { + triples = append(triples, memory.Triple{Subject: "AI", Relation: "回应", Object: truncate(assistantContent, 200)}) + } + return triples +} + +func truncate(s string, max int) string { + if len(s) > max { + return s[:max] + "..." + } + return s +} + +func parseLines(data string) []string { + var lines []string + current := "" + for _, ch := range data { + if ch == '\n' { + if current != "" { + lines = append(lines, current) + } + current = "" + } else { + current += string(ch) + } + } + if current != "" { + lines = append(lines, current) + } + return lines +} + +func splitLine(line string) []string { + var parts []string + current := "" + for _, ch := range line { + if ch == '\t' { + parts = append(parts, current) + current = "" + } else { + current += string(ch) + } + } + if current != "" { + parts = append(parts, current) + } + return parts +} + +func (d *Distiller) GetRecentRecords(limit int) []RawRecord { + d.mu.Lock() + defer d.mu.Unlock() + n := len(d.records) + if n == 0 { + return nil + } + if limit > 0 && limit < n { + n = limit + } + result := make([]RawRecord, n) + copy(result, d.records[len(d.records)-n:]) + return result +} + +func (d *Distiller) Stats() map[string]interface{} { + d.mu.Lock() + defer d.mu.Unlock() + return map[string]interface{}{ + "raw_records": len(d.records), + "interval": d.cfg.Interval.String(), + "retention_days": d.cfg.RetentionDays, + } +} diff --git a/internal/memory/text/text.go b/internal/memory/text/text.go new file mode 100644 index 0000000..9111075 --- /dev/null +++ b/internal/memory/text/text.go @@ -0,0 +1,241 @@ +package text + +import ( + "bufio" + "encoding/json" + "fmt" + "log" + "os" + "path/filepath" + "sort" + "strings" + "sync" + "time" +) + +// Event — 原始 I/O 事件记录,写入 JSONL +type Event struct { + Timestamp int64 `json:"ts"` + Source string `json:"source"` + Input string `json:"input"` + Response string `json:"response,omitempty"` + ToolsUsed []string `json:"tools_used,omitempty"` + AgentID string `json:"agent_id,omitempty"` +} + +// Memory — 文本记忆:追加写 JSONL,按时间/大小旋转 +type Memory struct { + dir string + interval time.Duration + maxSize int64 + + mu sync.Mutex + current *os.File + encoder *json.Encoder + created time.Time + size int64 + + stopCh chan struct{} +} + +type Option func(*Memory) + +func WithRotationInterval(d time.Duration) Option { + return func(m *Memory) { m.interval = d } +} + +func WithMaxSizeBytes(n int64) Option { + return func(m *Memory) { m.maxSize = n } +} + +func New(dir string, opts ...Option) *Memory { + m := &Memory{ + dir: dir, + interval: 24 * time.Hour, + maxSize: 10 * 1024 * 1024, + stopCh: make(chan struct{}), + } + for _, opt := range opts { + opt(m) + } + return m +} + +func (m *Memory) Start() error { + if err := os.MkdirAll(m.dir, 0755); err != nil { + return fmt.Errorf("text memory dir: %w", err) + } + if err := m.openCurrent(); err != nil { + return err + } + go m.rotationLoop() + return nil +} + +func (m *Memory) Stop() { + close(m.stopCh) + m.mu.Lock() + if m.current != nil { + m.current.Close() + } + m.mu.Unlock() +} + +func (m *Memory) Append(evt Event) error { + m.mu.Lock() + defer m.mu.Unlock() + + if m.needRotate() { + m.rotateLocked() + } + + if err := m.encoder.Encode(evt); err != nil { + return fmt.Errorf("encode event: %w", err) + } + m.current.Sync() + return nil +} + +func (m *Memory) needRotate() bool { + return time.Since(m.created) > m.interval || m.size > m.maxSize +} + +func (m *Memory) rotateLocked() { + if m.current != nil { + m.current.Close() + } + m.openCurrent() +} + +func (m *Memory) openCurrent() error { + now := time.Now() + name := fmt.Sprintf("text_%s.jsonl", now.Format("2006-01-02_15-04-05")) + path := filepath.Join(m.dir, name) + + f, err := os.OpenFile(path, os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0644) + if err != nil { + return fmt.Errorf("open text log %s: %w", path, err) + } + + stat, _ := f.Stat() + m.current = f + m.encoder = json.NewEncoder(f) + m.created = now + m.size = stat.Size() + return nil +} + +func (m *Memory) rotationLoop() { + ticker := time.NewTicker(m.interval / 2) + defer ticker.Stop() + for { + select { + case <-ticker.C: + m.mu.Lock() + if m.needRotate() { + m.rotateLocked() + log.Printf("[text memory] rotated log file") + } + m.mu.Unlock() + case <-m.stopCh: + return + } + } +} + +// Replay — 从 JSONL 文件流式回放事件 +func (m *Memory) Replay(fn func(Event) error) error { + m.mu.Lock() + files, err := m.listFiles() + m.mu.Unlock() + if err != nil { + return err + } + + for _, fpath := range files { + if err := m.replayFile(fpath, fn); err != nil { + return err + } + } + return nil +} + +func (m *Memory) replayFile(path string, fn func(Event) error) error { + f, err := os.Open(path) + if err != nil { + return err + } + defer f.Close() + + scanner := bufio.NewScanner(f) + scanner.Buffer(make([]byte, 1024*1024), 1024*1024) + for scanner.Scan() { + line := strings.TrimSpace(scanner.Text()) + if line == "" { + continue + } + var evt Event + if err := json.Unmarshal([]byte(line), &evt); err != nil { + continue + } + if err := fn(evt); err != nil { + return err + } + } + return scanner.Err() +} + +func (m *Memory) listFiles() ([]string, error) { + entries, err := os.ReadDir(m.dir) + if err != nil { + return nil, err + } + + var files []string + for _, e := range entries { + if strings.HasPrefix(e.Name(), "text_") && strings.HasSuffix(e.Name(), ".jsonl") { + files = append(files, filepath.Join(m.dir, e.Name())) + } + } + sort.Strings(files) + return files, nil +} + +// RecentEvents — 返回最近 n 条事件(跨所有文件的最新事件) +func (m *Memory) RecentEvents(n int) ([]Event, error) { + var all []Event + err := m.Replay(func(evt Event) error { + all = append(all, evt) + return nil + }) + if err != nil { + return nil, err + } + if len(all) > n { + all = all[len(all)-n:] + } + return all, nil +} + +func (m *Memory) FileCount() int { + m.mu.Lock() + defer m.mu.Unlock() + files, err := m.listFiles() + if err != nil { + return 0 + } + return len(files) +} + +func (m *Memory) Stats() map[string]interface{} { + m.mu.Lock() + defer m.mu.Unlock() + files, _ := m.listFiles() + return map[string]interface{}{ + "file_count": len(files), + "current_size": m.size, + "rotation_bytes": m.maxSize, + "rotation_interval": m.interval.String(), + "dir": m.dir, + } +} diff --git a/internal/memory/vector/store.go b/internal/memory/vector/store.go new file mode 100644 index 0000000..3f3b422 --- /dev/null +++ b/internal/memory/vector/store.go @@ -0,0 +1,302 @@ +package vector + +import ( + "math" + "sort" + "strings" + "sync" +) + +// Vectorizer 接口:将文本转为向量 +type Vectorizer interface { + Vectorize(text string) Vector +} + +// Vector 是带权特征映射:feature → weight +type Vector map[string]float64 + +// Store 向量存储,支持近似查询 +type Store struct { + mu sync.RWMutex + docs []DocVector + dim int + index *InvertedIndex +} + +type DocVector struct { + ID string + Vector Vector + Text string + Meta map[string]string +} + +func NewStore() *Store { + return &Store{ + index: NewInvertedIndex(), + } +} + +func (s *Store) Insert(id, text string, vec Vector, meta map[string]string) { + s.mu.Lock() + defer s.mu.Unlock() + + s.docs = append(s.docs, DocVector{ + ID: id, Vector: vec, Text: text, Meta: meta, + }) + s.index.Add(id, vec) +} + +func (s *Store) Remove(id string) { + s.mu.Lock() + defer s.mu.Unlock() + + filtered := make([]DocVector, 0, len(s.docs)) + for _, d := range s.docs { + if d.ID != id { + filtered = append(filtered, d) + } + } + s.docs = filtered + s.index.Remove(id) +} + +func (s *Store) Search(query Vector, topK int) []DocVector { + s.mu.RLock() + defer s.mu.RUnlock() + + if len(s.docs) == 0 || len(query) == 0 { + return nil + } + + candidates := s.index.Search(query, len(s.docs)) + type scored struct { + doc DocVector + score float64 + } + + var results []scored + seen := make(map[string]bool) + for _, id := range candidates { + if seen[id] { + continue + } + seen[id] = true + for _, d := range s.docs { + if d.ID == id { + score := CosineSimilarity(query, d.Vector) + if score > 0 { + results = append(results, scored{d, score}) + } + break + } + } + } + + sort.Slice(results, func(i, j int) bool { + return results[i].score > results[j].score + }) + + if len(results) > topK { + results = results[:topK] + } + + out := make([]DocVector, len(results)) + for i, r := range results { + out[i] = r.doc + } + return out +} + +func (s *Store) Size() int { + s.mu.RLock() + defer s.mu.RUnlock() + return len(s.docs) +} + +func (s *Store) All() []DocVector { + s.mu.RLock() + defer s.mu.RUnlock() + out := make([]DocVector, len(s.docs)) + copy(out, s.docs) + return out +} + +// TFIDFVectorizer 使用字符 bigram + TF-IDF +type TFIDFVectorizer struct { + mu sync.RWMutex + docFreq map[string]float64 // feature → 文档频率 + totalDocs int + maxNGram int +} + +func NewTFIDFVectorizer(maxNGram int) *TFIDFVectorizer { + if maxNGram <= 0 { + maxNGram = 2 + } + return &TFIDFVectorizer{ + docFreq: make(map[string]float64), + maxNGram: maxNGram, + } +} + +func (v *TFIDFVectorizer) Train(docs []string) { + v.mu.Lock() + defer v.mu.Unlock() + + v.docFreq = make(map[string]float64) + v.totalDocs = len(docs) + + seen := make(map[string]map[string]bool) + for _, doc := range docs { + features := extractNGrams(doc, v.maxNGram) + key := doc + if seen[key] == nil { + seen[key] = make(map[string]bool) + } + for _, f := range features { + if !seen[key][f] { + seen[key][f] = true + v.docFreq[f]++ + } + } + } +} + +func (v *TFIDFVectorizer) Vectorize(text string) Vector { + v.mu.RLock() + defer v.mu.RUnlock() + + features := extractNGrams(text, v.maxNGram) + tf := make(map[string]float64) + for _, f := range features { + tf[f]++ + } + maxTF := 0.0 + for _, c := range tf { + if c > maxTF { + maxTF = c + } + } + + vec := make(Vector) + for f, count := range tf { + tfNorm := count / maxTF + idf := 1.0 + if v.totalDocs > 0 { + df := v.docFreq[f] + if df > 0 { + idf = math.Log(float64(v.totalDocs+1)/df+1) + 1 + } + } + vec[f] = tfNorm * idf + } + return vec +} + +// extractNGrams 提取 n-gram 特征(主要用于中文) +func extractNGrams(text string, maxN int) []string { + runes := []rune(strings.ToLower(text)) + var features []string + seen := make(map[string]bool) + + for n := 1; n <= maxN; n++ { + for i := 0; i <= len(runes)-n; i++ { + gram := string(runes[i : i+n]) + gram = strings.TrimSpace(gram) + if gram == "" { + continue + } + if !seen[gram] { + seen[gram] = true + features = append(features, gram) + } + } + } + return features +} + +func CosineSimilarity(a, b Vector) float64 { + var dot, normA, normB float64 + for f, va := range a { + dot += va * b[f] + normA += va * va + } + for _, vb := range b { + normB += vb * vb + } + if normA == 0 || normB == 0 { + return 0 + } + return dot / (math.Sqrt(normA) * math.Sqrt(normB)) +} + +// InvertedIndex 倒排索引,加速向量搜索 +type InvertedIndex struct { + mu sync.RWMutex + postings map[string]map[string]float64 // feature → {docID: weight} +} + +func NewInvertedIndex() *InvertedIndex { + return &InvertedIndex{ + postings: make(map[string]map[string]float64), + } +} + +func (idx *InvertedIndex) Add(docID string, vec Vector) { + idx.mu.Lock() + defer idx.mu.Unlock() + + for feature, weight := range vec { + if idx.postings[feature] == nil { + idx.postings[feature] = make(map[string]float64) + } + idx.postings[feature][docID] = weight + } +} + +func (idx *InvertedIndex) Remove(docID string) { + idx.mu.Lock() + defer idx.mu.Unlock() + + for feature, postings := range idx.postings { + delete(postings, docID) + if len(postings) == 0 { + delete(idx.postings, feature) + } + } +} + +func (idx *InvertedIndex) Search(query Vector, maxResults int) []string { + idx.mu.RLock() + defer idx.mu.RUnlock() + + scores := make(map[string]float64) + for feature, qw := range query { + if postings, ok := idx.postings[feature]; ok { + for docID, dw := range postings { + scores[docID] += qw * dw + } + } + } + + type pair struct { + id string + score float64 + } + var sorted []pair + for id, score := range scores { + sorted = append(sorted, pair{id, score}) + } + sort.Slice(sorted, func(i, j int) bool { + return sorted[i].score > sorted[j].score + }) + + if len(sorted) > maxResults { + sorted = sorted[:maxResults] + } + out := make([]string, len(sorted)) + for i, p := range sorted { + out[i] = p.id + } + return out +} diff --git a/internal/network/monitor.go b/internal/network/monitor.go new file mode 100644 index 0000000..fa6d8b1 --- /dev/null +++ b/internal/network/monitor.go @@ -0,0 +1,174 @@ +package network + +import ( + "context" + "fmt" + "net" + "net/http" + "sync" + "time" + + "gitcode.com/JianFeeeee/HomeAgent/pkg/types" +) + +type Monitor struct { + mu sync.RWMutex + client *http.Client + interval time.Duration + endpoints []string + status []EndpointStatus +} + +type EndpointStatus struct { + URL string + Reachable bool + Latency time.Duration + LastCheck time.Time + Error string +} + +func NewMonitor(interval time.Duration) *Monitor { + return &Monitor{ + client: &http.Client{ + Timeout: 10 * time.Second, + Transport: &http.Transport{ + DialContext: (&net.Dialer{ + Timeout: 5 * time.Second, + KeepAlive: 30 * time.Second, + }).DialContext, + TLSHandshakeTimeout: 5 * time.Second, + ResponseHeaderTimeout: 5 * time.Second, + DisableKeepAlives: false, + MaxIdleConns: 2, + IdleConnTimeout: 30 * time.Second, + }, + }, + interval: interval, + } +} + +func (m *Monitor) Start(ctx context.Context, endpoints []string) { + m.mu.Lock() + m.endpoints = endpoints + m.status = make([]EndpointStatus, len(endpoints)) + for i, ep := range endpoints { + m.status[i] = EndpointStatus{URL: ep, Reachable: false} + } + m.mu.Unlock() + + ticker := time.NewTicker(m.interval) + defer ticker.Stop() + + m.checkAll(ctx) + + for { + select { + case <-ticker.C: + m.checkAll(ctx) + case <-ctx.Done(): + return + } + } +} + +func (m *Monitor) CheckOnce(ctx context.Context, endpoint string) EndpointStatus { + start := time.Now() + + req, err := http.NewRequestWithContext(ctx, "HEAD", endpoint, nil) + if err != nil { + return EndpointStatus{URL: endpoint, Reachable: false, Error: err.Error(), LastCheck: time.Now()} + } + + resp, err := m.client.Do(req) + latency := time.Since(start) + if err != nil { + return EndpointStatus{URL: endpoint, Reachable: false, Latency: latency, Error: err.Error(), LastCheck: time.Now()} + } + resp.Body.Close() + + return EndpointStatus{ + URL: endpoint, + Reachable: resp.StatusCode < 500, + Latency: latency, + LastCheck: time.Now(), + } +} + +func (m *Monitor) checkAll(ctx context.Context) { + m.mu.RLock() + endpoints := m.endpoints + m.mu.RUnlock() + + var wg sync.WaitGroup + results := make([]EndpointStatus, len(endpoints)) + + for i, ep := range endpoints { + wg.Add(1) + go func(idx int, url string) { + defer wg.Done() + results[idx] = m.CheckOnce(ctx, url) + }(i, ep) + } + + wg.Wait() + + m.mu.Lock() + m.status = results + m.mu.Unlock() +} + +func (m *Monitor) Status() []EndpointStatus { + m.mu.RLock() + defer m.mu.RUnlock() + + status := make([]EndpointStatus, len(m.status)) + copy(status, m.status) + return status +} + +func (m *Monitor) AllReachable() bool { + m.mu.RLock() + defer m.mu.RUnlock() + + for _, s := range m.status { + if !s.Reachable { + return false + } + } + return len(m.status) > 0 +} + +func (m *Monitor) AggregateResult() types.NetworkCheckResult { + m.mu.RLock() + defer m.mu.RUnlock() + + result := types.NetworkCheckResult{LLMAPIReachable: true, DNSResolving: true} + var totalLatency time.Duration + checked := 0 + + for _, s := range m.status { + if !s.Reachable { + result.LLMAPIReachable = false + result.Error = fmt.Sprintf("endpoint %s unreachable: %s", s.URL, s.Error) + } + if s.Latency > 0 { + totalLatency += s.Latency + checked++ + } + } + + if checked > 0 { + result.Latency = totalLatency / time.Duration(checked) + } + + result.DNSResolving = m.checkDNS() + return result +} + +func (m *Monitor) checkDNS() bool { + _, err := net.LookupHost("google.com") + if err != nil { + _, err = net.LookupHost("baidu.com") + } + return err == nil +} diff --git a/internal/onebot/client.go b/internal/onebot/client.go new file mode 100644 index 0000000..d953022 --- /dev/null +++ b/internal/onebot/client.go @@ -0,0 +1,244 @@ +package onebot + +import ( + "encoding/json" + "fmt" + "log" + "sync" + "sync/atomic" + "time" + + "github.com/gorilla/websocket" +) + +// EventHandler 处理 OneBot 推送的事件 +type EventHandler func(event *Event) + +// Client 是 OneBot 反向 WebSocket 客户端 +// 连接到 OneBot 兼容前端(如 go-cqhttp、Lagrange)的 WS 地址 +type Client struct { + url string + accessToken string + conn *websocket.Conn + mu sync.Mutex + done chan struct{} + + eventHandler EventHandler + + // 等待响应的 Action 调用 + pending sync.Map + echoCount int64 + + connected atomic.Bool +} + +// NewClient 创建 OneBot 反向 WS 客户端 +// wsURL: ws://host:port/onebot/v11/ws +func NewClient(wsURL, accessToken string) *Client { + return &Client{ + url: wsURL, + accessToken: accessToken, + done: make(chan struct{}), + } +} + +// SetEventHandler 注册事件处理函数 +func (c *Client) SetEventHandler(h EventHandler) { + c.eventHandler = h +} + +// Connected 返回是否已连接 +func (c *Client) Connected() bool { + return c.connected.Load() +} + +// Connect 连接到 OneBot 前端(阻塞直到连接建立或失败) +func (c *Client) Connect() error { + header := make(map[string][]string) + if c.accessToken != "" { + header["Authorization"] = []string{"Bearer " + c.accessToken} + } + + conn, _, err := websocket.DefaultDialer.Dial(c.url, header) + if err != nil { + return err + } + + c.mu.Lock() + if c.conn != nil { + c.conn.Close() + } + c.conn = conn + c.mu.Unlock() + + c.connected.Store(true) + log.Printf("[onebot] connected to %s", c.url) + + go c.readLoop() + return nil +} + +// reconnect 自动重连 +func (c *Client) reconnect() { + c.connected.Store(false) + backoff := time.Second + + for { + select { + case <-c.done: + return + case <-time.After(backoff): + log.Printf("[onebot] reconnecting in %v...", backoff) + if err := c.Connect(); err != nil { + log.Printf("[onebot] reconnect failed: %v, retry", err) + backoff *= 2 + if backoff > 30*time.Second { + backoff = 30 * time.Second + } + continue + } + log.Printf("[onebot] reconnected") + return + } + } +} + +func (c *Client) readLoop() { + defer c.connected.Store(false) + defer c.mu.Lock() + defer c.mu.Unlock() + defer func() { + if c.conn != nil { + c.conn.Close() + } + }() + + for { + _, message, err := c.conn.ReadMessage() + if err != nil { + log.Printf("[onebot] read error: %v", err) + go c.reconnect() + return + } + + // 尝试解析为 ActionResponse(有 echo 字段) + var resp ActionResponse + if err := json.Unmarshal(message, &resp); err == nil && resp.Echo != "" { + if ch, ok := c.pending.Load(resp.Echo); ok { + ch.(chan *ActionResponse) <- &resp + } + continue + } + + // 解析为 Event + var evt Event + if err := json.Unmarshal(message, &evt); err != nil { + log.Printf("[onebot] parse error: %v", err) + continue + } + + if c.eventHandler != nil { + c.eventHandler(&evt) + } + } +} + +// SendAction 发送一个 OneBot API 请求并等待响应 +func (c *Client) SendAction(action string, params map[string]interface{}, timeout time.Duration) (*ActionResponse, error) { + echo := atomic.AddInt64(&c.echoCount, 1) + echoStr := formatInt64(echo) + + msg := Action{ + Action: action, + Params: params, + Echo: echoStr, + } + + data, err := json.Marshal(msg) + if err != nil { + return nil, err + } + + ch := make(chan *ActionResponse, 1) + c.pending.Store(echoStr, ch) + defer c.pending.Delete(echoStr) + + c.mu.Lock() + if c.conn == nil { + c.mu.Unlock() + return nil, fmt.Errorf("not connected") + } + err = c.conn.WriteMessage(websocket.TextMessage, data) + c.mu.Unlock() + + if err != nil { + return nil, err + } + + if timeout <= 0 { + timeout = 10 * time.Second + } + + select { + case resp := <-ch: + if resp.Status == "failed" { + return resp, fmt.Errorf("onebot action %s failed: retcode=%d", action, resp.RetCode) + } + return resp, nil + case <-time.After(timeout): + return nil, fmt.Errorf("onebot action %s timeout", action) + } +} + +// SendPrivateMessage 发送私聊消息(便捷方法) +func (c *Client) SendPrivateMessage(userID int64, message interface{}, autoEscape bool) (*ActionResponse, error) { + params := map[string]interface{}{ + "user_id": userID, + "message": message, + "auto_escape": autoEscape, + } + return c.SendAction("send_private_msg", params, 0) +} + +// SendGroupMessage 发送群消息(便捷方法) +func (c *Client) SendGroupMessage(groupID int64, message interface{}, autoEscape bool) (*ActionResponse, error) { + params := map[string]interface{}{ + "group_id": groupID, + "message": message, + "auto_escape": autoEscape, + } + return c.SendAction("send_group_msg", params, 0) +} + +// GetLoginInfo 获取登录号信息 +func (c *Client) GetLoginInfo() (*ActionResponse, error) { + return c.SendAction("get_login_info", nil, 0) +} + +// GetGroupMemberInfo 获取群成员信息 +func (c *Client) GetGroupMemberInfo(groupID, userID int64) (*ActionResponse, error) { + return c.SendAction("get_group_member_info", map[string]interface{}{ + "group_id": groupID, + "user_id": userID, + }, 0) +} + +// GetGroupList 获取群列表 +func (c *Client) GetGroupList() (*ActionResponse, error) { + return c.SendAction("get_group_list", nil, 0) +} + +// Close 关闭连接 +func (c *Client) Close() error { + close(c.done) + c.mu.Lock() + defer c.mu.Unlock() + if c.conn != nil { + return c.conn.Close() + } + return nil +} + +func formatInt64(n int64) string { + return fmt.Sprintf("%d", n) +} diff --git a/internal/onebot/device.go b/internal/onebot/device.go new file mode 100644 index 0000000..9f23d72 --- /dev/null +++ b/internal/onebot/device.go @@ -0,0 +1,222 @@ +package onebot + +import ( + "encoding/json" + "fmt" + "log" + "strings" + + agentIO "gitcode.com/JianFeeeee/HomeAgent/internal/agent/io" +) + +// Device 将 OneBot 客户端包装为 IO 抽象层的 Device +// 作为 QQ 通道与 HomeAgent 之间的桥梁 +type Device struct { + name string + desc string + client *Client + iom *agentIO.IOManager +} + +// NewDevice 创建 OneBot IO 设备 +// name: 设备名称(如 "qq") +// wsURL: OneBot 前端 WebSocket 地址(如 "ws://127.0.0.1:6700") +// accessToken: OneBot 鉴权令牌(可选) +func NewDevice(name, wsURL, accessToken string, iom *agentIO.IOManager) *Device { + return &Device{ + name: name, + desc: fmt.Sprintf("OneBot 标准 QQ 通道 (%s)", wsURL), + client: NewClient(wsURL, accessToken), + iom: iom, + } +} + +// Name 返回设备名称 +func (d *Device) Name() string { return d.name } + +// Type 返回设备类型(双向 IO) +func (d *Device) Type() agentIO.DeviceType { return agentIO.DeviceIO } + +// Description 返回设备描述 +func (d *Device) Description() string { return d.desc } + +// OutputCapabilities 返回输出能力(文本+文件+图片) +func (d *Device) OutputCapabilities() agentIO.OutputCapability { + return agentIO.CapText | agentIO.CapFile | agentIO.CapImage +} + +// Tools 返回 OneBot 标准 API 的工具定义 +// AI 可以通过这些工具调用 OneBot 功能 +func (d *Device) Tools() []agentIO.ToolDef { + return []agentIO.ToolDef{ + { + Name: d.name + "_send_private_msg", + Description: "发送 QQ 私聊消息", + Parameters: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "user_id": map[string]interface{}{"type": "integer", "description": "目标 QQ 号"}, + "message": map[string]interface{}{"type": "string", "description": "消息内容(支持 CQ 码,如 [CQ:image,file=xxx.jpg])"}, + "auto_escape": map[string]interface{}{"type": "boolean", "description": "是否作为纯文本发送(不解析 CQ 码)"}, + }, + "required": []string{"user_id", "message"}, + }, + }, + { + Name: d.name + "_send_group_msg", + Description: "发送 QQ 群消息", + Parameters: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "group_id": map[string]interface{}{"type": "integer", "description": "目标群号"}, + "message": map[string]interface{}{"type": "string", "description": "消息内容(支持 CQ 码)"}, + "auto_escape": map[string]interface{}{"type": "boolean", "description": "是否作为纯文本发送"}, + }, + "required": []string{"group_id", "message"}, + }, + }, + { + Name: d.name + "_get_group_member_info", + Description: "获取 QQ 群成员信息", + Parameters: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "group_id": map[string]interface{}{"type": "integer", "description": "群号"}, + "user_id": map[string]interface{}{"type": "integer", "description": "QQ 号"}, + }, + "required": []string{"group_id", "user_id"}, + }, + }, + { + Name: d.name + "_get_group_list", + Description: "获取 QQ 群列表", + Parameters: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{}, + }, + }, + } +} + +// Execute 执行 OneBot 工具调用 +func (d *Device) Execute(tool string, args map[string]interface{}) (interface{}, error) { + if !d.client.Connected() { + return nil, fmt.Errorf("onebot 未连接到前端") + } + + // 去掉名称前缀以匹配方法名 + method := strings.TrimPrefix(tool, d.name+"_") + + switch method { + case "send_private_msg": + userID, _ := toInt64(args["user_id"]) + message, _ := args["message"].(string) + autoEscape, _ := args["auto_escape"].(bool) + return d.client.SendPrivateMessage(userID, message, autoEscape) + + case "send_group_msg": + groupID, _ := toInt64(args["group_id"]) + message, _ := args["message"].(string) + autoEscape, _ := args["auto_escape"].(bool) + return d.client.SendGroupMessage(groupID, message, autoEscape) + + case "get_group_member_info": + groupID, _ := toInt64(args["group_id"]) + userID, _ := toInt64(args["user_id"]) + return d.client.GetGroupMemberInfo(groupID, userID) + + case "get_group_list": + return d.client.GetGroupList() + + default: + return nil, fmt.Errorf("unknown onebot tool: %s", tool) + } +} + +// Start 连接到 OneBot 前端 +func (d *Device) Start() error { + // 非阻塞连接 + go func() { + if err := d.client.Connect(); err != nil { + log.Printf("[onebot] %s initial connect failed, will retry: %v", d.name, err) + } + }() + + // 注册事件处理:OneBot 事件 → IO InputEvent + d.client.SetEventHandler(func(evt *Event) { + d.handleEvent(evt) + }) + + return nil +} + +// Stop 断开连接 +func (d *Device) Stop() error { + return d.client.Close() +} + +// handleEvent 将 OneBot 事件转换为 IO InputEvent +func (d *Device) handleEvent(evt *Event) { + if d.iom == nil { + return + } + + switch evt.PostType { + case "message": + var text string + if evt.RawMessage != "" { + text = evt.RawMessage + } else if s, ok := evt.Message.(string); ok { + text = s + } + + if text == "" { + return + } + + // 构造输入源标识 + source := d.name + payload := map[string]interface{}{ + "content": text, + "source": source, + "user_id": evt.UserID, + "sender": evt.Sender, + } + if evt.MessageType == "group" { + payload["group_id"] = evt.GroupID + payload["label"] = fmt.Sprintf("group:%d:%d", evt.GroupID, evt.UserID) + } else { + payload["label"] = fmt.Sprintf("private:%d", evt.UserID) + } + + d.iom.InjectText(source, text) + // 同时注册输出路由:QQ 消息默认回复到 QQ 通道 + d.iom.RegisterOutputRoute(source, d.name) + + case "notice": + log.Printf("[onebot] notice from %s: type=%s", d.name, evt.NoticeType) + + case "request": + log.Printf("[onebot] request from %s: type=%s flag=%s", d.name, evt.RequestType, evt.Flag) + + case "meta_event": + if evt.MetaEventType == "heartbeat" { + log.Printf("[onebot] %s heartbeat: online=%v", d.name, evt.Status != nil && evt.Status.Online) + } + } +} + +func toInt64(v interface{}) (int64, bool) { + switch n := v.(type) { + case int64: + return n, true + case float64: + return int64(n), true + case int: + return int64(n), true + case json.Number: + i, err := n.Int64() + return i, err == nil + } + return 0, false +} diff --git a/internal/onebot/types.go b/internal/onebot/types.go new file mode 100644 index 0000000..0bc0e95 --- /dev/null +++ b/internal/onebot/types.go @@ -0,0 +1,94 @@ +// Package onebot 实现 OneBot V11 标准协议(Reverse WebSocket 通信) +// 参考: https://github.com/botuniverse/onebot-11 +package onebot + +import "fmt" + +// Action 是 OneBot 标准 API 请求 +type Action struct { + Action string `json:"action"` + Params map[string]interface{} `json:"params,omitempty"` + Echo string `json:"echo,omitempty"` +} + +// ActionResponse 是 OneBot 标准 API 响应 +type ActionResponse struct { + Status string `json:"status"` + RetCode int `json:"retcode"` + Data interface{} `json:"data"` + Echo string `json:"echo,omitempty"` +} + +// Event 是 OneBot 推送的事件 +type Event struct { + Time int64 `json:"time"` + SelfID int64 `json:"self_id"` + PostType string `json:"post_type"` // message, notice, request, meta_event + DetailType string `json:"-"` + + // 消息事件字段 + MessageType string `json:"message_type,omitempty"` // private, group + SubType string `json:"sub_type,omitempty"` + MessageID int64 `json:"message_id,omitempty"` + UserID int64 `json:"user_id,omitempty"` + GroupID int64 `json:"group_id,omitempty"` + Message interface{} `json:"message,omitempty"` // string 或 []MessageSegment + RawMessage string `json:"raw_message,omitempty"` + Font int `json:"font,omitempty"` + Sender *Sender `json:"sender,omitempty"` + + // 通知事件字段 + NoticeType string `json:"notice_type,omitempty"` + + // 请求事件字段 + RequestType string `json:"request_type,omitempty"` + Flag string `json:"flag,omitempty"` + Comment string `json:"comment,omitempty"` + + // 元事件字段 + MetaEventType string `json:"meta_event_type,omitempty"` + Interval int64 `json:"interval,omitempty"` + Status *Status `json:"status,omitempty"` +} + +type Sender struct { + UserID int64 `json:"user_id"` + Nickname string `json:"nickname"` + Sex string `json:"sex,omitempty"` + Age int `json:"age,omitempty"` + Card string `json:"card,omitempty"` // 群名片 + Area string `json:"area,omitempty"` + Level string `json:"level,omitempty"` + Role string `json:"role,omitempty"` // owner, admin, member + Title string `json:"title,omitempty"` +} + +type Status struct { + AppInitialized bool `json:"app_initialized"` + AppEnabled bool `json:"app_enabled"` + PluginsGood bool `json:"plugins_good"` + AppGood bool `json:"app_good"` + Online bool `json:"online"` + Good bool `json:"good"` +} + +// MessageSegment 表示 OneBot 消息段(数组格式) +type MessageSegment struct { + Type string `json:"type"` + Data map[string]string `json:"data"` +} + +// MessageText 快速构造纯文本消息段 +func MessageText(text string) MessageSegment { + return MessageSegment{Type: "text", Data: map[string]string{"text": text}} +} + +// MessageImage 构造图片消息段 +func MessageImage(file string) MessageSegment { + return MessageSegment{Type: "image", Data: map[string]string{"file": file}} +} + +// MessageAt 构造 @ 消息段 +func MessageAt(userID int64) MessageSegment { + return MessageSegment{Type: "at", Data: map[string]string{"qq": fmt.Sprintf("%d", userID)}} +} diff --git a/internal/plugin/plugin.go b/internal/plugin/plugin.go new file mode 100644 index 0000000..5cdbcbd --- /dev/null +++ b/internal/plugin/plugin.go @@ -0,0 +1,899 @@ +package plugin + +import ( + "encoding/json" + "fmt" + "log" + "os" + "path/filepath" + "sort" + "strings" + "sync" + "time" + + agentIO "gitcode.com/JianFeeeee/HomeAgent/internal/agent/io" +) + +type PluginType string + +const ( + PluginTypeSKILL PluginType = "skill" + PluginTypeNative PluginType = "native" +) + +// IOConfig 定义插件作为 IO 通道时的配置 +// 每个插件通过此配置声明自己的 I/O 端口 +type IOConfig struct { + Type string `json:"type"` // "input" / "output" / "io" + InputRoute string `json:"input_route"` // 输入源标识,如 "qq", "email" + OutputRoute string `json:"output_route"` // 输出通道标识,默认等于 InputRoute + OutputCaps []string `json:"output_caps"` // 支持的输出能力: "text","file","image","audio","structured" +} + +type Plugin interface { + Name() string + PluginType() PluginType + Description() string + Version() string + Tools() []ToolDef + Enabled() bool + SetEnabled(bool) + IOConfig() *IOConfig + Device() agentIO.Device // 内嵌的 IO 设备,nil 表示纯技能插件 +} + +// ToolDef 复用 IO 抽象层的定义,确保 Plugin 和 Device 使用同一类型 +type ToolDef = agentIO.ToolDef + +type SKILLPlugin struct { + mu sync.RWMutex + name string + description string + version string + author string + enabled bool + rawContent string + sourceDir string + toolDefs []ToolDef + ioConfig *IOConfig +} + +func LoadSKILL(path string) (*SKILLPlugin, error) { + info, err := os.Stat(path) + if err != nil { + return nil, fmt.Errorf("stat %s: %w", path, err) + } + + name := filepath.Base(path) + p := &SKILLPlugin{ + name: name, + sourceDir: path, + enabled: true, + } + + if info.IsDir() { + skillFile := filepath.Join(path, "SKILL.md") + if data, err := os.ReadFile(skillFile); err == nil { + p.rawContent = string(data) + p.description = extractDescription(p.rawContent) + p.version = extractField(p.rawContent, "version") + p.author = extractField(p.rawContent, "author") + p.ioConfig = extractIOConfig(p.rawContent) + } + + metaFile := filepath.Join(path, "skill.json") + if data, err := os.ReadFile(metaFile); err == nil { + var meta struct { + Name string `json:"name"` + Description string `json:"description"` + Version string `json:"version"` + Author string `json:"author"` + IO *IOConfig `json:"io,omitempty"` + } + if err := json.Unmarshal(data, &meta); err == nil { + if meta.Name != "" { + p.name = meta.Name + } + if meta.Description != "" { + p.description = meta.Description + } + if meta.Version != "" { + p.version = meta.Version + } + if meta.Author != "" { + p.author = meta.Author + } + if meta.IO != nil { + p.ioConfig = meta.IO + } + } + } + } else if filepath.Ext(path) == ".md" { + data, err := os.ReadFile(path) + if err != nil { + return nil, err + } + p.rawContent = string(data) + p.name = name[:len(name)-3] + p.description = extractDescription(p.rawContent) + p.ioConfig = extractIOConfig(p.rawContent) + } + + if p.rawContent != "" { + p.toolDefs = extractToolDefs(p.rawContent) + } + + log.Printf("[plugin] loaded SKILL: %s v%s", p.name, p.version) + return p, nil +} + +func (p *SKILLPlugin) Name() string { return p.name } +func (p *SKILLPlugin) PluginType() PluginType { return PluginTypeSKILL } +func (p *SKILLPlugin) Description() string { return p.description } +func (p *SKILLPlugin) Version() string { return p.version } +func (p *SKILLPlugin) Enabled() bool { p.mu.RLock(); defer p.mu.RUnlock(); return p.enabled } +func (p *SKILLPlugin) SetEnabled(v bool) { p.mu.Lock(); defer p.mu.Unlock(); p.enabled = v } +func (p *SKILLPlugin) Tools() []ToolDef { return p.toolDefs } +func (p *SKILLPlugin) IOConfig() *IOConfig { return p.ioConfig } +func (p *SKILLPlugin) Device() agentIO.Device { return nil } // SKILL 插件无原生 IO 设备 +func (p *SKILLPlugin) RawContent() string { return p.rawContent } + +// NativeFactory 是内置插件构造器,用于需要原生 Go 实现的插件(如 OneBot QQ) +// 返回 agentIO.Device 以直接注册到 IO 管理层 +type NativeFactory func(name string, config map[string]interface{}, iom *agentIO.IOManager) (agentIO.Device, error) + +type Registry struct { + mu sync.RWMutex + plugins map[string]Plugin + ioMgr *agentIO.IOManager + factories map[string]NativeFactory // 名称匹配的插件使用原生实现 +} + +func NewRegistry() *Registry { + return &Registry{ + plugins: make(map[string]Plugin), + factories: make(map[string]NativeFactory), + } +} + +// RegisterNative 注册内置原生插件工厂。当从 plugins/ 加载插件时, +// 如果插件名称匹配已注册的工厂,优先使用原生设备注册。 +// 例如: r.RegisterNative("qq", onebot.NewDeviceFactory) +func (r *Registry) RegisterNative(name string, factory NativeFactory) { + r.mu.Lock() + defer r.mu.Unlock() + r.factories[name] = factory +} + +// SetIOManager 绑定 IO 管理器,启用 IO 设备自动注册 +func (r *Registry) SetIOManager(mgr *agentIO.IOManager) { + r.mu.Lock() + defer r.mu.Unlock() + r.ioMgr = mgr +} + +func (r *Registry) Register(p Plugin) { + r.mu.Lock() + defer r.mu.Unlock() + r.plugins[p.Name()] = p + log.Printf("[plugin] registered: %s (%s)", p.Name(), p.PluginType()) + + // 注册内嵌 IO 设备到 IOManager + if r.ioMgr != nil { + dev := p.Device() + if dev != nil { + // 原生设备直接注册 + if err := r.ioMgr.RegisterDevice(dev); err != nil { + log.Printf("[plugin] register native device %s: %v", p.Name(), err) + return + } + } else if p.IOConfig() != nil { + // 有 IO 配置但无原生设备 → 用 PluginDevice 包装 + dev = NewPluginDevice(p) + if err := r.ioMgr.RegisterDevice(dev); err != nil { + log.Printf("[plugin] register plugin device %s: %v", p.Name(), err) + return + } + } else { + return // 纯技能插件,无 IO 通道 + } + + // 自动注册输出路由 + if cfg := p.IOConfig(); cfg != nil { + log.Printf("[plugin] io device %s active (type=%s, caps=%v)", + p.Name(), cfg.Type, cfg.OutputCaps) + if cfg.InputRoute != "" { + outputRoute := cfg.OutputRoute + if outputRoute == "" { + outputRoute = cfg.InputRoute + } + r.ioMgr.RegisterOutputRoute(cfg.InputRoute, outputRoute) + log.Printf("[plugin] route: %s → %s", cfg.InputRoute, outputRoute) + } + } + } +} + +// nativePlugin 包装原生 IO 设备为 Plugin 接口 +// Plugin 是容器,Device 是组件 +type nativePlugin struct { + name string + dev agentIO.Device + cfg *IOConfig +} + +func (p *nativePlugin) Name() string { return p.name } +func (p *nativePlugin) PluginType() PluginType { return PluginTypeNative } +func (p *nativePlugin) Description() string { return p.dev.Description() } +func (p *nativePlugin) Version() string { return "1.0.0" } +func (p *nativePlugin) Tools() []ToolDef { return p.dev.Tools() } +func (p *nativePlugin) Enabled() bool { return true } +func (p *nativePlugin) SetEnabled(v bool) {} +func (p *nativePlugin) IOConfig() *IOConfig { return p.cfg } +func (p *nativePlugin) Device() agentIO.Device { return p.dev } + +// Reload 原子化重载 plugins/ 目录: +// 1. 扫描磁盘加载新插件,构造并启动新 IO 设备 +// 2. 原子替换 IOManager 的设备表与路由表 +// 3. 停止并清理旧设备 +// 全程无通信中断:旧设备持续服务直到路由切换完成 +func (r *Registry) Reload(dir string) (string, error) { + entries, err := os.ReadDir(dir) + if err != nil { + if os.IsNotExist(err) { + return "插件目录不存在", nil + } + return "", fmt.Errorf("扫描插件目录: %w", err) + } + + // 1. 扫描磁盘,并行加载新插件 + type loadedPlugin struct { + name string + p Plugin + err error + } + var loaded []loadedPlugin + + for _, entry := range entries { + name := entry.Name() + if !entry.IsDir() && !strings.HasSuffix(name, ".md") { + continue + } + if entry.IsDir() { + // 略过非插件目录(不含 SKILL.md) + if _, err := os.Stat(filepath.Join(dir, name, "SKILL.md")); os.IsNotExist(err) { + continue + } + } else { + name = name[:len(name)-3] + } + + r.mu.RLock() + _, exists := r.plugins[name] + r.mu.RUnlock() + if exists { + continue // 已存在,跳过(后续可用 diff 检测变更) + } + + r.mu.RLock() + factory, hasFactory := r.factories[name] + r.mu.RUnlock() + + if hasFactory { + dev, err := factory(name, r.configFor(name, dir), r.ioMgr) + if err != nil { + loaded = append(loaded, loadedPlugin{name: name, err: err}) + continue + } + np := &nativePlugin{name: name, dev: dev} + np.cfg = extractIOConfig(r.readFile(filepath.Join(dir, name))) + loaded = append(loaded, loadedPlugin{name: name, p: np}) + } else { + p, err := LoadSKILL(filepath.Join(dir, entry.Name())) + if err != nil { + loaded = append(loaded, loadedPlugin{name: name, err: err}) + continue + } + loaded = append(loaded, loadedPlugin{name: name, p: p}) + } + } + + // 2. 启动新设备的 IO 通道 + newDevices := make(map[string]agentIO.Device) + newRoutes := make(map[string]string) + for _, lp := range loaded { + if lp.err != nil { + log.Printf("[plugin] skip %s: %v", lp.name, lp.err) + continue + } + dev := lp.p.Device() + if dev == nil && lp.p.IOConfig() != nil { + dev = NewPluginDevice(lp.p) + } + if dev != nil { + dev.Start() // 新设备预先启动 + newDevices[lp.name] = dev + if cfg := lp.p.IOConfig(); cfg != nil { + if cfg.InputRoute != "" { + out := cfg.OutputRoute + if out == "" { + out = cfg.InputRoute + } + newRoutes[cfg.InputRoute] = out + } + } + } + } + + // 3. 原子切换 + r.mu.Lock() + var oldPlugins map[string]Plugin + if r.ioMgr != nil { + // 获取旧设备并原子替换 + oldDevices := r.ioMgr.AtomicSwapDevices(newDevices, newRoutes) + // 停止旧设备 + for _, dev := range oldDevices { + go dev.Stop() + } + } else { + for _, dev := range newDevices { + dev.Stop() + } + } + // 替换插件表 + oldPlugins = r.plugins + r.plugins = make(map[string]Plugin) + for _, lp := range loaded { + if lp.p != nil { + r.plugins[lp.name] = lp.p + } + } + pluginCount := len(r.plugins) + ioCount := len(newDevices) + r.mu.Unlock() + + // 4. 清理旧插件资源 + for _, p := range oldPlugins { + _ = p + } + + log.Printf("[plugin] atomic reload: %d plugins, %d io channels", pluginCount, ioCount) + return fmt.Sprintf("插件重载完成: %d 个插件, %d 个 IO 通道", pluginCount, ioCount), nil +} + +func (r *Registry) Unregister(name string) { + r.mu.Lock() + defer r.mu.Unlock() + + if _, ok := r.plugins[name]; ok && r.ioMgr != nil { + r.ioMgr.UnregisterDevice(name) + log.Printf("[plugin] unregistered io device: %s", name) + } + delete(r.plugins, name) +} + +func (r *Registry) Get(name string) Plugin { + r.mu.RLock() + defer r.mu.RUnlock() + return r.plugins[name] +} + +func (r *Registry) List() []Plugin { + r.mu.RLock() + defer r.mu.RUnlock() + list := make([]Plugin, 0, len(r.plugins)) + for _, p := range r.plugins { + list = append(list, p) + } + sort.Slice(list, func(i, j int) bool { + return list[i].Name() < list[j].Name() + }) + return list +} + +func (r *Registry) ListEnabled() []Plugin { + r.mu.RLock() + defer r.mu.RUnlock() + var list []Plugin + for _, p := range r.plugins { + if p.Enabled() { + list = append(list, p) + } + } + return list +} + +func (r *Registry) GetAllToolDefs() []ToolDef { + r.mu.RLock() + defer r.mu.RUnlock() + var defs []ToolDef + for _, p := range r.plugins { + if !p.Enabled() { + continue + } + defs = append(defs, p.Tools()...) + } + return defs +} + +// HotReload 定期扫描插件目录,检测新增/变更/删除的插件并动态注册/注销 +// interval=0 表示只执行一次扫描 +func (r *Registry) HotReload(dir string, interval time.Duration, done <-chan struct{}) { + if interval <= 0 { + r.scanAndSync(dir) + return + } + ticker := time.NewTicker(interval) + defer ticker.Stop() + for { + select { + case <-ticker.C: + r.scanAndSync(dir) + case <-done: + return + } + } +} + +// scanAndSync 扫描插件目录并与当前注册表同步 +func (r *Registry) scanAndSync(dir string) { + entries, err := os.ReadDir(dir) + if err != nil { + if !os.IsNotExist(err) { + log.Printf("[plugin] scan error: %v", err) + } + return + } + + // 收集当前磁盘上的插件名 + diskSet := make(map[string]bool) + for _, entry := range entries { + name := entry.Name() + if entry.IsDir() { + diskSet[name] = true + } else if strings.HasSuffix(name, ".md") { + diskSet[name[:len(name)-3]] = true + } + } + + r.mu.Lock() + // 移除已不存在的插件 + for name := range r.plugins { + if !diskSet[name] { + if r.ioMgr != nil { + r.ioMgr.UnregisterDevice(name) + } + delete(r.plugins, name) + log.Printf("[plugin] hot-unload: %s", name) + } + } + r.mu.Unlock() + + // 加载新增的插件 + for _, entry := range entries { + path := filepath.Join(dir, entry.Name()) + name := entry.Name() + if entry.IsDir() { + name = entry.Name() + } else if strings.HasSuffix(name, ".md") { + name = name[:len(name)-3] + } else { + continue + } + + r.mu.RLock() + exists := r.plugins[name] + r.mu.RUnlock() + if exists != nil { + continue + } + + // 检查是否有原生工厂 + r.mu.RLock() + factory, hasFactory := r.factories[name] + r.mu.RUnlock() + + if hasFactory { + // 使用原生设备,构造 Plugin 容器 + dev, err := factory(name, r.configFor(name, dir), r.ioMgr) + if err != nil { + log.Printf("[plugin] native factory %s: %v", name, err) + continue + } + np := &nativePlugin{name: name, dev: dev} + np.cfg = extractIOConfig(r.readFile(path)) + r.Register(np) + log.Printf("[plugin] hot-load native: %s", name) + } else { + var p Plugin + p, err = LoadSKILL(path) + if err != nil { + log.Printf("[plugin] hot-load skip %s: %v", entry.Name(), err) + continue + } + r.Register(p) + log.Printf("[plugin] hot-load: %s", p.Name()) + } + } +} + +// configFor 读取插件目录的 skill.json 作为原生工厂的配置 +func (r *Registry) configFor(name, dir string) map[string]interface{} { + cfg := map[string]interface{}{} + path := filepath.Join(dir, name, "skill.json") + data, err := os.ReadFile(path) + if err != nil { + return cfg + } + var meta map[string]interface{} + if json.Unmarshal(data, &meta) == nil { + for k, v := range meta { + cfg[k] = v + } + } + return cfg +} + +// readFile 读取插件目录的 SKILL.md +func (r *Registry) readFile(path string) string { + if info, err := os.Stat(path); err == nil && info.IsDir() { + data, _ := os.ReadFile(filepath.Join(path, "SKILL.md")) + return string(data) + } + data, _ := os.ReadFile(path) + return string(data) +} + +func (r *Registry) LoadDir(dir string) error { + entries, err := os.ReadDir(dir) + if err != nil { + if os.IsNotExist(err) { + return nil + } + return err + } + for _, entry := range entries { + path := filepath.Join(dir, entry.Name()) + name := entry.Name() + if entry.IsDir() { + // 原生插件优先 + r.mu.RLock() + factory, hasFactory := r.factories[name] + r.mu.RUnlock() + if hasFactory { + dev, err := factory(name, r.configFor(name, dir), r.ioMgr) + if err != nil { + log.Printf("[plugin] native factory %s: %v", name, err) + continue + } + r.mu.Lock() + r.plugins[name] = &nativePlugin{name: name, dev: dev} + r.mu.Unlock() + if r.ioMgr != nil { + r.ioMgr.RegisterDevice(dev) + log.Printf("[plugin] native: %s", name) + } + continue + } + p, err := LoadSKILL(path) + if err != nil { + log.Printf("[plugin] skip dir %s: %v", entry.Name(), err) + continue + } + r.Register(p) + } else if filepath.Ext(entry.Name()) == ".md" { + p, err := LoadSKILL(path) + if err != nil { + log.Printf("[plugin] skip file %s: %v", entry.Name(), err) + continue + } + r.Register(p) + } + } + return nil +} + +// extractIOConfig parses IO port metadata from SKILL.md content +// Supported fields: +// io_type: input|output|io +// io_input_route: qq +// io_output_route: qq (optional, defaults to input_route) +// io_output_caps: text,file,image +func extractIOConfig(content string) *IOConfig { + ioType := extractField(content, "io_type") + if ioType == "" { + return nil + } + cfg := &IOConfig{ + Type: ioType, + InputRoute: extractField(content, "io_input_route"), + } + if cfg.InputRoute == "" { + cfg.InputRoute = extractField(content, "io_route") + } + cfg.OutputRoute = extractField(content, "io_output_route") + if cfg.OutputRoute == "" { + cfg.OutputRoute = cfg.InputRoute + } + capsStr := extractField(content, "io_output_caps") + if capsStr != "" { + for _, c := range strings.Split(capsStr, ",") { + cfg.OutputCaps = append(cfg.OutputCaps, strings.TrimSpace(c)) + } + } + return cfg +} + +// PluginDevice 将 Plugin 包装为 IO Device,实现热插拔 +type PluginDevice struct { + plugin Plugin + caps agentIO.OutputCapability +} + +func NewPluginDevice(p Plugin) *PluginDevice { + var caps agentIO.OutputCapability + if cfg := p.IOConfig(); cfg != nil { + // 有 IO 配置时使用声明的能力 + for _, c := range cfg.OutputCaps { + switch strings.ToLower(c) { + case "text": + caps |= agentIO.CapText + case "file": + caps |= agentIO.CapFile + case "image": + caps |= agentIO.CapImage + case "audio": + caps |= agentIO.CapAudio + case "structured": + caps |= agentIO.CapStructured + } + } + } + // 无 IO 配置时 caps == 0 → 纯工具插件,不暴露为输出通道 + return &PluginDevice{plugin: p, caps: caps} +} + +func (d *PluginDevice) Name() string { return d.plugin.Name() } +func (d *PluginDevice) Type() agentIO.DeviceType { + if cfg := d.plugin.IOConfig(); cfg != nil { + switch cfg.Type { + case "input": + return agentIO.DeviceInput + case "output": + return agentIO.DeviceOutput + case "io": + return agentIO.DeviceIO + } + } + // 无 IO 配置 → 纯工具插件,归为 DeviceInput(无输出能力) + return agentIO.DeviceInput +} +func (d *PluginDevice) Description() string { return d.plugin.Description() } + +func (d *PluginDevice) Tools() []agentIO.ToolDef { + pts := d.plugin.Tools() + defs := make([]agentIO.ToolDef, 0, len(pts)) + for _, t := range pts { + defs = append(defs, agentIO.ToolDef{ + Name: t.Name, + Description: t.Description, + Parameters: t.Parameters, + }) + } + return defs +} + +func (d *PluginDevice) Execute(tool string, args map[string]interface{}) (interface{}, error) { + for _, t := range d.plugin.Tools() { + if t.Name == tool { + if t.Handler != nil { + return t.Handler(args) + } + return nil, fmt.Errorf("plugin %s: tool %s has no handler", d.plugin.Name(), tool) + } + } + return nil, fmt.Errorf("plugin %s: unknown tool %s", d.plugin.Name(), tool) +} + +func (d *PluginDevice) Start() error { return nil } +func (d *PluginDevice) Stop() error { return nil } +func (d *PluginDevice) OutputCapabilities() agentIO.OutputCapability { return d.caps } + +// extractDescription returns the first non-empty, non-header line +func extractDescription(content string) string { + for _, line := range splitLines(content) { + line = trimSpace(line) + if line != "" && !hasPrefix(line, "#") { + return line + } + } + return "" +} + +// extractField finds `field: value` pattern in content +func extractField(content string, field string) string { + prefix := field + ":" + for _, line := range splitLines(content) { + trimmed := trimSpace(line) + if hasPrefix(toLower(trimmed), prefix) { + return trimSpace(trimPrefix(trimmed, prefix)) + } + } + return "" +} + +// extractToolDefs 从 SKILL.md 中解析工具定义(OpenClaw 格式) +// +// 格式: +// ## tool_name +// 工具描述 +// - param1: 参数描述 +// - param2: 参数描述 +// +// 也支持: +// ### Tool: tool_name +// 格式(三级标题) +func extractToolDefs(content string) []ToolDef { + lines := splitLines(content) + var defs []ToolDef + var currentTool *ToolDef + inCodeBlock := false + + for i := 0; i < len(lines); i++ { + line := lines[i] + trimmed := trimSpace(line) + + // 跳过代码块 + if strings.HasPrefix(trimmed, "```") { + inCodeBlock = !inCodeBlock + continue + } + if inCodeBlock { + continue + } + + // 检测工具定义开始: ## tool_name 或 ### Tool: tool_name + if strings.HasPrefix(trimmed, "## ") && !strings.HasPrefix(trimmed, "### ") { + // 结束上一个工具 + if currentTool != nil && currentTool.Name != "" { + defs = append(defs, *currentTool) + } + currentTool = &ToolDef{} + namePart := strings.TrimPrefix(trimmed, "## ") + // 跳过已知的非工具章节 + if isNonToolSection(namePart) { + currentTool = nil + continue + } + currentTool.Name = namePart + currentTool.Parameters = map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{}, + } + continue + } + + // 也支持 ### Tool: name 格式 + if strings.HasPrefix(trimmed, "### Tool: ") { + if currentTool != nil && currentTool.Name != "" { + defs = append(defs, *currentTool) + } + currentTool = &ToolDef{} + namePart := strings.TrimPrefix(trimmed, "### Tool: ") + currentTool.Name = namePart + currentTool.Parameters = map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{}, + } + continue + } + + if currentTool == nil { + continue + } + + // 如果 name 为空且已进入工具段,跳过 + if currentTool.Name == "" { + continue + } + + // 描述行:第一个非空、非标题、非列表行 + if currentTool.Description == "" && trimmed != "" && + !strings.HasPrefix(trimmed, "- ") && + !strings.HasPrefix(trimmed, "#") { + currentTool.Description = trimmed + continue + } + + // 参数行:- param: description + if strings.HasPrefix(trimmed, "- ") { + paramLine := strings.TrimPrefix(trimmed, "- ") + colonIdx := strings.Index(paramLine, ":") + if colonIdx > 0 { + paramName := strings.TrimSpace(paramLine[:colonIdx]) + paramDesc := strings.TrimSpace(paramLine[colonIdx+1:]) + if paramName != "" { + props := currentTool.Parameters["properties"].(map[string]interface{}) + props[paramName] = map[string]interface{}{ + "type": "string", + "description": paramDesc, + } + } + } + } + } + + // 收尾最后一个工具 + if currentTool != nil && currentTool.Name != "" { + defs = append(defs, *currentTool) + } + + return defs +} + +// isNonToolSection 判断是否为非工具章节(如 Usage、Examples、Installation 等) +func isNonToolSection(name string) bool { + lower := strings.ToLower(name) + skip := []string{ + "tools", "usage", "examples", "installation", "setup", + "configuration", "overview", "description", "notes", + "parameters", "return", "returns", "options", "syntax", + } + for _, s := range skip { + if lower == s || strings.HasPrefix(lower, s+" ") || strings.HasPrefix(lower, s+":") { + return true + } + } + return false +} + +func splitLines(s string) []string { + var lines []string + start := 0 + for i := 0; i <= len(s); i++ { + if i == len(s) || s[i] == '\n' { + if i > start { + lines = append(lines, s[start:i]) + } + start = i + 1 + } + } + return lines +} + +func trimSpace(s string) string { + start, end := 0, len(s) + for start < end && (s[start] == ' ' || s[start] == '\t' || s[start] == '\r') { + start++ + } + for end > start && (s[end-1] == ' ' || s[end-1] == '\t' || s[end-1] == '\r') { + end-- + } + return s[start:end] +} + +func hasPrefix(s, prefix string) bool { + if len(s) < len(prefix) { + return false + } + return s[:len(prefix)] == prefix +} + +func toLower(s string) string { + b := make([]byte, len(s)) + for i := 0; i < len(s); i++ { + if s[i] >= 'A' && s[i] <= 'Z' { + b[i] = s[i] + 32 + } else { + b[i] = s[i] + } + } + return string(b) +} + +func trimPrefix(s, prefix string) string { + if hasPrefix(s, prefix) { + return s[len(prefix):] + } + return s +} diff --git a/internal/skill/manager.go b/internal/skill/manager.go new file mode 100644 index 0000000..87ca0db --- /dev/null +++ b/internal/skill/manager.go @@ -0,0 +1,218 @@ +package skill + +import ( + "encoding/json" + "fmt" + "os" + "path/filepath" + "sort" + "strings" + "sync" +) + +type Skill struct { + Name string `json:"name"` + Description string `json:"description"` + Version string `json:"version"` + Author string `json:"author,omitempty"` + Entry string `json:"entry,omitempty"` + Source string `json:"source,omitempty"` + Enabled bool `json:"enabled"` + RawContent string `json:"-"` +} + +type Manager struct { + mu sync.RWMutex + skillsDir string + skills map[string]*Skill +} + +func NewManager(skillsDir string) *Manager { + return &Manager{ + skillsDir: skillsDir, + skills: make(map[string]*Skill), + } +} + +func (m *Manager) Init() error { + if err := os.MkdirAll(m.skillsDir, 0755); err != nil { + return fmt.Errorf("create skills dir: %w", err) + } + return m.loadAll() +} + +func (m *Manager) loadAll() error { + entries, err := os.ReadDir(m.skillsDir) + if err != nil { + if os.IsNotExist(err) { + return nil + } + return err + } + + for _, entry := range entries { + if !entry.IsDir() { + continue + } + skillDir := filepath.Join(m.skillsDir, entry.Name()) + skill, err := m.loadSkill(skillDir) + if err != nil { + continue + } + m.skills[skill.Name] = skill + } + + return nil +} + +func (m *Manager) loadSkill(dir string) (*Skill, error) { + skill := &Skill{ + Name: filepath.Base(dir), + Enabled: true, + } + + skillFilePath := filepath.Join(dir, "SKILL.md") + if data, err := os.ReadFile(skillFilePath); err == nil { + skill.RawContent = string(data) + skill.Description = extractDescription(skill.RawContent) + skill.Version = extractField(skill.RawContent, "version") + skill.Author = extractField(skill.RawContent, "author") + } + + metaPath := filepath.Join(dir, "skill.json") + if data, err := os.ReadFile(metaPath); err == nil { + var meta struct { + Name string `json:"name"` + Description string `json:"description"` + Version string `json:"version"` + Author string `json:"author"` + Entry string `json:"entry"` + } + if err := json.Unmarshal(data, &meta); err == nil { + if meta.Name != "" { + skill.Name = meta.Name + } + if meta.Description != "" { + skill.Description = meta.Description + } + if meta.Version != "" { + skill.Version = meta.Version + } + if meta.Author != "" { + skill.Author = meta.Author + } + if meta.Entry != "" { + skill.Entry = meta.Entry + } + } + } + + return skill, nil +} + +func (m *Manager) Install(name string, content string) error { + m.mu.Lock() + defer m.mu.Unlock() + + skillDir := filepath.Join(m.skillsDir, name) + if err := os.MkdirAll(skillDir, 0755); err != nil { + return fmt.Errorf("create skill dir: %w", err) + } + + if err := os.WriteFile(filepath.Join(skillDir, "SKILL.md"), []byte(content), 0644); err != nil { + return fmt.Errorf("write SKILL.md: %w", err) + } + + skill, err := m.loadSkill(skillDir) + if err != nil { + return fmt.Errorf("load installed skill: %w", err) + } + + m.skills[name] = skill + return nil +} + +func (m *Manager) Uninstall(name string) error { + m.mu.Lock() + defer m.mu.Unlock() + + if _, ok := m.skills[name]; !ok { + return fmt.Errorf("skill %s not found", name) + } + + skillDir := filepath.Join(m.skillsDir, name) + if err := os.RemoveAll(skillDir); err != nil { + return fmt.Errorf("remove skill dir: %w", err) + } + + delete(m.skills, name) + return nil +} + +func (m *Manager) List() []*Skill { + m.mu.RLock() + defer m.mu.RUnlock() + + skills := make([]*Skill, 0, len(m.skills)) + for _, s := range m.skills { + skills = append(skills, s) + } + + sort.Slice(skills, func(i, j int) bool { + return skills[i].Name < skills[j].Name + }) + + return skills +} + +func (m *Manager) Get(name string) *Skill { + m.mu.RLock() + defer m.mu.RUnlock() + return m.skills[name] +} + +func (m *Manager) Toggle(name string, enabled bool) error { + m.mu.Lock() + defer m.mu.Unlock() + + if _, ok := m.skills[name]; !ok { + return fmt.Errorf("skill %s not found", name) + } + m.skills[name].Enabled = enabled + return nil +} + +func (m *Manager) GetInjectedPrompt() string { + m.mu.RLock() + defer m.mu.RUnlock() + + var parts []string + for _, s := range m.skills { + if s.Enabled && s.RawContent != "" { + parts = append(parts, fmt.Sprintf("=== Skill: %s ===\n%s", s.Name, s.RawContent)) + } + } + return strings.Join(parts, "\n\n") +} + +func extractDescription(content string) string { + lines := strings.Split(content, "\n") + for _, line := range lines { + line = strings.TrimSpace(line) + if line != "" && !strings.HasPrefix(line, "#") { + return line + } + } + return "" +} + +func extractField(content string, field string) string { + prefix := fmt.Sprintf("%s:", field) + for _, line := range strings.Split(content, "\n") { + trimmed := strings.TrimSpace(line) + if strings.HasPrefix(strings.ToLower(trimmed), prefix) { + return strings.TrimSpace(strings.TrimPrefix(trimmed, prefix)) + } + } + return "" +} diff --git a/internal/snapshot/manager.go b/internal/snapshot/manager.go new file mode 100644 index 0000000..5ed098a --- /dev/null +++ b/internal/snapshot/manager.go @@ -0,0 +1,196 @@ +package snapshot + +import ( + "context" + "fmt" + "log" + "os" + "path/filepath" + "sort" + "sync" + "time" + + "gitcode.com/JianFeeeee/HomeAgent/internal/container" + "gitcode.com/JianFeeeee/HomeAgent/pkg/types" +) + +type Manager struct { + mu sync.RWMutex + dataDir string + container *container.Manager + snapshots map[types.AgentID][]types.Snapshot +} + +func NewManager(dataDir string, cm *container.Manager) *Manager { + return &Manager{ + dataDir: filepath.Join(dataDir, "snapshots"), + container: cm, + snapshots: make(map[types.AgentID][]types.Snapshot), + } +} + +func (m *Manager) Create(ctx context.Context, agentID types.AgentID, containerID string, reason string) (*types.Snapshot, error) { + snapDir := filepath.Join(m.dataDir, string(agentID)) + if err := os.MkdirAll(snapDir, 0755); err != nil { + return nil, fmt.Errorf("create snapshot dir: %w", err) + } + + snapID := types.SnapshotID(fmt.Sprintf("snap_%s_%d", agentID, time.Now().UnixNano())) + imageTag := fmt.Sprintf("homeagent/snap-%s:%s", agentID, snapID) + imagePath := filepath.Join(snapDir, string(snapID)+".tar") + + if err := m.container.Commit(ctx, containerID, imageTag); err != nil { + return nil, fmt.Errorf("commit container: %w", err) + } + if err := m.container.SaveImage(ctx, imageTag, imagePath); err != nil { + return nil, fmt.Errorf("save image: %w", err) + } + + info, err := os.Stat(imagePath) + var size int64 + if err == nil { + size = info.Size() + } + + snap := types.Snapshot{ + ID: snapID, + AgentID: agentID, + CreatedAt: time.Now(), + Reason: reason, + Size: size, + DockerImage: imageTag, + Valid: true, + } + + m.mu.Lock() + m.snapshots[agentID] = append(m.snapshots[agentID], snap) + m.mu.Unlock() + + log.Printf("[snapshot] created %s for agent %s (reason: %s, size: %d bytes)", snapID, agentID, reason, size) + + m.enforceLimit(agentID) + + return &snap, nil +} + +func (m *Manager) Restore(ctx context.Context, agentID types.AgentID, containerID string, snapID types.SnapshotID) error { + m.mu.RLock() + snapshots := m.snapshots[agentID] + var target *types.Snapshot + for _, s := range snapshots { + if s.ID == snapID && s.Valid { + target = &s + break + } + } + m.mu.RUnlock() + + if target == nil { + return fmt.Errorf("snapshot %s not found or invalid", snapID) + } + + snapDir := filepath.Join(m.dataDir, string(agentID)) + imagePath := filepath.Join(snapDir, string(snapID)+".tar") + + if _, err := os.Stat(imagePath); os.IsNotExist(err) { + return fmt.Errorf("snapshot file %s not found", imagePath) + } + + if err := m.container.Stop(ctx, containerID); err != nil { + log.Printf("[snapshot] warning: stop container during restore: %v", err) + } + + if err := m.container.Remove(ctx, containerID); err != nil { + return fmt.Errorf("remove container for restore: %w", err) + } + + if err := m.container.LoadImage(ctx, imagePath); err != nil { + return fmt.Errorf("load snapshot image: %w", err) + } + + log.Printf("[snapshot] restored agent %s to snapshot %s", agentID, snapID) + return nil +} + +func (m *Manager) List(agentID types.AgentID) []types.Snapshot { + m.mu.RLock() + defer m.mu.RUnlock() + + snapshots := m.snapshots[agentID] + result := make([]types.Snapshot, len(snapshots)) + copy(result, snapshots) + + sort.Slice(result, func(i, j int) bool { + return result[i].CreatedAt.After(result[j].CreatedAt) + }) + + return result +} + +func (m *Manager) Latest(agentID types.AgentID) *types.Snapshot { + snapshots := m.List(agentID) + if len(snapshots) == 0 { + return nil + } + return &snapshots[0] +} + +func (m *Manager) MarkInvalid(agentID types.AgentID, snapID types.SnapshotID) { + m.mu.Lock() + defer m.mu.Unlock() + + for i, s := range m.snapshots[agentID] { + if s.ID == snapID { + m.snapshots[agentID][i].Valid = false + return + } + } +} + +func (m *Manager) enforceLimit(agentID types.AgentID) { + m.mu.Lock() + defer m.mu.Unlock() + + snapshots := m.snapshots[agentID] + if len(snapshots) <= 20 { + return + } + + sort.Slice(snapshots, func(i, j int) bool { + return snapshots[i].CreatedAt.Before(snapshots[j].CreatedAt) + }) + + toRemove := len(snapshots) - 20 + for i := 0; i < toRemove; i++ { + s := snapshots[i] + snapDir := filepath.Join(m.dataDir, string(agentID)) + imagePath := filepath.Join(snapDir, string(s.ID)+".tar") + os.Remove(imagePath) + } + + m.snapshots[agentID] = snapshots[toRemove:] +} + +func (m *Manager) Cleanup(agentID types.AgentID, keep int) { + m.mu.Lock() + defer m.mu.Unlock() + + snapshots := m.snapshots[agentID] + if len(snapshots) <= keep { + return + } + + sort.Slice(snapshots, func(i, j int) bool { + return snapshots[i].CreatedAt.Before(snapshots[j].CreatedAt) + }) + + toRemove := len(snapshots) - keep + for i := 0; i < toRemove; i++ { + s := snapshots[i] + snapDir := filepath.Join(m.dataDir, string(agentID)) + imagePath := filepath.Join(snapDir, string(s.ID)+".tar") + os.Remove(imagePath) + } + + m.snapshots[agentID] = snapshots[toRemove:] +} diff --git a/internal/supervisor/daemon.go b/internal/supervisor/daemon.go new file mode 100644 index 0000000..19f35e0 --- /dev/null +++ b/internal/supervisor/daemon.go @@ -0,0 +1,234 @@ +package supervisor + +import ( + "context" + "fmt" + "log" + "sync" + "time" + + "gitcode.com/JianFeeeee/HomeAgent/internal/network" + "gitcode.com/JianFeeeee/HomeAgent/internal/tracker" + "gitcode.com/JianFeeeee/HomeAgent/pkg/types" +) + +const directAgentID types.AgentID = "main" + +type Daemon struct { + cfg *types.Config + nm *network.Monitor + trk *tracker.Tracker + agents map[types.AgentID]*agentInstance + mu sync.RWMutex + ctx context.Context + cancel context.CancelFunc +} + +type agentInstance struct { + cfg *types.AgentConfig + state types.AgentState + health types.HealthStatus + lastHB time.Time + failCount int + useTracker bool +} + +func New(cfg *types.Config) *Daemon { + ctx, cancel := context.WithCancel(context.Background()) + nm := network.NewMonitor(cfg.Daemon.CheckInterval) + + return &Daemon{ + cfg: cfg, + nm: nm, + agents: make(map[types.AgentID]*agentInstance), + ctx: ctx, + cancel: cancel, + } +} + +func (d *Daemon) SetTracker(trk *tracker.Tracker) { + d.trk = trk +} + +func (d *Daemon) Start() error { + log.Println("[homed] starting HomeAgent daemon") + + go d.nm.Start(d.ctx, d.cfg.Defaults.LLMEndpoints) + + go d.healthLoop() + + log.Println("[homed] daemon started successfully") + return nil +} + +func (d *Daemon) Shutdown() { + log.Println("[homed] shutting down...") + d.cancel() + + d.mu.RLock() + defer d.mu.RUnlock() + + for id, agent := range d.agents { + if agent.state == types.AgentStateRunning { + log.Printf("[homed] stopping agent %s", id) + agent.state = types.AgentStateStopped + } + } +} + +func (d *Daemon) RegisterAgent(id types.AgentID) { + d.mu.Lock() + defer d.mu.Unlock() + + useTrk := d.trk != nil + maxRetries := d.cfg.Defaults.RollbackPolicy.MaxRetries + if maxRetries <= 0 { + maxRetries = 3 + } + cfg := &types.AgentConfig{ID: id} + cfg.RollbackPolicy.MaxRetries = maxRetries + + d.agents[id] = &agentInstance{ + cfg: cfg, + state: types.AgentStateRunning, + health: types.HealthHealthy, + lastHB: time.Now(), + useTracker: useTrk, + } + log.Printf("[homed] agent %s registered (tracker=%v, maxRetries=%d)", id, useTrk, maxRetries) +} + +func (d *Daemon) healthLoop() { + ticker := time.NewTicker(d.cfg.Daemon.HeartbeatInterval) + defer ticker.Stop() + + for { + select { + case <-ticker.C: + d.checkAllAgents() + case <-d.ctx.Done(): + return + } + } +} + +func (d *Daemon) checkAllAgents() { + d.mu.RLock() + agents := make(map[types.AgentID]*agentInstance) + for id, a := range d.agents { + agents[id] = a + } + d.mu.RUnlock() + + for id, agent := range agents { + d.checkAgent(id, agent) + } +} + +func (d *Daemon) checkAgent(id types.AgentID, agent *agentInstance) { + netStatus := d.nm.AggregateResult() + + if !netStatus.LLMAPIReachable { + agent.health = types.HealthDegraded + agent.failCount++ + log.Printf("[homed] agent %s: LLM API unreachable (fail %d)", id, agent.failCount) + } else { + agent.health = types.HealthHealthy + agent.failCount = 0 + } + + agent.lastHB = time.Now() + + if agent.failCount >= agent.cfg.RollbackPolicy.MaxRetries { + d.handleFailure(id, agent) + } +} + +func (d *Daemon) handleFailure(id types.AgentID, agent *agentInstance) { + log.Printf("[homed] agent %s failed %d times, initiating recovery", id, agent.failCount) + + if agent.useTracker && d.trk != nil { + log.Printf("[homed] rolling back agent %s via change tracker", id) + if err := d.trk.Rollback(); err != nil { + log.Printf("[homed] tracker rollback failed: %v — restarting agent", err) + d.restartAgent(id, agent) + return + } + log.Printf("[homed] agent %s tracker rollback complete", id) + agent.failCount = 0 + return + } + + d.restartAgent(id, agent) +} + +func (d *Daemon) restartAgent(id types.AgentID, agent *agentInstance) { + log.Printf("[homed] resetting agent %s", id) + + agent.state = types.AgentStateStopped + d.RegisterAgent(id) + + agent.failCount = 0 + log.Printf("[homed] agent %s reset", id) +} + +func (d *Daemon) GetAgentStatus(id types.AgentID) (*AgentStatus, error) { + d.mu.RLock() + agent, ok := d.agents[id] + d.mu.RUnlock() + + if !ok { + return nil, fmt.Errorf("agent %s not found", id) + } + + netStatus := d.nm.AggregateResult() + + trackerStats := map[string]interface{}{"active": false} + if d.trk != nil { + trackerStats = d.trk.Stats() + } + + return &AgentStatus{ + ID: id, + State: agent.state, + Health: agent.health, + Uptime: time.Since(agent.lastHB), + Network: netStatus, + TrackerStats: trackerStats, + }, nil +} + +func (d *Daemon) ListAgents() []AgentStatus { + d.mu.RLock() + defer d.mu.RUnlock() + + var statuses []AgentStatus + for id, agent := range d.agents { + statuses = append(statuses, AgentStatus{ + ID: id, + State: agent.state, + Health: agent.health, + }) + } + return statuses +} + +func (d *Daemon) PreActionSnapshot(id types.AgentID) (*types.Snapshot, error) { + return nil, fmt.Errorf("snapshot not supported in direct mode — use tracker instead") +} + +func (d *Daemon) RollbackAgent(id types.AgentID, snapID types.SnapshotID) error { + if d.trk != nil { + return d.trk.Rollback() + } + return fmt.Errorf("no tracker available for rollback") +} + +type AgentStatus struct { + ID types.AgentID `json:"id"` + State types.AgentState `json:"state"` + Health types.HealthStatus `json:"health"` + Uptime time.Duration `json:"uptime,omitempty"` + Network types.NetworkCheckResult `json:"network,omitempty"` + TrackerStats map[string]interface{} `json:"tracker_stats,omitempty"` +} diff --git a/internal/tokenizer/jieba.go b/internal/tokenizer/jieba.go new file mode 100644 index 0000000..244e5cf --- /dev/null +++ b/internal/tokenizer/jieba.go @@ -0,0 +1,81 @@ +package tokenizer + +import ( + "strings" + "sync" + + jieba "github.com/yanyiwu/gojieba" +) + +type Jieba struct { + mu sync.Mutex + handle *jieba.Jieba +} + +var ( + global *Jieba + once sync.Once +) + +func Global() *Jieba { + once.Do(func() { + global = &Jieba{ + handle: jieba.NewJieba(), + } + }) + return global +} + +func (j *Jieba) Close() { + j.mu.Lock() + defer j.mu.Unlock() + if j.handle != nil { + j.handle.Free() + j.handle = nil + } +} + +func (j *Jieba) ExtractKeywords(text string, topK int) []string { + j.mu.Lock() + defer j.mu.Unlock() + + words := j.handle.ExtractWithWeight(text, topK) + result := make([]string, 0, len(words)) + seen := make(map[string]bool) + + for _, w := range words { + if seen[w.Word] { + continue + } + if len([]rune(w.Word)) < 2 { + continue + } + seen[w.Word] = true + result = append(result, w.Word) + } + + return result +} + +func (j *Jieba) Cut(text string) []string { + j.mu.Lock() + defer j.mu.Unlock() + + return j.handle.Cut(text, true) +} + +func (j *Jieba) Tag(text string) map[string]string { + j.mu.Lock() + defer j.mu.Unlock() + + words := j.handle.Tag(text) + result := make(map[string]string, len(words)) + for _, pair := range words { + if idx := strings.Index(pair, "/"); idx > 0 { + result[pair[:idx]] = pair[idx+1:] + } else { + result[pair] = "" + } + } + return result +} diff --git a/internal/tracker/changeset.go b/internal/tracker/changeset.go new file mode 100644 index 0000000..f382e6b --- /dev/null +++ b/internal/tracker/changeset.go @@ -0,0 +1,130 @@ +package tracker + +import ( + "crypto/sha256" + "encoding/hex" + "fmt" + "os" + "path/filepath" + "time" +) + +type ChangeType string + +const ( + ChangeFileCreated ChangeType = "created" + ChangeFileModified ChangeType = "modified" + ChangeFileDeleted ChangeType = "deleted" +) + +type FileChange struct { + Path string `json:"path"` + Type ChangeType `json:"type"` + SizeBefore int64 `json:"size_before,omitempty"` + SizeAfter int64 `json:"size_after,omitempty"` + HashBefore string `json:"hash_before,omitempty"` + HashAfter string `json:"hash_after,omitempty"` + Content []byte `json:"-"` // stored separately, not in JSON +} + +type ChangeSet struct { + ID string `json:"id"` + Action string `json:"action"` + Timestamp time.Time `json:"timestamp"` + Files []FileChange `json:"files"` +} + +func NewChangeSet(action string) *ChangeSet { + return &ChangeSet{ + ID: fmt.Sprintf("cs_%d", time.Now().UnixNano()), + Action: action, + Timestamp: time.Now(), + } +} + +func fileHash(path string) (string, int64, error) { + data, err := os.ReadFile(path) + if err != nil { + return "", 0, err + } + h := sha256.Sum256(data) + return hex.EncodeToString(h[:]), int64(len(data)), nil +} + +func fileInfo(path string) (size int64, modTime time.Time, err error) { + info, err := os.Stat(path) + if err != nil { + return 0, time.Time{}, err + } + return info.Size(), info.ModTime(), nil +} + +type FSState struct { + Files map[string]FileChange `json:"files"` + Root string `json:"root"` +} + +func captureFSState(root string) (*FSState, error) { + state := &FSState{ + Files: make(map[string]FileChange), + Root: root, + } + err := filepath.Walk(root, func(path string, info os.FileInfo, err error) error { + if err != nil || info.IsDir() { + return err + } + rel, _ := filepath.Rel(root, path) + hash, size, err := fileHash(path) + if err != nil { + return nil + } + state.Files[rel] = FileChange{ + Path: rel, + HashAfter: hash, + SizeAfter: size, + } + return nil + }) + return state, err +} + +func diffStates(before, after *FSState) []FileChange { + var changes []FileChange + seen := make(map[string]bool) + + for path, afterFile := range after.Files { + seen[path] = true + if beforeFile, ok := before.Files[path]; ok { + if beforeFile.HashAfter != afterFile.HashAfter { + changes = append(changes, FileChange{ + Path: path, + Type: ChangeFileModified, + HashBefore: beforeFile.HashAfter, + HashAfter: afterFile.HashAfter, + SizeBefore: beforeFile.SizeAfter, + SizeAfter: afterFile.SizeAfter, + }) + } + } else { + changes = append(changes, FileChange{ + Path: path, + Type: ChangeFileCreated, + HashAfter: afterFile.HashAfter, + SizeAfter: afterFile.SizeAfter, + }) + } + } + + for path := range before.Files { + if !seen[path] { + changes = append(changes, FileChange{ + Path: path, + Type: ChangeFileDeleted, + HashBefore: before.Files[path].HashAfter, + SizeBefore: before.Files[path].SizeAfter, + }) + } + } + + return changes +} diff --git a/internal/tracker/tracker.go b/internal/tracker/tracker.go new file mode 100644 index 0000000..5e267e5 --- /dev/null +++ b/internal/tracker/tracker.go @@ -0,0 +1,223 @@ +package tracker + +import ( + "encoding/json" + "fmt" + "log" + "os" + "os/exec" + "path/filepath" + "sync" +) + +type Tracker struct { + mu sync.Mutex + dataDir string + workDir string + lowerDir string + upperDir string + mergeDir string + mounted bool + active bool + before *FSState + changeSets []*ChangeSet +} + +func NewTracker(dataDir, workDir string) *Tracker { + return &Tracker{ + dataDir: dataDir, + workDir: workDir, + lowerDir: filepath.Join(workDir, "lower"), + upperDir: filepath.Join(workDir, "upper"), + mergeDir: filepath.Join(workDir, "merged"), + changeSets: make([]*ChangeSet, 0), + } +} + +func (t *Tracker) Init() error { + for _, d := range []string{t.lowerDir, t.upperDir, t.mergeDir} { + if err := os.MkdirAll(d, 0755); err != nil { + return fmt.Errorf("create overlay dir %s: %w", d, err) + } + } + log.Printf("[tracker] initialized (work=%s)", t.workDir) + return nil +} + +func (t *Tracker) Start() error { + t.mu.Lock() + defer t.mu.Unlock() + + if t.mounted { + return nil + } + + if err := t.mountOverlay(); err != nil { + return fmt.Errorf("mount overlay: %w", err) + } + t.mounted = true + t.active = true + + t.before = t.capture() + + log.Printf("[tracker] overlay mounted at %s", t.mergeDir) + return nil +} + +func (t *Tracker) Stop() error { + t.mu.Lock() + defer t.mu.Unlock() + + if !t.mounted { + return nil + } + + if err := t.umountOverlay(); err != nil { + return fmt.Errorf("umount overlay: %w", err) + } + t.mounted = false + t.active = false + return nil +} + +func (t *Tracker) PreAction(action string) *ChangeSet { + t.mu.Lock() + defer t.mu.Unlock() + + cs := NewChangeSet(action) + t.before = t.capture() + return cs +} + +func (t *Tracker) PostAction(action string) *ChangeSet { + t.mu.Lock() + defer t.mu.Unlock() + + after := t.capture() + changes := diffStates(t.before, after) + + cs := NewChangeSet(action) + cs.Files = changes + if len(changes) > 0 { + t.changeSets = append(t.changeSets, cs) + t.saveChangeSet(cs) + log.Printf("[tracker] action=%s changed=%d files", action, len(changes)) + for _, f := range changes { + log.Printf(" %s: %s", f.Type, f.Path) + } + } + + t.before = t.capture() + return cs +} + +func (t *Tracker) HasChanges() bool { + return len(t.changeSets) > 0 +} + +func (t *Tracker) ChangeSets() []*ChangeSet { + t.mu.Lock() + defer t.mu.Unlock() + result := make([]*ChangeSet, len(t.changeSets)) + copy(result, t.changeSets) + return result +} + +func (t *Tracker) capture() *FSState { + state, err := captureFSState(t.upperDir) + if err != nil { + return &FSState{Files: make(map[string]FileChange), Root: t.upperDir} + } + return state +} + +func (t *Tracker) mountOverlay() error { + workDir := filepath.Join(t.workDir, "work") + os.MkdirAll(workDir, 0755) + + args := []string{ + "-t", "overlay", + "overlay", + "-o", fmt.Sprintf("lowerdir=%s,upperdir=%s,workdir=%s", t.lowerDir, t.upperDir, workDir), + t.mergeDir, + } + + cmd := exec.Command("mount", args...) + if output, err := cmd.CombinedOutput(); err != nil { + return fmt.Errorf("mount overlayfs failed: %s: %w", string(output), err) + } + return nil +} + +func (t *Tracker) umountOverlay() error { + cmd := exec.Command("umount", t.mergeDir) + if output, err := cmd.CombinedOutput(); err != nil { + return fmt.Errorf("umount overlayfs failed: %s: %w", string(output), err) + } + return nil +} + +func (t *Tracker) saveChangeSet(cs *ChangeSet) { + dir := filepath.Join(t.dataDir, "changesets") + os.MkdirAll(dir, 0755) + + path := filepath.Join(dir, cs.ID+".json") + data, err := json.MarshalIndent(cs, "", " ") + if err != nil { + log.Printf("[tracker] save changeset %s: %v", cs.ID, err) + return + } + if err := os.WriteFile(path, data, 0644); err != nil { + log.Printf("[tracker] write changeset %s: %v", cs.ID, err) + } +} + +func (t *Tracker) Rollback() error { + t.mu.Lock() + defer t.mu.Unlock() + + if t.mounted { + if err := t.umountOverlay(); err != nil { + return fmt.Errorf("umount for rollback: %w", err) + } + } + + if err := os.RemoveAll(t.upperDir); err != nil { + return fmt.Errorf("remove upper: %w", err) + } + if err := os.RemoveAll(filepath.Join(t.workDir, "work")); err != nil { + return fmt.Errorf("remove work: %w", err) + } + + if err := os.MkdirAll(t.upperDir, 0755); err != nil { + return fmt.Errorf("recreate upper: %w", err) + } + + t.changeSets = nil + t.before = nil + t.mounted = false + + log.Printf("[tracker] rollback complete") + return nil +} + +func (t *Tracker) MergeDir() string { + return t.mergeDir +} + +func (t *Tracker) Stats() map[string]interface{} { + t.mu.Lock() + defer t.mu.Unlock() + totalChanges := 0 + for _, cs := range t.changeSets { + totalChanges += len(cs.Files) + } + return map[string]interface{}{ + "mounted": t.mounted, + "active": t.active, + "change_sets": len(t.changeSets), + "total_changes": totalChanges, + "merge_dir": t.mergeDir, + "upper_dir": t.upperDir, + } +} diff --git a/pkg/types/types.go b/pkg/types/types.go new file mode 100644 index 0000000..27a4559 --- /dev/null +++ b/pkg/types/types.go @@ -0,0 +1,118 @@ +package types + +import "time" + +type AgentState int + +const ( + AgentStateStopped AgentState = 0 + AgentStateRunning AgentState = 1 + AgentStateDegraded AgentState = 2 + AgentStatePanic AgentState = 3 +) + +type HealthStatus int + +const ( + HealthUnknown HealthStatus = 0 + HealthHealthy HealthStatus = 1 + HealthUnstable HealthStatus = 2 + HealthDown HealthStatus = 3 + HealthDegraded HealthStatus = 4 +) + +type AgentID string + +type SnapshotID string + +type Snapshot struct { + ID SnapshotID `json:"id"` + AgentID AgentID `json:"agent_id"` + CreatedAt time.Time `json:"created_at"` + Reason string `json:"reason"` + Size int64 `json:"size_bytes"` + DockerImage string `json:"docker_image,omitempty"` + Valid bool `json:"valid"` +} + +type Heartbeat struct { + AgentID AgentID `json:"agent_id"` + Timestamp time.Time `json:"timestamp"` + State AgentState `json:"state"` + Health HealthStatus `json:"health"` + Uptime time.Duration `json:"uptime"` + LLMConnected bool `json:"llm_connected"` + Error string `json:"error,omitempty"` +} + +type NetworkCheckResult struct { + LLMAPIReachable bool `json:"llm_api_reachable"` + DNSResolving bool `json:"dns_resolving"` + Latency time.Duration `json:"latency_ms"` + Error string `json:"error,omitempty"` +} + +type SnapshotPolicy struct { + Interval time.Duration `json:"interval"` + MaxSnapshots int `json:"max_snapshots"` + PreAction bool `json:"pre_action"` + PostAction bool `json:"post_action"` +} + +type RollbackPolicy struct { + MaxRetries int `json:"max_retries"` + HealthThreshold HealthStatus `json:"health_threshold"` + CooldownPeriod time.Duration `json:"cooldown_period"` + AutoRollback bool `json:"auto_rollback"` +} + +type AgentConfig struct { + ID AgentID `json:"id"` + Image string `json:"image"` + Name string `json:"name"` + LLMEndpoints []string `json:"llm_endpoints"` + SnapshotPolicy SnapshotPolicy `json:"snapshot_policy"` + RollbackPolicy RollbackPolicy `json:"rollback_policy"` + ResourceLimit ResourceLimit `json:"resource_limit"` + OpenClawEnabled bool `json:"openclaw_enabled"` +} + +type ResourceLimit struct { + CPU string `json:"cpu"` + Memory string `json:"memory"` + Disk string `json:"disk"` + Network bool `json:"network"` +} + +type OperationLog struct { + ID string `json:"id"` + AgentID AgentID `json:"agent_id"` + Timestamp time.Time `json:"timestamp"` + Action string `json:"action"` + SnapshotID SnapshotID `json:"snapshot_id,omitempty"` + Success bool `json:"success"` +} + +type LLMConfig struct { + Provider string `json:"provider"` + Model string `json:"model"` + BaseURL string `json:"base_url"` + APIKey string `json:"api_key"` + Temperature float64 `json:"temperature"` + MaxTokens int `json:"max_tokens"` +} + +type Config struct { + Daemon DaemonConfig `json:"daemon"` + LLM LLMConfig `json:"llm"` + Defaults AgentConfig `json:"defaults"` + Agents []AgentConfig `json:"agents"` +} + +type DaemonConfig struct { + ListenAddr string `json:"listen_addr"` + DataDir string `json:"data_dir"` + HeartbeatInterval time.Duration `json:"heartbeat_interval"` + CheckInterval time.Duration `json:"check_interval"` + LogLevel string `json:"log_level"` +} diff --git a/plugins/qq/SKILL.md b/plugins/qq/SKILL.md new file mode 100644 index 0000000..537649c --- /dev/null +++ b/plugins/qq/SKILL.md @@ -0,0 +1,31 @@ +# QQ 通知插件 +version: 1.0.0 +author: HomeAgent +io_type: io +io_input_route: qq +io_output_route: qq +io_output_caps: text,file,image + +QQ 通道插件,通过 OneBot 标准协议与 QQ 前端(go-cqhttp / Lagrange)协同工作。 +OneBot 前端处理 QQ 协议登录与消息收发,本插件将 OneBot 事件转换为 IO 输入, +将 AI 输出转换为 OneBot 动作调用。 + +## qq_send_private_msg +发送 QQ 私聊消息 +- user_id: 目标 QQ 号 +- message: 消息内容(支持 CQ 码,如 [CQ:image,file=xxx.jpg]) +- auto_escape: 是否作为纯文本发送 + +## qq_send_group_msg +发送 QQ 群消息 +- group_id: 目标群号 +- message: 消息内容(支持 CQ 码) +- auto_escape: 是否作为纯文本发送 + +## qq_get_group_member_info +获取 QQ 群成员信息 +- group_id: 群号 +- user_id: QQ 号 + +## qq_get_group_list +获取 QQ 群列表 diff --git a/plugins/qq/skill.json b/plugins/qq/skill.json new file mode 100644 index 0000000..5e768cc --- /dev/null +++ b/plugins/qq/skill.json @@ -0,0 +1,13 @@ +{ + "name": "qq", + "description": "QQ 通道插件,通过 OneBot 标准协议与 QQ 前端协同工作", + "version": "1.0.0", + "author": "HomeAgent", + "entry": "ws://127.0.0.1:6700", + "io": { + "type": "io", + "input_route": "qq", + "output_route": "qq", + "output_caps": ["text", "file", "image"] + } +}