mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-09-22 18:08:04 +00:00
适配多个 LLM 源 (anthropic/gemini/mistral/groq/github) + SQLite 配置收敛 + 测试插件
This commit is contained in:
@ -1,133 +1,433 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"gitcode.com/JianFeeeee/HomeAgent/pkg/types"
|
||||
_ "github.com/mattn/go-sqlite3"
|
||||
)
|
||||
|
||||
type ConfigRegistry struct {
|
||||
mu sync.RWMutex
|
||||
values map[string]interface{}
|
||||
persistPath string
|
||||
dirty bool
|
||||
mu sync.RWMutex
|
||||
db *sql.DB
|
||||
dbPath string
|
||||
}
|
||||
|
||||
func NewConfigRegistry(persistPath string) *ConfigRegistry {
|
||||
r := &ConfigRegistry{
|
||||
values: make(map[string]interface{}),
|
||||
persistPath: persistPath,
|
||||
func NewConfigRegistry(dbPath string) *ConfigRegistry {
|
||||
if dbPath == "" {
|
||||
dbPath = ":memory:"
|
||||
}
|
||||
r.load()
|
||||
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()
|
||||
if _, exists := r.values[key]; !exists {
|
||||
r.values[key] = value
|
||||
}
|
||||
r.db.Exec(`INSERT OR IGNORE INTO config (key, value) VALUES (?, ?)`, key, fmt.Sprint(value))
|
||||
}
|
||||
|
||||
func (r *ConfigRegistry) RegisterDefault(key string, value interface{}) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
if _, exists := r.values[key]; !exists {
|
||||
r.values[key] = value
|
||||
}
|
||||
r.Register(key, value)
|
||||
}
|
||||
|
||||
func (r *ConfigRegistry) Get(key string) (interface{}, error) {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
v, ok := r.values[key]
|
||||
if !ok {
|
||||
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)
|
||||
}
|
||||
return v, nil
|
||||
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()
|
||||
r.values[key] = value
|
||||
r.dirty = true
|
||||
return nil
|
||||
_, 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
|
||||
for k := range r.values {
|
||||
if prefix == "" || strings.HasPrefix(k, prefix) {
|
||||
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)
|
||||
}
|
||||
}
|
||||
sort.Strings(keys)
|
||||
return keys
|
||||
}
|
||||
|
||||
func (r *ConfigRegistry) Delete(key string) error {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
delete(r.values, key)
|
||||
r.dirty = true
|
||||
return nil
|
||||
_, 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()
|
||||
cp := make(map[string]interface{})
|
||||
for k, v := range r.values {
|
||||
cp[k] = v
|
||||
result := make(map[string]interface{})
|
||||
rows, err := r.db.Query(`SELECT key, value FROM config ORDER BY key`)
|
||||
if err != nil {
|
||||
return result
|
||||
}
|
||||
return cp
|
||||
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 {
|
||||
r.mu.RLock()
|
||||
if !r.dirty {
|
||||
r.mu.RUnlock()
|
||||
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()
|
||||
}
|
||||
|
||||
// SeedFrom 从 *types.Config 批量导入默认值到 config 表(仅空表时写入)
|
||||
func (r *ConfigRegistry) SeedFrom(cfg *types.Config) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
if r.persistPath == "" {
|
||||
return nil
|
||||
// 检查是否已有数据
|
||||
var count int
|
||||
r.db.QueryRow(`SELECT COUNT(*) FROM config`).Scan(&count)
|
||||
if count > 0 {
|
||||
return
|
||||
}
|
||||
os.MkdirAll(filepath.Dir(r.persistPath), 0755)
|
||||
data, err := json.MarshalIndent(r.values, "", " ")
|
||||
|
||||
tx, err := r.db.Begin()
|
||||
if err != nil {
|
||||
return fmt.Errorf("marshal config: %w", err)
|
||||
return
|
||||
}
|
||||
if err := os.WriteFile(r.persistPath, data, 0644); err != nil {
|
||||
return fmt.Errorf("write config: %w", err)
|
||||
defer tx.Rollback()
|
||||
|
||||
stmt, err := tx.Prepare(`INSERT OR IGNORE INTO config (key, value) VALUES (?, ?)`)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
r.dirty = false
|
||||
return nil
|
||||
defer stmt.Close()
|
||||
|
||||
set := func(k, v string) { stmt.Exec(k, v) }
|
||||
|
||||
// daemon
|
||||
set("core.daemon.listen_addr", cfg.Daemon.ListenAddr)
|
||||
set("core.daemon.data_dir", cfg.Daemon.DataDir)
|
||||
set("core.daemon.heartbeat_interval", cfg.Daemon.HeartbeatInterval.String())
|
||||
set("core.daemon.check_interval", cfg.Daemon.CheckInterval.String())
|
||||
set("core.daemon.log_level", cfg.Daemon.LogLevel)
|
||||
|
||||
// llm
|
||||
set("core.llm.provider", cfg.LLM.Provider)
|
||||
set("core.llm.model", cfg.LLM.Model)
|
||||
set("core.llm.base_url", cfg.LLM.BaseURL)
|
||||
set("core.llm.adapter", cfg.LLM.Adapter)
|
||||
set("core.llm.temperature", strconv.FormatFloat(cfg.LLM.Temperature, 'f', 2, 64))
|
||||
set("core.llm.max_tokens", strconv.Itoa(cfg.LLM.MaxTokens))
|
||||
|
||||
// llm sources
|
||||
for _, src := range cfg.LLM.Sources {
|
||||
p := "core.llm.sources." + src.Name
|
||||
set(p+".base_url", src.BaseURL)
|
||||
set(p+".model", src.Model)
|
||||
set(p+".adapter", src.Adapter)
|
||||
set(p+".adapter_path", src.AdapterPath)
|
||||
}
|
||||
|
||||
// defaults
|
||||
set("core.defaults.image", cfg.Defaults.Image)
|
||||
set("core.defaults.openclaw_enabled", strconv.FormatBool(cfg.Defaults.OpenClawEnabled))
|
||||
set("core.defaults.snapshot.interval", cfg.Defaults.SnapshotPolicy.Interval.String())
|
||||
set("core.defaults.snapshot.max_snapshots", strconv.Itoa(cfg.Defaults.SnapshotPolicy.MaxSnapshots))
|
||||
set("core.defaults.snapshot.pre_action", strconv.FormatBool(cfg.Defaults.SnapshotPolicy.PreAction))
|
||||
set("core.defaults.snapshot.post_action", strconv.FormatBool(cfg.Defaults.SnapshotPolicy.PostAction))
|
||||
set("core.defaults.rollback.max_retries", strconv.Itoa(cfg.Defaults.RollbackPolicy.MaxRetries))
|
||||
set("core.defaults.rollback.health_threshold", strconv.Itoa(int(cfg.Defaults.RollbackPolicy.HealthThreshold)))
|
||||
set("core.defaults.rollback.cooldown_period", cfg.Defaults.RollbackPolicy.CooldownPeriod.String())
|
||||
set("core.defaults.rollback.auto_rollback", strconv.FormatBool(cfg.Defaults.RollbackPolicy.AutoRollback))
|
||||
set("core.defaults.resource.cpu", cfg.Defaults.ResourceLimit.CPU)
|
||||
set("core.defaults.resource.memory", cfg.Defaults.ResourceLimit.Memory)
|
||||
set("core.defaults.resource.disk", cfg.Defaults.ResourceLimit.Disk)
|
||||
set("core.defaults.resource.network", strconv.FormatBool(cfg.Defaults.ResourceLimit.Network))
|
||||
set("core.agent.max_tool_turns", "10")
|
||||
set("core.agent.max_context_size", "30")
|
||||
set("core.agent.distill_interval", "30m")
|
||||
|
||||
tx.Commit()
|
||||
}
|
||||
|
||||
func (r *ConfigRegistry) load() {
|
||||
if r.persistPath == "" {
|
||||
return
|
||||
}
|
||||
data, err := os.ReadFile(r.persistPath)
|
||||
// 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
|
||||
return defaultVal
|
||||
}
|
||||
var vals map[string]interface{}
|
||||
if err := json.Unmarshal(data, &vals); err != nil {
|
||||
return
|
||||
}
|
||||
r.values = vals
|
||||
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.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)
|
||||
|
||||
// 重建 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", ""),
|
||||
Adapter: read(p+".adapter", ""),
|
||||
AdapterPath: read(p+".adapter_path", ""),
|
||||
})
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
@ -1,9 +1,11 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"gitcode.com/JianFeeeee/HomeAgent/pkg/types"
|
||||
)
|
||||
|
||||
func TestRegistryBasic(t *testing.T) {
|
||||
@ -33,23 +35,25 @@ func TestRegistryBasic(t *testing.T) {
|
||||
|
||||
func TestRegistryPersist(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "settings.json")
|
||||
path := filepath.Join(dir, "config.db")
|
||||
|
||||
r := NewConfigRegistry(path)
|
||||
r.Register("core.log_level", "debug")
|
||||
r.Set("plugin.test.key", 42)
|
||||
r.Set("plugin.test.key", "42")
|
||||
if err := r.Flush(); err != nil {
|
||||
t.Fatalf("Flush: %v", err)
|
||||
}
|
||||
r.Close()
|
||||
|
||||
r2 := NewConfigRegistry(path)
|
||||
val, err := r2.Get("plugin.test.key")
|
||||
if err != nil {
|
||||
t.Fatalf("Get after reload: %v", err)
|
||||
}
|
||||
if v, _ := val.(float64); v != 42 {
|
||||
if v, _ := val.(string); v != "42" {
|
||||
t.Fatalf("expected 42, got %v", val)
|
||||
}
|
||||
r2.Close()
|
||||
}
|
||||
|
||||
func TestRegistryDelete(t *testing.T) {
|
||||
@ -65,7 +69,7 @@ func TestRegistryDelete(t *testing.T) {
|
||||
|
||||
func TestRegistryDump(t *testing.T) {
|
||||
r := NewConfigRegistry("")
|
||||
r.Register("x", 1)
|
||||
r.Register("x", "1")
|
||||
r.Register("y", "two")
|
||||
dump := r.Dump()
|
||||
if len(dump) != 2 {
|
||||
@ -83,23 +87,171 @@ func TestRegistryUnknownKey(t *testing.T) {
|
||||
|
||||
func TestRegistryFlush(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "settings.json")
|
||||
path := filepath.Join(dir, "config.db")
|
||||
r := NewConfigRegistry(path)
|
||||
r.Set("k", "v")
|
||||
if err := r.Flush(); err != nil {
|
||||
t.Fatalf("Flush: %v", err)
|
||||
}
|
||||
data, _ := os.ReadFile(path)
|
||||
if len(data) == 0 {
|
||||
t.Fatal("expected persisted data")
|
||||
r.Close()
|
||||
|
||||
// Reopen and verify persistence
|
||||
r2 := NewConfigRegistry(path)
|
||||
val, err := r2.Get("k")
|
||||
if err != nil {
|
||||
t.Fatalf("Get after flush: %v", err)
|
||||
}
|
||||
if v, _ := val.(string); v != "v" {
|
||||
t.Fatalf("expected v, got %v", val)
|
||||
}
|
||||
r2.Close()
|
||||
}
|
||||
|
||||
func TestRegistryFlushIdempotent(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "settings.json")
|
||||
path := filepath.Join(dir, "config.db")
|
||||
r := NewConfigRegistry(path)
|
||||
r.Set("k", "v")
|
||||
r.Flush()
|
||||
r.Flush() // second flush should not error
|
||||
r.Close()
|
||||
}
|
||||
|
||||
func TestPluginConfig(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "config.db")
|
||||
r := NewConfigRegistry(path)
|
||||
|
||||
ps := r.PluginConfig("test_deepseek")
|
||||
if err := ps.Set("api_key", "sk-test123"); err != nil {
|
||||
t.Fatalf("PluginSettings.Set: %v", err)
|
||||
}
|
||||
|
||||
val, err := ps.Get("api_key")
|
||||
if err != nil {
|
||||
t.Fatalf("PluginSettings.Get: %v", err)
|
||||
}
|
||||
if v, _ := val.(string); v != "sk-test123" {
|
||||
t.Fatalf("expected sk-test123, got %v", val)
|
||||
}
|
||||
|
||||
keys, err := ps.List("")
|
||||
if err != nil {
|
||||
t.Fatalf("PluginSettings.List: %v", err)
|
||||
}
|
||||
if len(keys) != 1 || keys[0] != "api_key" {
|
||||
t.Fatalf("expected [api_key], got %v", keys)
|
||||
}
|
||||
|
||||
// Core table should not contain plugin data
|
||||
coreKeys := r.List("")
|
||||
for _, k := range coreKeys {
|
||||
if k == "api_key" {
|
||||
t.Fatal("plugin key leaked into core config table")
|
||||
}
|
||||
}
|
||||
|
||||
r.Close()
|
||||
}
|
||||
|
||||
func TestSeedFromToConfig(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "config.db")
|
||||
|
||||
cfg := &types.Config{
|
||||
Daemon: types.DaemonConfig{
|
||||
ListenAddr: ":9090",
|
||||
DataDir: "/tmp/test",
|
||||
HeartbeatInterval: 10 * time.Second,
|
||||
CheckInterval: 20 * time.Second,
|
||||
LogLevel: "debug",
|
||||
},
|
||||
LLM: types.LLMConfig{
|
||||
Provider: "deepseek",
|
||||
Model: "deepseek-v4-flash",
|
||||
BaseURL: "https://api.deepseek.com",
|
||||
Adapter: "deepseek",
|
||||
Temperature: 0.5,
|
||||
MaxTokens: 2048,
|
||||
Sources: []types.LLMSource{
|
||||
{Name: "deepseek", BaseURL: "https://api.deepseek.com", Model: "deepseek-v4-flash", Adapter: "deepseek", AdapterPath: "adapters/deepseek.lua"},
|
||||
{Name: "openai", BaseURL: "https://api.openai.com/v1", Model: "gpt-4o", Adapter: "openai", AdapterPath: "adapters/openai.lua"},
|
||||
},
|
||||
},
|
||||
Defaults: types.AgentConfig{
|
||||
Image: "test-image",
|
||||
OpenClawEnabled: true,
|
||||
},
|
||||
}
|
||||
|
||||
r := NewConfigRegistry(path)
|
||||
r.SeedFrom(cfg)
|
||||
|
||||
// Verify DB was seeded
|
||||
if len(r.List("")) == 0 {
|
||||
t.Fatal("SeedFrom produced empty DB")
|
||||
}
|
||||
|
||||
// Reconstruct config from DB
|
||||
cfg2 := r.ToConfig()
|
||||
|
||||
if cfg2.Daemon.ListenAddr != ":9090" {
|
||||
t.Fatalf("expected :9090, got %s", cfg2.Daemon.ListenAddr)
|
||||
}
|
||||
if cfg2.Daemon.LogLevel != "debug" {
|
||||
t.Fatalf("expected debug, got %s", cfg2.Daemon.LogLevel)
|
||||
}
|
||||
if cfg2.LLM.Provider != "deepseek" {
|
||||
t.Fatalf("expected deepseek, got %s", cfg2.LLM.Provider)
|
||||
}
|
||||
if cfg2.LLM.MaxTokens != 2048 {
|
||||
t.Fatalf("expected 2048, got %d", cfg2.LLM.MaxTokens)
|
||||
}
|
||||
if len(cfg2.LLM.Sources) != 2 {
|
||||
t.Fatalf("expected 2 sources, got %d", len(cfg2.LLM.Sources))
|
||||
}
|
||||
if cfg2.LLM.Sources[0].AdapterPath != "adapters/deepseek.lua" {
|
||||
t.Fatalf("expected adapters/deepseek.lua, got %s", cfg2.LLM.Sources[0].AdapterPath)
|
||||
}
|
||||
|
||||
// Second SeedFrom should be no-op (DB already has data)
|
||||
r.SeedFrom(cfg)
|
||||
if len(r.List("")) != len(r.List("")) {
|
||||
t.Fatal("second SeedFrom changed DB count")
|
||||
}
|
||||
|
||||
r.Close()
|
||||
}
|
||||
|
||||
func TestGetHelpers(t *testing.T) {
|
||||
r := NewConfigRegistry("")
|
||||
r.Set("str_key", "hello")
|
||||
r.Set("int_key", "42")
|
||||
r.Set("dur_key", "5m")
|
||||
r.Set("bool_key", "true")
|
||||
|
||||
if got := r.GetString("str_key", ""); got != "hello" {
|
||||
t.Fatalf("GetString: expected hello, got %s", got)
|
||||
}
|
||||
if got := r.GetString("nonexistent", "fallback"); got != "fallback" {
|
||||
t.Fatalf("GetString fallback: expected fallback, got %s", got)
|
||||
}
|
||||
if got := r.GetInt("int_key", 0); got != 42 {
|
||||
t.Fatalf("GetInt: expected 42, got %d", got)
|
||||
}
|
||||
if got := r.GetInt("nonexistent", 99); got != 99 {
|
||||
t.Fatalf("GetInt fallback: expected 99, got %d", got)
|
||||
}
|
||||
if got := r.GetDuration("dur_key", 0); got != 5*time.Minute {
|
||||
t.Fatalf("GetDuration: expected 5m, got %v", got)
|
||||
}
|
||||
if got := r.GetDuration("nonexistent", 30*time.Second); got != 30*time.Second {
|
||||
t.Fatalf("GetDuration fallback: expected 30s, got %v", got)
|
||||
}
|
||||
if got := r.GetBool("bool_key", false); got != true {
|
||||
t.Fatalf("GetBool: expected true, got %v", got)
|
||||
}
|
||||
if got := r.GetBool("nonexistent", true); got != true {
|
||||
t.Fatalf("GetBool fallback: expected true, got %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user