Files
HomeAgent/internal/config/registry.go
root 239a22899b fix: 模型思考模式配置 + Unicode 截断 + 审计修复 (13 files)
模型模式:
- 新增 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 退出路径
2026-07-03 20:46:04 +08:00

451 lines
14 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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
}