mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-09-21 17:38:10 +00:00
模型模式: - 新增 LLMConfig/Source.ThinkingEnabled 配置,通过 ExtraBody 控制 DeepSeek thinking mode,默认关闭 - SeedDefaults/ToConfig 读写 core.llm.thinking_enabled - deepseek.lua 移除硬编码 temperature=0 Unicode 截断: - truncateStr 改按 rune 计数,修复中文截断乱码 审计修复 (Critical): - graph.go: defer rows.Close 在 for 循环 → 显式 Close (连接池泄漏) - cli/openclaw/plugin.go: bare type assertion → comma-ok (panic) - channel.go: payload["type"].(string) → comma-ok (panic) - webui/handler.go: .(string) → fmt.Sprint (panic) - agent.go: 添加 nil provider 错误返回 审计修复 (High): - events/bus.go: copy handler slice under RLock (data race) - webui/handler.go: SSE 通过 channel 串行化写入 (data race) - timer/plugin.go: time.Sleep → select with stopCh (Stop 阻塞) - provider.go: stream ch <- 添加 select ctx.Done (goroutine 泄漏) - main.go: outputCh goroutine 添加 ctx.Done 退出路径
451 lines
14 KiB
Go
451 lines
14 KiB
Go
package config
|
||
|
||
import (
|
||
"database/sql"
|
||
"fmt"
|
||
"sort"
|
||
"strconv"
|
||
"strings"
|
||
"sync"
|
||
"time"
|
||
|
||
"gitcode.com/JianFeeeee/HomeAgent/pkg/types"
|
||
_ "github.com/mattn/go-sqlite3"
|
||
)
|
||
|
||
type ConfigRegistry struct {
|
||
mu sync.RWMutex
|
||
db *sql.DB
|
||
dbPath string
|
||
}
|
||
|
||
func NewConfigRegistry(dbPath string) *ConfigRegistry {
|
||
if dbPath == "" {
|
||
dbPath = ":memory:"
|
||
}
|
||
db, err := sql.Open("sqlite3", dbPath)
|
||
if err != nil {
|
||
panic(fmt.Sprintf("open config db: %v", err))
|
||
}
|
||
// WAL 模式提升并发
|
||
db.Exec("PRAGMA journal_mode=WAL")
|
||
r := &ConfigRegistry{db: db, dbPath: dbPath}
|
||
r.initCoreTable()
|
||
return r
|
||
}
|
||
|
||
func (r *ConfigRegistry) initCoreTable() {
|
||
r.db.Exec(`CREATE TABLE IF NOT EXISTS config (
|
||
key TEXT PRIMARY KEY,
|
||
value TEXT NOT NULL
|
||
)`)
|
||
}
|
||
|
||
func (r *ConfigRegistry) ensurePluginTable(name string) {
|
||
table := r.pluginTableName(name)
|
||
r.db.Exec(fmt.Sprintf(`CREATE TABLE IF NOT EXISTS %s (
|
||
key TEXT PRIMARY KEY,
|
||
value TEXT NOT NULL
|
||
)`, table))
|
||
}
|
||
|
||
func (r *ConfigRegistry) pluginTableName(name string) string {
|
||
safe := strings.Map(func(c rune) rune {
|
||
if (c >= 'a' && c <= 'z') || (c >= '0' && c <= '9') || c == '_' {
|
||
return c
|
||
}
|
||
return '_'
|
||
}, strings.ToLower(name))
|
||
return "config_" + safe
|
||
}
|
||
|
||
func (r *ConfigRegistry) Register(key string, value interface{}) {
|
||
r.mu.Lock()
|
||
defer r.mu.Unlock()
|
||
r.db.Exec(`INSERT OR IGNORE INTO config (key, value) VALUES (?, ?)`, key, fmt.Sprint(value))
|
||
}
|
||
|
||
func (r *ConfigRegistry) RegisterDefault(key string, value interface{}) {
|
||
r.Register(key, value)
|
||
}
|
||
|
||
func (r *ConfigRegistry) Get(key string) (interface{}, error) {
|
||
r.mu.RLock()
|
||
defer r.mu.RUnlock()
|
||
var val string
|
||
err := r.db.QueryRow(`SELECT value FROM config WHERE key = ?`, key).Scan(&val)
|
||
if err == sql.ErrNoRows {
|
||
return nil, fmt.Errorf("config key %q not found", key)
|
||
}
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
return val, nil
|
||
}
|
||
|
||
func (r *ConfigRegistry) Set(key string, value interface{}) error {
|
||
r.mu.Lock()
|
||
defer r.mu.Unlock()
|
||
_, err := r.db.Exec(`INSERT OR REPLACE INTO config (key, value) VALUES (?, ?)`, key, fmt.Sprint(value))
|
||
return err
|
||
}
|
||
|
||
func (r *ConfigRegistry) List(prefix string) []string {
|
||
r.mu.RLock()
|
||
defer r.mu.RUnlock()
|
||
var keys []string
|
||
q := `SELECT key FROM config WHERE key LIKE ? ORDER BY key`
|
||
like := prefix + "%"
|
||
rows, err := r.db.Query(q, like)
|
||
if err != nil {
|
||
return nil
|
||
}
|
||
defer rows.Close()
|
||
for rows.Next() {
|
||
var k string
|
||
if err := rows.Scan(&k); err == nil {
|
||
keys = append(keys, k)
|
||
}
|
||
}
|
||
return keys
|
||
}
|
||
|
||
func (r *ConfigRegistry) Delete(key string) error {
|
||
r.mu.Lock()
|
||
defer r.mu.Unlock()
|
||
_, err := r.db.Exec(`DELETE FROM config WHERE key = ?`, key)
|
||
return err
|
||
}
|
||
|
||
func (r *ConfigRegistry) Dump() map[string]interface{} {
|
||
r.mu.RLock()
|
||
defer r.mu.RUnlock()
|
||
result := make(map[string]interface{})
|
||
rows, err := r.db.Query(`SELECT key, value FROM config ORDER BY key`)
|
||
if err != nil {
|
||
return result
|
||
}
|
||
defer rows.Close()
|
||
for rows.Next() {
|
||
var k, v string
|
||
if err := rows.Scan(&k, &v); err == nil {
|
||
result[k] = v
|
||
}
|
||
}
|
||
return result
|
||
}
|
||
|
||
func (r *ConfigRegistry) Flush() error {
|
||
if r.dbPath == "" || r.dbPath == ":memory:" {
|
||
return nil
|
||
}
|
||
// SQLite 自动持久化;显式 checkpoint 确保一致性
|
||
r.mu.RLock()
|
||
_, err := r.db.Exec("PRAGMA wal_checkpoint(TRUNCATE)")
|
||
r.mu.RUnlock()
|
||
return err
|
||
}
|
||
|
||
func (r *ConfigRegistry) Close() error {
|
||
return r.db.Close()
|
||
}
|
||
|
||
// SeedDefaults 用硬编码默认值填充 config 表(仅空表时写入),不再依赖 YAML
|
||
func (r *ConfigRegistry) SeedDefaults(dataDir string) {
|
||
r.mu.Lock()
|
||
defer r.mu.Unlock()
|
||
|
||
var count int
|
||
r.db.QueryRow(`SELECT COUNT(*) FROM config`).Scan(&count)
|
||
if count > 0 {
|
||
return
|
||
}
|
||
|
||
tx, err := r.db.Begin()
|
||
if err != nil {
|
||
return
|
||
}
|
||
defer tx.Rollback()
|
||
|
||
stmt, err := tx.Prepare(`INSERT OR IGNORE INTO config (key, value) VALUES (?, ?)`)
|
||
if err != nil {
|
||
return
|
||
}
|
||
defer stmt.Close()
|
||
|
||
set := func(k, v string) { stmt.Exec(k, v) }
|
||
|
||
// daemon
|
||
set("core.daemon.listen_addr", ":8080")
|
||
set("core.daemon.data_dir", dataDir)
|
||
set("core.daemon.heartbeat_interval", "15s")
|
||
set("core.daemon.check_interval", "30s")
|
||
set("core.daemon.log_level", "info")
|
||
|
||
// llm
|
||
set("core.llm.provider", "deepseek")
|
||
set("core.llm.model", "deepseek-v4-flash")
|
||
set("core.llm.base_url", "https://api.deepseek.com")
|
||
set("core.llm.api_key", "")
|
||
set("core.llm.adapter", "deepseek")
|
||
set("core.llm.temperature", "0.7")
|
||
set("core.llm.max_tokens", "4096")
|
||
set("core.llm.thinking_enabled", "false")
|
||
|
||
// llm sources
|
||
sources := map[string]map[string]string{
|
||
"deepseek": {"base_url": "https://api.deepseek.com", "model": "deepseek-v4-flash", "api_key": "", "thinking_enabled": "false", "adapter": "deepseek", "adapter_path": "adapters/deepseek.lua"},
|
||
"openai": {"base_url": "https://api.openai.com/v1", "model": "gpt-4o", "api_key": "", "thinking_enabled": "false", "adapter": "openai", "adapter_path": "adapters/openai.lua"},
|
||
"anthropic": {"base_url": "https://api.anthropic.com", "model": "claude-sonnet-4-20250514", "api_key": "", "thinking_enabled": "false", "adapter": "anthropic", "adapter_path": "adapters/anthropic.lua"},
|
||
"gemini": {"base_url": "https://generativelanguage.googleapis.com", "model": "gemini-2.0-flash", "api_key": "", "thinking_enabled": "false", "adapter": "gemini", "adapter_path": "adapters/gemini.lua"},
|
||
"mistral": {"base_url": "https://api.mistral.ai", "model": "mistral-large-latest", "api_key": "", "thinking_enabled": "false", "adapter": "mistral", "adapter_path": "adapters/mistral.lua"},
|
||
"groq": {"base_url": "https://api.groq.com", "model": "llama3-70b-8192", "api_key": "", "thinking_enabled": "false", "adapter": "groq", "adapter_path": "adapters/groq.lua"},
|
||
"github": {"base_url": "https://models.inference.ai.azure.com", "model": "gpt-4o", "api_key": "", "thinking_enabled": "false", "adapter": "github", "adapter_path": "adapters/github.lua"},
|
||
"ollama": {"base_url": "http://localhost:11434", "model": "llama3", "api_key": "", "thinking_enabled": "false", "adapter": "ollama", "adapter_path": "adapters/ollama.lua"},
|
||
}
|
||
for name, props := range sources {
|
||
p := "core.llm.sources." + name
|
||
set(p+".base_url", props["base_url"])
|
||
set(p+".model", props["model"])
|
||
set(p+".api_key", props["api_key"])
|
||
set(p+".thinking_enabled", props["thinking_enabled"])
|
||
set(p+".adapter", props["adapter"])
|
||
set(p+".adapter_path", props["adapter_path"])
|
||
}
|
||
|
||
// defaults
|
||
set("core.defaults.image", "homeagent/agent-base:latest")
|
||
set("core.defaults.openclaw_enabled", "true")
|
||
set("core.defaults.snapshot.interval", "10m")
|
||
set("core.defaults.snapshot.max_snapshots", "20")
|
||
set("core.defaults.snapshot.pre_action", "true")
|
||
set("core.defaults.snapshot.post_action", "false")
|
||
set("core.defaults.rollback.max_retries", "3")
|
||
set("core.defaults.rollback.health_threshold", "3")
|
||
set("core.defaults.rollback.cooldown_period", "30s")
|
||
set("core.defaults.rollback.auto_rollback", "true")
|
||
set("core.defaults.resource.cpu", "2")
|
||
set("core.defaults.resource.memory", "2g")
|
||
set("core.defaults.resource.disk", "10g")
|
||
set("core.defaults.resource.network", "true")
|
||
set("core.agent.max_tool_turns", "10")
|
||
set("core.agent.max_context_size", "30")
|
||
set("core.agent.distill_interval", "30m")
|
||
|
||
tx.Commit()
|
||
}
|
||
|
||
// helpers — 所有值存为 TEXT,解析时自动转换
|
||
|
||
func (r *ConfigRegistry) GetString(key, defaultVal string) string {
|
||
r.mu.RLock()
|
||
defer r.mu.RUnlock()
|
||
var v string
|
||
err := r.db.QueryRow(`SELECT value FROM config WHERE key = ?`, key).Scan(&v)
|
||
if err != nil {
|
||
return defaultVal
|
||
}
|
||
return v
|
||
}
|
||
|
||
func (r *ConfigRegistry) GetInt(key string, defaultVal int) int {
|
||
r.mu.RLock()
|
||
defer r.mu.RUnlock()
|
||
var v string
|
||
err := r.db.QueryRow(`SELECT value FROM config WHERE key = ?`, key).Scan(&v)
|
||
if err != nil {
|
||
return defaultVal
|
||
}
|
||
n, err := strconv.Atoi(v)
|
||
if err != nil {
|
||
return defaultVal
|
||
}
|
||
return n
|
||
}
|
||
|
||
func (r *ConfigRegistry) GetDuration(key string, defaultVal time.Duration) time.Duration {
|
||
r.mu.RLock()
|
||
defer r.mu.RUnlock()
|
||
var v string
|
||
err := r.db.QueryRow(`SELECT value FROM config WHERE key = ?`, key).Scan(&v)
|
||
if err != nil {
|
||
return defaultVal
|
||
}
|
||
d, err := time.ParseDuration(v)
|
||
if err != nil {
|
||
return defaultVal
|
||
}
|
||
return d
|
||
}
|
||
|
||
func (r *ConfigRegistry) GetBool(key string, defaultVal bool) bool {
|
||
r.mu.RLock()
|
||
defer r.mu.RUnlock()
|
||
var v string
|
||
err := r.db.QueryRow(`SELECT value FROM config WHERE key = ?`, key).Scan(&v)
|
||
if err != nil {
|
||
return defaultVal
|
||
}
|
||
b, err := strconv.ParseBool(v)
|
||
if err != nil {
|
||
return defaultVal
|
||
}
|
||
return b
|
||
}
|
||
|
||
// ToConfig 从 config 表重建 *types.Config(数据库为真实源,YAML 仅作初始 seed)
|
||
func (r *ConfigRegistry) ToConfig() *types.Config {
|
||
cfg := &types.Config{}
|
||
dump := r.Dump()
|
||
|
||
read := func(key, def string) string {
|
||
if v, ok := dump[key]; ok {
|
||
if s, ok := v.(string); ok && s != "" {
|
||
return s
|
||
}
|
||
}
|
||
return def
|
||
}
|
||
readInt := func(key string, def int) int {
|
||
s := read(key, "")
|
||
if s == "" {
|
||
return def
|
||
}
|
||
n, err := strconv.Atoi(s)
|
||
if err != nil {
|
||
return def
|
||
}
|
||
return n
|
||
}
|
||
readDur := func(key string, def time.Duration) time.Duration {
|
||
s := read(key, "")
|
||
if s == "" {
|
||
return def
|
||
}
|
||
d, err := time.ParseDuration(s)
|
||
if err != nil {
|
||
return def
|
||
}
|
||
return d
|
||
}
|
||
readBool := func(key string, def bool) bool {
|
||
s := read(key, "")
|
||
if s == "" {
|
||
return def
|
||
}
|
||
b, err := strconv.ParseBool(s)
|
||
if err != nil {
|
||
return def
|
||
}
|
||
return b
|
||
}
|
||
|
||
cfg.Daemon.ListenAddr = read("core.daemon.listen_addr", cfg.Daemon.ListenAddr)
|
||
cfg.Daemon.DataDir = read("core.daemon.data_dir", cfg.Daemon.DataDir)
|
||
cfg.Daemon.HeartbeatInterval = readDur("core.daemon.heartbeat_interval", cfg.Daemon.HeartbeatInterval)
|
||
cfg.Daemon.CheckInterval = readDur("core.daemon.check_interval", cfg.Daemon.CheckInterval)
|
||
cfg.Daemon.LogLevel = read("core.daemon.log_level", cfg.Daemon.LogLevel)
|
||
|
||
cfg.LLM.Provider = read("core.llm.provider", cfg.LLM.Provider)
|
||
cfg.LLM.Model = read("core.llm.model", cfg.LLM.Model)
|
||
cfg.LLM.BaseURL = read("core.llm.base_url", cfg.LLM.BaseURL)
|
||
cfg.LLM.APIKey = read("core.llm.api_key", cfg.LLM.APIKey)
|
||
cfg.LLM.Adapter = read("core.llm.adapter", cfg.LLM.Adapter)
|
||
cfg.LLM.Temperature = float64(readInt("core.llm.temperature", int(cfg.LLM.Temperature*100))) / 100
|
||
cfg.LLM.MaxTokens = readInt("core.llm.max_tokens", cfg.LLM.MaxTokens)
|
||
cfg.LLM.ThinkingEnabled = readBool("core.llm.thinking_enabled", cfg.LLM.ThinkingEnabled)
|
||
|
||
// 重建 sources —— 从 DB 中按前缀扫描,按名称排序保证确定性
|
||
sourceNames := make([]string, 0)
|
||
for k := range dump {
|
||
if strings.HasPrefix(k, "core.llm.sources.") && strings.HasSuffix(k, ".base_url") {
|
||
name := strings.TrimPrefix(k, "core.llm.sources.")
|
||
name = strings.TrimSuffix(name, ".base_url")
|
||
sourceNames = append(sourceNames, name)
|
||
}
|
||
}
|
||
sort.Strings(sourceNames)
|
||
for _, name := range sourceNames {
|
||
p := "core.llm.sources." + name
|
||
cfg.LLM.Sources = append(cfg.LLM.Sources, types.LLMSource{
|
||
Name: name,
|
||
BaseURL: read(p+".base_url", ""),
|
||
Model: read(p+".model", ""),
|
||
APIKey: read(p+".api_key", ""),
|
||
Adapter: read(p+".adapter", ""),
|
||
AdapterPath: read(p+".adapter_path", ""),
|
||
ThinkingEnabled: readBool(p+".thinking_enabled", false),
|
||
})
|
||
}
|
||
|
||
cfg.Defaults.Image = read("core.defaults.image", cfg.Defaults.Image)
|
||
cfg.Defaults.OpenClawEnabled = readBool("core.defaults.openclaw_enabled", cfg.Defaults.OpenClawEnabled)
|
||
cfg.Defaults.SnapshotPolicy.Interval = readDur("core.defaults.snapshot.interval", cfg.Defaults.SnapshotPolicy.Interval)
|
||
cfg.Defaults.SnapshotPolicy.MaxSnapshots = readInt("core.defaults.snapshot.max_snapshots", cfg.Defaults.SnapshotPolicy.MaxSnapshots)
|
||
cfg.Defaults.SnapshotPolicy.PreAction = readBool("core.defaults.snapshot.pre_action", cfg.Defaults.SnapshotPolicy.PreAction)
|
||
cfg.Defaults.SnapshotPolicy.PostAction = readBool("core.defaults.snapshot.post_action", cfg.Defaults.SnapshotPolicy.PostAction)
|
||
cfg.Defaults.RollbackPolicy.MaxRetries = readInt("core.defaults.rollback.max_retries", cfg.Defaults.RollbackPolicy.MaxRetries)
|
||
cfg.Defaults.RollbackPolicy.HealthThreshold = types.HealthStatus(readInt("core.defaults.rollback.health_threshold", int(cfg.Defaults.RollbackPolicy.HealthThreshold)))
|
||
cfg.Defaults.RollbackPolicy.CooldownPeriod = readDur("core.defaults.rollback.cooldown_period", cfg.Defaults.RollbackPolicy.CooldownPeriod)
|
||
cfg.Defaults.RollbackPolicy.AutoRollback = readBool("core.defaults.rollback.auto_rollback", cfg.Defaults.RollbackPolicy.AutoRollback)
|
||
cfg.Defaults.ResourceLimit.CPU = read("core.defaults.resource.cpu", cfg.Defaults.ResourceLimit.CPU)
|
||
cfg.Defaults.ResourceLimit.Memory = read("core.defaults.resource.memory", cfg.Defaults.ResourceLimit.Memory)
|
||
cfg.Defaults.ResourceLimit.Disk = read("core.defaults.resource.disk", cfg.Defaults.ResourceLimit.Disk)
|
||
cfg.Defaults.ResourceLimit.Network = readBool("core.defaults.resource.network", cfg.Defaults.ResourceLimit.Network)
|
||
|
||
return cfg
|
||
}
|
||
func (r *ConfigRegistry) PluginConfig(name string) *PluginSettings {
|
||
r.ensurePluginTable(name)
|
||
return &PluginSettings{
|
||
registry: r,
|
||
table: r.pluginTableName(name),
|
||
}
|
||
}
|
||
|
||
// PluginSettings 实现 sdk.SettingsAPI,作用域为单个插件表
|
||
type PluginSettings struct {
|
||
registry *ConfigRegistry
|
||
table string
|
||
}
|
||
|
||
func (p *PluginSettings) Get(key string) (interface{}, error) {
|
||
p.registry.mu.RLock()
|
||
defer p.registry.mu.RUnlock()
|
||
var val string
|
||
err := p.registry.db.QueryRow(fmt.Sprintf(`SELECT value FROM %s WHERE key = ?`, p.table), key).Scan(&val)
|
||
if err == sql.ErrNoRows {
|
||
return nil, fmt.Errorf("config key %q not found", key)
|
||
}
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
return val, nil
|
||
}
|
||
|
||
func (p *PluginSettings) Set(key string, value interface{}) error {
|
||
p.registry.mu.Lock()
|
||
defer p.registry.mu.Unlock()
|
||
_, err := p.registry.db.Exec(fmt.Sprintf(`INSERT OR REPLACE INTO %s (key, value) VALUES (?, ?)`, p.table), key, fmt.Sprint(value))
|
||
return err
|
||
}
|
||
|
||
func (p *PluginSettings) List(prefix string) ([]string, error) {
|
||
p.registry.mu.RLock()
|
||
defer p.registry.mu.RUnlock()
|
||
var keys []string
|
||
q := fmt.Sprintf(`SELECT key FROM %s WHERE key LIKE ? ORDER BY key`, p.table)
|
||
rows, err := p.registry.db.Query(q, prefix+"%")
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
defer rows.Close()
|
||
for rows.Next() {
|
||
var k string
|
||
if err := rows.Scan(&k); err == nil {
|
||
keys = append(keys, k)
|
||
}
|
||
}
|
||
return keys, nil
|
||
}
|