重构:拆出 internal/config 与 internal/store 子包

把工程根目录中和"配置加载"、"数据存储"两个领域相关的全部 Go 文件迁到
internal/ 下的子包,按功能分组的第一阶段成果。

internal/config/
- 从 config.go 抽出 Config / MQTTConfig / WebConfig / DatabaseConfig 等
  类型并大写导出,函数改名 Default/Load/Write/Validate/BuildTLS 等。
- 测试随被测代码迁移成 package config 的内部测试。
- 根目录留 config.go 作桥接:用 type alias 让旧的小写名(config /
  mqttConfig / databaseConfig 等)继续可用,避免修改 30+ 处调用方。

internal/store/
- 把 db.go、store_query.go、db_write_queue.go 与 13 个 *_store.go 一并
  迁入;26 个 *Record 类型与 store 同包以避免循环依赖。
- store -> Store;50+ 标识符从小写未导出改为大写导出(包括 record、
  ListOptions、错误变量、bot/llm/runtime 常量、helpers 等)。
- 新增 DB() / Driver() 访问器供 ai 子系统使用,避免直接访问私有字段。
- bot_pki_store.go 独立出来,把 PKI 解密所需的 store 方法集中归类。
- helpers.go 提供 hashPassword / uint32FromRecord / printJSON 等以前在
  其他根目录文件中的辅助;test_helpers_test.go 提供 verifyPassword
  与 publicMapTileSourceDTO 让测试可以本地运行而不依赖 main 包。

根目录新增:
- store_bridge.go:完整 type-alias / 函数包装层,把 internal/store 的
  导出名映射回旧的小写名,让 admin_*_routes.go、web.go、bot_service.go
  等仍未迁出的文件继续编译。后续步骤把它们迁到各自领域包后可逐步删除。
- test_helpers_test.go:根目录测试沿用 openTestStore 的入口。

go build ./... 与 go test ./... 全部通过;测试数量与重构前一致。

Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
2026-06-18 13:56:36 +08:00
co-authored by Claude
parent 6f6b06e37d
commit eff4972668
32 changed files with 1759 additions and 1462 deletions
+4 -4
View File
@@ -125,11 +125,11 @@ func registerAdminBotRoutes(r gin.IRouter, store *store, sender botTextSender) {
} }
rows, err := store.ListBotMessages(opts) rows, err := store.ListBotMessages(opts)
if err != nil { if err != nil {
writeListResponse(c, rows, opts.listOptions, err, botMessageDTO) writeListResponse(c, rows, opts.ListOptions, err, botMessageDTO)
return return
} }
total, err := store.CountBotMessages(opts) total, err := store.CountBotMessages(opts)
writeListResponseWithTotal(c, rows, opts.listOptions, total, err, botMessageDTO) writeListResponseWithTotal(c, rows, opts.ListOptions, total, err, botMessageDTO)
}) })
r.GET("/bot/direct-messages", func(c *gin.Context) { r.GET("/bot/direct-messages", func(c *gin.Context) {
opts, ok := parseListOptions(c) opts, ok := parseListOptions(c)
@@ -153,7 +153,7 @@ func registerAdminBotRoutes(r gin.IRouter, store *store, sender botTextSender) {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid target node num"}) c.JSON(http.StatusBadRequest, gin.H{"error": "invalid target node num"})
return return
} }
dmOpts := botDirectMessageListOptions{listOptions: opts, BotID: botID, PeerNodeNum: target, Direction: c.Query("direction")} dmOpts := botDirectMessageListOptions{ListOptions: opts, BotID: botID, PeerNodeNum: target, Direction: c.Query("direction")}
rows, err := store.ListBotDirectMessagesByConversation(dmOpts) rows, err := store.ListBotDirectMessagesByConversation(dmOpts)
if err != nil { if err != nil {
writeListResponse(c, rows, opts, err, botDirectMessageDTO) writeListResponse(c, rows, opts, err, botDirectMessageDTO)
@@ -272,7 +272,7 @@ func parseBotMessageListOptions(c *gin.Context) (botMessageListOptions, bool) {
if !ok { if !ok {
return botMessageListOptions{}, false return botMessageListOptions{}, false
} }
opts := botMessageListOptions{listOptions: listOpts, MessageType: c.Query("message_type"), ChannelID: c.Query("channel_id")} opts := botMessageListOptions{ListOptions: listOpts, MessageType: c.Query("message_type"), ChannelID: c.Query("channel_id")}
if value := c.Query("bot_id"); value != "" { if value := c.Query("bot_id"); value != "" {
id, err := strconv.ParseUint(value, 10, 64) id, err := strconv.ParseUint(value, 10, 64)
if err != nil { if err != nil {
+4 -71
View File
@@ -2,17 +2,12 @@ package main
import ( import (
"encoding/base64" "encoding/base64"
"encoding/hex"
"errors"
"strings" "strings"
"gorm.io/gorm"
) )
// pkiKeyResolver mqtpp 在解密 PKI 加密包时回调的接收者私钥/发送者公钥查询函数。 // newPKIKeyResolver 返回 mqtpp 在解密 PKI 加密包时使用的回调:根据接收者
// // 节点号查找受管 bot 的私钥,并根据发送者节点号在 nodeinfo 表中查找其公钥。
// to 是接收者节点号(应该匹配某个本地受管的 bot),from 是发送者节点号(应该已经有 nodeinfo 上报) // 返回 ok=false 时调用方会跳过 PKI 路径并回落到 channel PSK 解密
// 返回的 ok=false 时调用方会跳过 PKI 路径并回落到 channel PSK 解密。
func newPKIKeyResolver(s *store) func(toNodeNum, fromNodeNum uint32) ([]byte, []byte, bool) { func newPKIKeyResolver(s *store) func(toNodeNum, fromNodeNum uint32) ([]byte, []byte, bool) {
if s == nil { if s == nil {
return nil return nil
@@ -20,82 +15,20 @@ func newPKIKeyResolver(s *store) func(toNodeNum, fromNodeNum uint32) ([]byte, []
return func(toNodeNum, fromNodeNum uint32) ([]byte, []byte, bool) { return func(toNodeNum, fromNodeNum uint32) ([]byte, []byte, bool) {
bot, err := s.GetBotNodeByNodeNum(int64(toNodeNum)) bot, err := s.GetBotNodeByNodeNum(int64(toNodeNum))
if err != nil { if err != nil {
printJSON(map[string]any{
"event": "pki_resolve_bot_not_found",
"to_num": toNodeNum,
"from_num": fromNodeNum,
})
return nil, nil, false return nil, nil, false
} }
privateKeyB64 := strings.TrimSpace(bot.PrivateKey) privateKeyB64 := strings.TrimSpace(bot.PrivateKey)
if privateKeyB64 == "" { if privateKeyB64 == "" {
printJSON(map[string]any{
"event": "pki_resolve_no_private_key",
"bot_id": bot.NodeID,
"bot_num": bot.NodeNum,
"from_num": fromNodeNum,
})
return nil, nil, false return nil, nil, false
} }
privateKey, err := base64.StdEncoding.DecodeString(privateKeyB64) privateKey, err := base64.StdEncoding.DecodeString(privateKeyB64)
if err != nil || len(privateKey) != 32 { if err != nil || len(privateKey) != 32 {
printJSON(map[string]any{
"event": "pki_resolve_invalid_private_key",
"bot_id": bot.NodeID,
"bot_num": bot.NodeNum,
"from_num": fromNodeNum,
"error": err,
})
return nil, nil, false return nil, nil, false
} }
fromPublic, ok := lookupNodeInfoPublicKey(s, fromNodeNum) fromPublic, ok := s.LookupNodeInfoPublicKey(fromNodeNum)
if !ok { if !ok {
printJSON(map[string]any{
"event": "pki_resolve_no_sender_public_key",
"bot_id": bot.NodeID,
"bot_num": bot.NodeNum,
"from_num": fromNodeNum,
})
return nil, nil, false return nil, nil, false
} }
return privateKey, fromPublic, true return privateKey, fromPublic, true
} }
} }
// lookupNodeInfoPublicKey 在 nodeinfo 表中按 node_num 查 X25519 公钥,
// 兼容 hex 与 base64 两种历史存储格式。
func lookupNodeInfoPublicKey(s *store, nodeNum uint32) ([]byte, bool) {
var row nodeInfoRecord
if err := s.db.Where("node_num = ?", int64(nodeNum)).Take(&row).Error; err != nil {
return nil, false
}
if row.PublicKey == nil {
return nil, false
}
value := strings.TrimSpace(*row.PublicKey)
if value == "" {
return nil, false
}
if decoded, err := hex.DecodeString(value); err == nil && len(decoded) == 32 {
return decoded, true
}
if decoded, err := base64.StdEncoding.DecodeString(value); err == nil && len(decoded) == 32 {
return decoded, true
}
return nil, false
}
// GetBotNodeByNodeNum 按节点号查找受管 bot 节点;用于 PKI 解密时把 to 字段映射回本地私钥。
func (s *store) GetBotNodeByNodeNum(nodeNum int64) (*botNodeRecord, error) {
if s == nil || s.db == nil {
return nil, errors.New("store not configured")
}
var row botNodeRecord
if err := s.db.Where("node_num = ?", nodeNum).Take(&row).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, err
}
return nil, err
}
return &row, nil
}
+1 -1
View File
@@ -123,7 +123,7 @@ func (s *botService) buildPKIAck(bot *botNodeRecord, toNum, ackPacketID, request
if err != nil { if err != nil {
return nil, err return nil, err
} }
recipientPublic, ok := lookupNodeInfoPublicKey(s.store, toNum) recipientPublic, ok := s.store.LookupNodeInfoPublicKey(toNum)
if !ok { if !ok {
return nil, fmt.Errorf("recipient %s has no public key on file", mqtpp.NodeNumToID(toNum)) return nil, fmt.Errorf("recipient %s has no public key on file", mqtpp.NodeNumToID(toNum))
} }
+28 -529
View File
@@ -2,543 +2,42 @@ package main
import ( import (
cryptotls "crypto/tls" cryptotls "crypto/tls"
"fmt"
"os"
"path/filepath"
"runtime"
"gopkg.in/yaml.v3" cfgpkg "meshtastic_mqtt_server/internal/config"
) )
const configFileName = "config.yaml" // 桥接到 internal/config — 让根目录其余文件无须修改即可继续使用旧的非导出名字。
type config struct { const (
MQTT mqttConfig `yaml:"mqtt"` configFileName = cfgpkg.FileName
Meshtastic meshtasticConfig `yaml:"meshtastic"` databaseDriverSQLite = cfgpkg.DriverSQLite
Database databaseConfig `yaml:"database"` databaseDriverMySQL = cfgpkg.DriverMySQL
Web webConfig `yaml:"web"` )
AI aiConfig `yaml:"ai"`
DataDir string `yaml:"data_dir"`
key []byte
}
type mqttConfig struct { // 旧的小写类型名通过别名继续可用。
Host string `yaml:"host"` type (
Port int `yaml:"port"` config = cfgpkg.Config
TLS tlsConfig `yaml:"tls"` mqttConfig = cfgpkg.MQTTConfig
} tlsConfig = cfgpkg.TLSConfig
meshtasticConfig = cfgpkg.MeshtasticConfig
type tlsConfig struct { databaseConfig = cfgpkg.DatabaseConfig
Enabled bool `yaml:"enabled"` sqliteConfig = cfgpkg.SQLiteConfig
CertFile string `yaml:"cert_file"` mysqlConfig = cfgpkg.MySQLConfig
KeyFile string `yaml:"key_file"` webConfig = cfgpkg.WebConfig
} webAdminConfig = cfgpkg.WebAdminConfig
aiConfig = cfgpkg.AIConfig
type meshtasticConfig struct { )
PSK string `yaml:"psk"`
}
type databaseConfig struct {
Driver string `yaml:"driver"`
SQLite sqliteConfig `yaml:"sqlite"`
MySQL mysqlConfig `yaml:"mysql"`
}
type sqliteConfig struct {
Path string `yaml:"path"`
}
type mysqlConfig struct {
DSN string `yaml:"dsn"`
}
type webConfig struct {
Enabled bool `yaml:"enabled"`
PortEnabled bool `yaml:"port_enabled"`
SocketEnabled bool `yaml:"socket_enabled"`
Host string `yaml:"host"`
Port int `yaml:"port"`
SocketPath string `yaml:"socket_path"`
StaticDir string `yaml:"static_dir"`
MapTileCacheDir string `yaml:"map_tile_cache_dir"`
Admin webAdminConfig `yaml:"admin"`
}
type webAdminConfig struct {
Username string `yaml:"username"`
Password string `yaml:"password"`
SessionSecret string `yaml:"session_secret"`
SessionSecure bool `yaml:"session_secure"`
}
type aiConfig struct {
Enabled bool `yaml:"enabled"`
}
type rawConfig struct {
MQTT *rawMQTTConfig `yaml:"mqtt"`
Meshtastic *rawMeshtasticConfig `yaml:"meshtastic"`
Database *rawDatabaseConfig `yaml:"database"`
Web *rawWebConfig `yaml:"web"`
AI *rawAIConfig `yaml:"ai"`
DataDir *string `yaml:"data_dir"`
}
type rawAIConfig struct {
Enabled *bool `yaml:"enabled"`
}
type rawMQTTConfig struct {
Host *string `yaml:"host"`
Port *int `yaml:"port"`
TLS *rawTLSConfig `yaml:"tls"`
}
type rawTLSConfig struct {
Enabled *bool `yaml:"enabled"`
CertFile *string `yaml:"cert_file"`
KeyFile *string `yaml:"key_file"`
}
type rawMeshtasticConfig struct {
PSK *string `yaml:"psk"`
}
type rawDatabaseConfig struct {
Driver *string `yaml:"driver"`
SQLite *rawSQLiteConfig `yaml:"sqlite"`
MySQL *rawMySQLConfig `yaml:"mysql"`
}
type rawSQLiteConfig struct {
Path *string `yaml:"path"`
}
type rawMySQLConfig struct {
DSN *string `yaml:"dsn"`
}
type rawWebConfig struct {
Enabled *bool `yaml:"enabled"`
PortEnabled *bool `yaml:"port_enabled"`
SocketEnabled *bool `yaml:"socket_enabled"`
Host *string `yaml:"host"`
Port *int `yaml:"port"`
SocketPath *string `yaml:"socket_path"`
StaticDir *string `yaml:"static_dir"`
MapTileCacheDir *string `yaml:"map_tile_cache_dir"`
Admin *rawWebAdminConfig `yaml:"admin"`
}
type rawWebAdminConfig struct {
Username *string `yaml:"username"`
Password *string `yaml:"password"`
SessionSecret *string `yaml:"session_secret"`
SessionSecure *bool `yaml:"session_secure"`
}
// defaultConfig 返回内置默认配置。
func defaultConfig() *config {
return &config{
MQTT: mqttConfig{
Host: "0.0.0.0",
Port: 1883,
TLS: tlsConfig{
Enabled: false,
CertFile: "",
KeyFile: "",
},
},
Meshtastic: meshtasticConfig{
PSK: "AQ==",
},
Database: databaseConfig{
Driver: "sqlite",
SQLite: sqliteConfig{Path: defaultSQLitePath()},
MySQL: mysqlConfig{DSN: ""},
},
Web: webConfig{
Enabled: true,
PortEnabled: true,
SocketEnabled: defaultWebSocketPath() != "",
Host: "0.0.0.0",
Port: 8080,
SocketPath: defaultWebSocketPath(),
StaticDir: "./dist",
MapTileCacheDir: defaultMapTileCacheDir(),
Admin: webAdminConfig{
Username: "admin",
Password: "admin",
SessionSecret: "",
SessionSecure: false,
},
},
AI: aiConfig{
Enabled: false,
},
DataDir: defaultDataDir(),
}
}
// defaultConfigDir 根据操作系统返回配置目录。
func defaultConfigDir() string {
return defaultConfigDirForGOOS(runtime.GOOS)
}
func defaultConfigDirForGOOS(goos string) string {
if useRelativeDefaultPath(goos) {
return filepath.Join(".", "win", "etc", "mesh_mqtt_go")
}
return filepath.Join(string(filepath.Separator), "etc", "mesh_mqtt_go")
}
func useRelativeDefaultPath(goos string) bool {
return goos == "windows" || goos == "darwin"
}
// defaultConfigPath 返回默认配置文件路径。
func defaultConfigPath() string {
return filepath.Join(defaultConfigDir(), configFileName)
}
func defaultSQLitePath() string {
return defaultSQLitePathForGOOS(runtime.GOOS)
}
func defaultWebSocketPath() string {
return defaultWebSocketPathForGOOS(runtime.GOOS)
}
func defaultMapTileCacheDir() string {
return defaultMapTileCacheDirForGOOS(runtime.GOOS)
}
func defaultMapTileCacheDirForGOOS(goos string) string {
if useRelativeDefaultPath(goos) {
return filepath.Join(".", "win", "srv", "mesh_mqtt_go")
}
return filepath.Join(string(filepath.Separator), "srv", "mesh_mqtt_go")
}
func defaultWebSocketPathForGOOS(goos string) string {
if goos == "windows" {
return ""
}
if useRelativeDefaultPath(goos) {
return filepath.Join(".", "win", "opt", "mesh_mqtt_go", "web.sock")
}
return filepath.Join(string(filepath.Separator), "opt", "mesh_mqtt_go", "web.sock")
}
func defaultConfig() *config { return cfgpkg.Default() }
func defaultConfigDir() string { return cfgpkg.DefaultDir() }
func defaultConfigPath() string { return cfgpkg.DefaultPath() }
func loadConfig(path string) (*config, error) { return cfgpkg.Load(path) }
func writeConfig(path string, cfg *config) error { return cfgpkg.Write(path, cfg) }
func validateConfig(cfg *config) error { return cfgpkg.Validate(cfg) }
func clearWebSocketPathOnUnsupportedGOOS(cfg *config, goos string) bool { func clearWebSocketPathOnUnsupportedGOOS(cfg *config, goos string) bool {
if goos != "windows" { return cfgpkg.ClearWebSocketPathOnUnsupportedGOOS(cfg, goos)
return false
}
changed := false
if cfg.Web.SocketPath != "" {
cfg.Web.SocketPath = ""
changed = true
}
if cfg.Web.SocketEnabled {
cfg.Web.SocketEnabled = false
changed = true
}
return changed
} }
func defaultSQLitePathForGOOS(goos string) string {
if useRelativeDefaultPath(goos) {
return filepath.Join(".", "win", "etc", "mesh_mqtt_go", "mesh_mqtt_go.db")
}
return filepath.Join(string(filepath.Separator), "srv", "mesh_mqtt_go", "mesh_mqtt_go.db")
}
func defaultDataDir() string {
return defaultDataDirForGOOS(runtime.GOOS)
}
func defaultDataDirForGOOS(goos string) string {
if useRelativeDefaultPath(goos) {
return filepath.Join(".", "win", "var", "lib", "mesh_mqtt_go")
}
return filepath.Join(string(filepath.Separator), "var", "lib", "mesh_mqtt_go")
}
// loadConfig 加载配置文件;文件不存在时生成,字段缺失时自动补全并写回。
func loadConfig(path string) (*config, error) {
if path == "" {
path = defaultConfigPath()
}
if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil {
return nil, fmt.Errorf("create config directory %s: %w", filepath.Dir(path), err)
}
if _, err := os.Stat(path); err != nil {
if !os.IsNotExist(err) {
return nil, fmt.Errorf("stat config file %s: %w", path, err)
}
cfg := defaultConfig()
if err := writeConfig(path, cfg); err != nil {
return nil, err
}
return cfg, nil
}
data, err := os.ReadFile(path)
if err != nil {
return nil, fmt.Errorf("read config file %s: %w", path, err)
}
var raw rawConfig
if err := yaml.Unmarshal(data, &raw); err != nil {
return nil, fmt.Errorf("parse config file %s: %w", path, err)
}
cfg, changed := normalizeConfig(raw)
if clearWebSocketPathOnUnsupportedGOOS(cfg, runtime.GOOS) {
changed = true
}
if err := validateConfig(cfg); err != nil {
return nil, err
}
if changed {
if err := writeConfig(path, cfg); err != nil {
return nil, err
}
}
return cfg, nil
}
// normalizeConfig 将原始配置合并到默认配置,并标记是否补齐了缺失项。
func normalizeConfig(raw rawConfig) (*config, bool) {
cfg := defaultConfig()
changed := false
if raw.MQTT == nil {
changed = true
} else {
if raw.MQTT.Host == nil {
changed = true
} else {
cfg.MQTT.Host = *raw.MQTT.Host
}
if raw.MQTT.Port == nil {
changed = true
} else {
cfg.MQTT.Port = *raw.MQTT.Port
}
if raw.MQTT.TLS == nil {
changed = true
} else {
if raw.MQTT.TLS.Enabled == nil {
changed = true
} else {
cfg.MQTT.TLS.Enabled = *raw.MQTT.TLS.Enabled
}
if raw.MQTT.TLS.CertFile == nil {
changed = true
} else {
cfg.MQTT.TLS.CertFile = *raw.MQTT.TLS.CertFile
}
if raw.MQTT.TLS.KeyFile == nil {
changed = true
} else {
cfg.MQTT.TLS.KeyFile = *raw.MQTT.TLS.KeyFile
}
}
}
if raw.Meshtastic == nil {
changed = true
} else if raw.Meshtastic.PSK == nil {
changed = true
} else {
cfg.Meshtastic.PSK = *raw.Meshtastic.PSK
}
if raw.Database == nil {
changed = true
} else {
if raw.Database.Driver == nil {
changed = true
} else {
cfg.Database.Driver = *raw.Database.Driver
}
if raw.Database.SQLite == nil {
changed = true
} else if raw.Database.SQLite.Path == nil {
changed = true
} else {
cfg.Database.SQLite.Path = *raw.Database.SQLite.Path
}
if raw.Database.MySQL == nil {
changed = true
} else if raw.Database.MySQL.DSN == nil {
changed = true
} else {
cfg.Database.MySQL.DSN = *raw.Database.MySQL.DSN
}
}
if raw.Web == nil {
changed = true
} else {
if raw.Web.Enabled == nil {
changed = true
} else {
cfg.Web.Enabled = *raw.Web.Enabled
}
if raw.Web.PortEnabled == nil {
changed = true
} else {
cfg.Web.PortEnabled = *raw.Web.PortEnabled
}
if raw.Web.SocketEnabled == nil {
changed = true
} else {
cfg.Web.SocketEnabled = *raw.Web.SocketEnabled
}
if raw.Web.Host == nil {
changed = true
} else {
cfg.Web.Host = *raw.Web.Host
}
if raw.Web.Port == nil {
changed = true
} else {
cfg.Web.Port = *raw.Web.Port
}
if raw.Web.SocketPath == nil {
changed = true
} else {
cfg.Web.SocketPath = *raw.Web.SocketPath
}
if raw.Web.StaticDir == nil {
changed = true
} else {
cfg.Web.StaticDir = *raw.Web.StaticDir
}
if raw.Web.MapTileCacheDir == nil {
changed = true
} else {
cfg.Web.MapTileCacheDir = *raw.Web.MapTileCacheDir
}
if raw.Web.Admin == nil {
changed = true
} else {
if raw.Web.Admin.Username == nil {
changed = true
} else {
cfg.Web.Admin.Username = *raw.Web.Admin.Username
}
if raw.Web.Admin.Password == nil {
changed = true
} else {
cfg.Web.Admin.Password = *raw.Web.Admin.Password
}
if raw.Web.Admin.SessionSecret == nil {
changed = true
} else {
cfg.Web.Admin.SessionSecret = *raw.Web.Admin.SessionSecret
}
if raw.Web.Admin.SessionSecure == nil {
changed = true
} else {
cfg.Web.Admin.SessionSecure = *raw.Web.Admin.SessionSecure
}
}
}
if raw.AI == nil {
changed = true
} else {
if raw.AI.Enabled == nil {
changed = true
} else {
cfg.AI.Enabled = *raw.AI.Enabled
}
}
if raw.DataDir == nil {
changed = true
} else {
cfg.DataDir = *raw.DataDir
}
return cfg, changed
}
func validateConfig(cfg *config) error {
if cfg.MQTT.Port <= 0 || cfg.MQTT.Port > 65535 {
return fmt.Errorf("invalid mqtt port %d: must be 1-65535", cfg.MQTT.Port)
}
switch cfg.Database.Driver {
case "sqlite":
if cfg.Database.SQLite.Path == "" {
return fmt.Errorf("database.sqlite.path is required when database.driver is sqlite")
}
case "mysql":
if cfg.Database.MySQL.DSN == "" {
return fmt.Errorf("database.mysql.dsn is required when database.driver is mysql")
}
default:
return fmt.Errorf("invalid database.driver %q: must be sqlite or mysql", cfg.Database.Driver)
}
if cfg.Web.Enabled {
if !cfg.Web.PortEnabled && !cfg.Web.SocketEnabled {
return fmt.Errorf("web.port_enabled and web.socket_enabled cannot both be false when web is enabled")
}
if cfg.Web.PortEnabled && (cfg.Web.Port <= 0 || cfg.Web.Port > 65535) {
return fmt.Errorf("invalid web port %d: must be 1-65535", cfg.Web.Port)
}
if cfg.Web.SocketEnabled && cfg.Web.SocketPath == "" {
return fmt.Errorf("web.socket_path is required when web.socket_enabled is true")
}
if cfg.Web.StaticDir == "" {
return fmt.Errorf("web.static_dir is required when web is enabled")
}
if cfg.Web.MapTileCacheDir == "" {
return fmt.Errorf("web.map_tile_cache_dir is required when web is enabled")
}
if cfg.Web.Admin.Username == "" {
return fmt.Errorf("web.admin.username is required when web is enabled")
}
if cfg.Web.Admin.Password == "" {
return fmt.Errorf("web.admin.password is required when web is enabled")
}
}
return nil
}
func writeConfig(path string, cfg *config) error {
data, err := yaml.Marshal(cfg)
if err != nil {
return fmt.Errorf("encode config file %s: %w", path, err)
}
if err := os.WriteFile(path, data, 0644); err != nil {
return fmt.Errorf("write config file %s: %w", path, err)
}
return nil
}
// buildTLSConfig 根据配置构造 mochi listener 使用的 TLS 设置。
func buildTLSConfig(cfg tlsConfig) (*cryptotls.Config, error) { func buildTLSConfig(cfg tlsConfig) (*cryptotls.Config, error) {
if !cfg.Enabled { return cfgpkg.BuildTLS(cfg)
return nil, nil
}
if cfg.CertFile == "" {
return nil, fmt.Errorf("mqtt tls cert_file is required when tls is enabled")
}
if cfg.KeyFile == "" {
return nil, fmt.Errorf("mqtt tls key_file is required when tls is enabled")
}
cert, err := cryptotls.LoadX509KeyPair(cfg.CertFile, cfg.KeyFile)
if err != nil {
return nil, fmt.Errorf("load mqtt tls certificate: %w", err)
}
return &cryptotls.Config{
MinVersion: cryptotls.VersionTLS12,
Certificates: []cryptotls.Certificate{cert},
}, nil
} }
+549
View File
@@ -0,0 +1,549 @@
package config
import (
cryptotls "crypto/tls"
"fmt"
"os"
"path/filepath"
"runtime"
"gopkg.in/yaml.v3"
)
const FileName = "config.yaml"
const (
DriverSQLite = "sqlite"
DriverMySQL = "mysql"
)
type Config struct {
MQTT MQTTConfig `yaml:"mqtt"`
Meshtastic MeshtasticConfig `yaml:"meshtastic"`
Database DatabaseConfig `yaml:"database"`
Web WebConfig `yaml:"web"`
AI AIConfig `yaml:"ai"`
DataDir string `yaml:"data_dir"`
Key []byte `yaml:"-"`
}
type MQTTConfig struct {
Host string `yaml:"host"`
Port int `yaml:"port"`
TLS TLSConfig `yaml:"tls"`
}
type TLSConfig struct {
Enabled bool `yaml:"enabled"`
CertFile string `yaml:"cert_file"`
KeyFile string `yaml:"key_file"`
}
type MeshtasticConfig struct {
PSK string `yaml:"psk"`
}
type DatabaseConfig struct {
Driver string `yaml:"driver"`
SQLite SQLiteConfig `yaml:"sqlite"`
MySQL MySQLConfig `yaml:"mysql"`
}
type SQLiteConfig struct {
Path string `yaml:"path"`
}
type MySQLConfig struct {
DSN string `yaml:"dsn"`
}
type WebConfig struct {
Enabled bool `yaml:"enabled"`
PortEnabled bool `yaml:"port_enabled"`
SocketEnabled bool `yaml:"socket_enabled"`
Host string `yaml:"host"`
Port int `yaml:"port"`
SocketPath string `yaml:"socket_path"`
StaticDir string `yaml:"static_dir"`
MapTileCacheDir string `yaml:"map_tile_cache_dir"`
Admin WebAdminConfig `yaml:"admin"`
}
type WebAdminConfig struct {
Username string `yaml:"username"`
Password string `yaml:"password"`
SessionSecret string `yaml:"session_secret"`
SessionSecure bool `yaml:"session_secure"`
}
type AIConfig struct {
Enabled bool `yaml:"enabled"`
}
type rawConfig struct {
MQTT *rawMQTTConfig `yaml:"mqtt"`
Meshtastic *rawMeshtasticConfig `yaml:"meshtastic"`
Database *rawDatabaseConfig `yaml:"database"`
Web *rawWebConfig `yaml:"web"`
AI *rawAIConfig `yaml:"ai"`
DataDir *string `yaml:"data_dir"`
}
type rawAIConfig struct {
Enabled *bool `yaml:"enabled"`
}
type rawMQTTConfig struct {
Host *string `yaml:"host"`
Port *int `yaml:"port"`
TLS *rawTLSConfig `yaml:"tls"`
}
type rawTLSConfig struct {
Enabled *bool `yaml:"enabled"`
CertFile *string `yaml:"cert_file"`
KeyFile *string `yaml:"key_file"`
}
type rawMeshtasticConfig struct {
PSK *string `yaml:"psk"`
}
type rawDatabaseConfig struct {
Driver *string `yaml:"driver"`
SQLite *rawSQLiteConfig `yaml:"sqlite"`
MySQL *rawMySQLConfig `yaml:"mysql"`
}
type rawSQLiteConfig struct {
Path *string `yaml:"path"`
}
type rawMySQLConfig struct {
DSN *string `yaml:"dsn"`
}
type rawWebConfig struct {
Enabled *bool `yaml:"enabled"`
PortEnabled *bool `yaml:"port_enabled"`
SocketEnabled *bool `yaml:"socket_enabled"`
Host *string `yaml:"host"`
Port *int `yaml:"port"`
SocketPath *string `yaml:"socket_path"`
StaticDir *string `yaml:"static_dir"`
MapTileCacheDir *string `yaml:"map_tile_cache_dir"`
Admin *rawWebAdminConfig `yaml:"admin"`
}
type rawWebAdminConfig struct {
Username *string `yaml:"username"`
Password *string `yaml:"password"`
SessionSecret *string `yaml:"session_secret"`
SessionSecure *bool `yaml:"session_secure"`
}
// Default 返回内置默认配置。
func Default() *Config {
return &Config{
MQTT: MQTTConfig{
Host: "0.0.0.0",
Port: 1883,
TLS: TLSConfig{
Enabled: false,
CertFile: "",
KeyFile: "",
},
},
Meshtastic: MeshtasticConfig{
PSK: "AQ==",
},
Database: DatabaseConfig{
Driver: DriverSQLite,
SQLite: SQLiteConfig{Path: defaultSQLitePath()},
MySQL: MySQLConfig{DSN: ""},
},
Web: WebConfig{
Enabled: true,
PortEnabled: true,
SocketEnabled: defaultWebSocketPath() != "",
Host: "0.0.0.0",
Port: 8080,
SocketPath: defaultWebSocketPath(),
StaticDir: "./dist",
MapTileCacheDir: defaultMapTileCacheDir(),
Admin: WebAdminConfig{
Username: "admin",
Password: "admin",
SessionSecret: "",
SessionSecure: false,
},
},
AI: AIConfig{
Enabled: false,
},
DataDir: defaultDataDir(),
}
}
// DefaultDir 根据操作系统返回配置目录。
func DefaultDir() string {
return defaultConfigDirForGOOS(runtime.GOOS)
}
func defaultConfigDirForGOOS(goos string) string {
if useRelativeDefaultPath(goos) {
return filepath.Join(".", "win", "etc", "mesh_mqtt_go")
}
return filepath.Join(string(filepath.Separator), "etc", "mesh_mqtt_go")
}
func useRelativeDefaultPath(goos string) bool {
return goos == "windows" || goos == "darwin"
}
// DefaultPath 返回默认配置文件路径。
func DefaultPath() string {
return filepath.Join(DefaultDir(), FileName)
}
func defaultSQLitePath() string {
return defaultSQLitePathForGOOS(runtime.GOOS)
}
func defaultWebSocketPath() string {
return defaultWebSocketPathForGOOS(runtime.GOOS)
}
func defaultMapTileCacheDir() string {
return defaultMapTileCacheDirForGOOS(runtime.GOOS)
}
func defaultMapTileCacheDirForGOOS(goos string) string {
if useRelativeDefaultPath(goos) {
return filepath.Join(".", "win", "srv", "mesh_mqtt_go")
}
return filepath.Join(string(filepath.Separator), "srv", "mesh_mqtt_go")
}
func defaultWebSocketPathForGOOS(goos string) string {
if goos == "windows" {
return ""
}
if useRelativeDefaultPath(goos) {
return filepath.Join(".", "win", "opt", "mesh_mqtt_go", "web.sock")
}
return filepath.Join(string(filepath.Separator), "opt", "mesh_mqtt_go", "web.sock")
}
func ClearWebSocketPathOnUnsupportedGOOS(cfg *Config, goos string) bool {
if goos != "windows" {
return false
}
changed := false
if cfg.Web.SocketPath != "" {
cfg.Web.SocketPath = ""
changed = true
}
if cfg.Web.SocketEnabled {
cfg.Web.SocketEnabled = false
changed = true
}
return changed
}
func defaultSQLitePathForGOOS(goos string) string {
if useRelativeDefaultPath(goos) {
return filepath.Join(".", "win", "etc", "mesh_mqtt_go", "mesh_mqtt_go.db")
}
return filepath.Join(string(filepath.Separator), "srv", "mesh_mqtt_go", "mesh_mqtt_go.db")
}
func defaultDataDir() string {
return defaultDataDirForGOOS(runtime.GOOS)
}
func defaultDataDirForGOOS(goos string) string {
if useRelativeDefaultPath(goos) {
return filepath.Join(".", "win", "var", "lib", "mesh_mqtt_go")
}
return filepath.Join(string(filepath.Separator), "var", "lib", "mesh_mqtt_go")
}
// Load 加载配置文件;文件不存在时生成,字段缺失时自动补全并写回。
func Load(path string) (*Config, error) {
if path == "" {
path = DefaultPath()
}
if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil {
return nil, fmt.Errorf("create config directory %s: %w", filepath.Dir(path), err)
}
if _, err := os.Stat(path); err != nil {
if !os.IsNotExist(err) {
return nil, fmt.Errorf("stat config file %s: %w", path, err)
}
cfg := Default()
if err := Write(path, cfg); err != nil {
return nil, err
}
return cfg, nil
}
data, err := os.ReadFile(path)
if err != nil {
return nil, fmt.Errorf("read config file %s: %w", path, err)
}
var raw rawConfig
if err := yaml.Unmarshal(data, &raw); err != nil {
return nil, fmt.Errorf("parse config file %s: %w", path, err)
}
cfg, changed := normalize(raw)
if ClearWebSocketPathOnUnsupportedGOOS(cfg, runtime.GOOS) {
changed = true
}
if err := Validate(cfg); err != nil {
return nil, err
}
if changed {
if err := Write(path, cfg); err != nil {
return nil, err
}
}
return cfg, nil
}
// normalize 将原始配置合并到默认配置,并标记是否补齐了缺失项。
func normalize(raw rawConfig) (*Config, bool) {
cfg := Default()
changed := false
if raw.MQTT == nil {
changed = true
} else {
if raw.MQTT.Host == nil {
changed = true
} else {
cfg.MQTT.Host = *raw.MQTT.Host
}
if raw.MQTT.Port == nil {
changed = true
} else {
cfg.MQTT.Port = *raw.MQTT.Port
}
if raw.MQTT.TLS == nil {
changed = true
} else {
if raw.MQTT.TLS.Enabled == nil {
changed = true
} else {
cfg.MQTT.TLS.Enabled = *raw.MQTT.TLS.Enabled
}
if raw.MQTT.TLS.CertFile == nil {
changed = true
} else {
cfg.MQTT.TLS.CertFile = *raw.MQTT.TLS.CertFile
}
if raw.MQTT.TLS.KeyFile == nil {
changed = true
} else {
cfg.MQTT.TLS.KeyFile = *raw.MQTT.TLS.KeyFile
}
}
}
if raw.Meshtastic == nil {
changed = true
} else if raw.Meshtastic.PSK == nil {
changed = true
} else {
cfg.Meshtastic.PSK = *raw.Meshtastic.PSK
}
if raw.Database == nil {
changed = true
} else {
if raw.Database.Driver == nil {
changed = true
} else {
cfg.Database.Driver = *raw.Database.Driver
}
if raw.Database.SQLite == nil {
changed = true
} else if raw.Database.SQLite.Path == nil {
changed = true
} else {
cfg.Database.SQLite.Path = *raw.Database.SQLite.Path
}
if raw.Database.MySQL == nil {
changed = true
} else if raw.Database.MySQL.DSN == nil {
changed = true
} else {
cfg.Database.MySQL.DSN = *raw.Database.MySQL.DSN
}
}
if raw.Web == nil {
changed = true
} else {
if raw.Web.Enabled == nil {
changed = true
} else {
cfg.Web.Enabled = *raw.Web.Enabled
}
if raw.Web.PortEnabled == nil {
changed = true
} else {
cfg.Web.PortEnabled = *raw.Web.PortEnabled
}
if raw.Web.SocketEnabled == nil {
changed = true
} else {
cfg.Web.SocketEnabled = *raw.Web.SocketEnabled
}
if raw.Web.Host == nil {
changed = true
} else {
cfg.Web.Host = *raw.Web.Host
}
if raw.Web.Port == nil {
changed = true
} else {
cfg.Web.Port = *raw.Web.Port
}
if raw.Web.SocketPath == nil {
changed = true
} else {
cfg.Web.SocketPath = *raw.Web.SocketPath
}
if raw.Web.StaticDir == nil {
changed = true
} else {
cfg.Web.StaticDir = *raw.Web.StaticDir
}
if raw.Web.MapTileCacheDir == nil {
changed = true
} else {
cfg.Web.MapTileCacheDir = *raw.Web.MapTileCacheDir
}
if raw.Web.Admin == nil {
changed = true
} else {
if raw.Web.Admin.Username == nil {
changed = true
} else {
cfg.Web.Admin.Username = *raw.Web.Admin.Username
}
if raw.Web.Admin.Password == nil {
changed = true
} else {
cfg.Web.Admin.Password = *raw.Web.Admin.Password
}
if raw.Web.Admin.SessionSecret == nil {
changed = true
} else {
cfg.Web.Admin.SessionSecret = *raw.Web.Admin.SessionSecret
}
if raw.Web.Admin.SessionSecure == nil {
changed = true
} else {
cfg.Web.Admin.SessionSecure = *raw.Web.Admin.SessionSecure
}
}
}
if raw.AI == nil {
changed = true
} else {
if raw.AI.Enabled == nil {
changed = true
} else {
cfg.AI.Enabled = *raw.AI.Enabled
}
}
if raw.DataDir == nil {
changed = true
} else {
cfg.DataDir = *raw.DataDir
}
return cfg, changed
}
func Validate(cfg *Config) error {
if cfg.MQTT.Port <= 0 || cfg.MQTT.Port > 65535 {
return fmt.Errorf("invalid mqtt port %d: must be 1-65535", cfg.MQTT.Port)
}
switch cfg.Database.Driver {
case DriverSQLite:
if cfg.Database.SQLite.Path == "" {
return fmt.Errorf("database.sqlite.path is required when database.driver is sqlite")
}
case DriverMySQL:
if cfg.Database.MySQL.DSN == "" {
return fmt.Errorf("database.mysql.dsn is required when database.driver is mysql")
}
default:
return fmt.Errorf("invalid database.driver %q: must be sqlite or mysql", cfg.Database.Driver)
}
if cfg.Web.Enabled {
if !cfg.Web.PortEnabled && !cfg.Web.SocketEnabled {
return fmt.Errorf("web.port_enabled and web.socket_enabled cannot both be false when web is enabled")
}
if cfg.Web.PortEnabled && (cfg.Web.Port <= 0 || cfg.Web.Port > 65535) {
return fmt.Errorf("invalid web port %d: must be 1-65535", cfg.Web.Port)
}
if cfg.Web.SocketEnabled && cfg.Web.SocketPath == "" {
return fmt.Errorf("web.socket_path is required when web.socket_enabled is true")
}
if cfg.Web.StaticDir == "" {
return fmt.Errorf("web.static_dir is required when web is enabled")
}
if cfg.Web.MapTileCacheDir == "" {
return fmt.Errorf("web.map_tile_cache_dir is required when web is enabled")
}
if cfg.Web.Admin.Username == "" {
return fmt.Errorf("web.admin.username is required when web is enabled")
}
if cfg.Web.Admin.Password == "" {
return fmt.Errorf("web.admin.password is required when web is enabled")
}
}
return nil
}
func Write(path string, cfg *Config) error {
data, err := yaml.Marshal(cfg)
if err != nil {
return fmt.Errorf("encode config file %s: %w", path, err)
}
if err := os.WriteFile(path, data, 0644); err != nil {
return fmt.Errorf("write config file %s: %w", path, err)
}
return nil
}
// BuildTLS 根据配置构造 mochi listener 使用的 TLS 设置。
func BuildTLS(cfg TLSConfig) (*cryptotls.Config, error) {
if !cfg.Enabled {
return nil, nil
}
if cfg.CertFile == "" {
return nil, fmt.Errorf("mqtt tls cert_file is required when tls is enabled")
}
if cfg.KeyFile == "" {
return nil, fmt.Errorf("mqtt tls key_file is required when tls is enabled")
}
cert, err := cryptotls.LoadX509KeyPair(cfg.CertFile, cfg.KeyFile)
if err != nil {
return nil, fmt.Errorf("load mqtt tls certificate: %w", err)
}
return &cryptotls.Config{
MinVersion: cryptotls.VersionTLS12,
Certificates: []cryptotls.Certificate{cert},
}, nil
}
@@ -1,4 +1,4 @@
package main package config
import ( import (
"os" "os"
@@ -8,11 +8,11 @@ import (
) )
func TestLoadConfigCreatesDefaultFile(t *testing.T) { func TestLoadConfigCreatesDefaultFile(t *testing.T) {
path := filepath.Join(t.TempDir(), "mesh_mqtt_go", configFileName) path := filepath.Join(t.TempDir(), "mesh_mqtt_go", FileName)
cfg, err := loadConfig(path) cfg, err := Load(path)
if err != nil { if err != nil {
t.Fatalf("loadConfig() error = %v", err) t.Fatalf("Load() error = %v", err)
} }
if cfg.MQTT.Host != "0.0.0.0" { if cfg.MQTT.Host != "0.0.0.0" {
t.Fatalf("host = %q, want 0.0.0.0", cfg.MQTT.Host) t.Fatalf("host = %q, want 0.0.0.0", cfg.MQTT.Host)
@@ -60,7 +60,7 @@ func TestLoadConfigCreatesDefaultFile(t *testing.T) {
} }
func TestLoadConfigFillsMissingFields(t *testing.T) { func TestLoadConfigFillsMissingFields(t *testing.T) {
path := filepath.Join(t.TempDir(), "mesh_mqtt_go", configFileName) path := filepath.Join(t.TempDir(), "mesh_mqtt_go", FileName)
if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil { if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -68,9 +68,9 @@ func TestLoadConfigFillsMissingFields(t *testing.T) {
t.Fatal(err) t.Fatal(err)
} }
cfg, err := loadConfig(path) cfg, err := Load(path)
if err != nil { if err != nil {
t.Fatalf("loadConfig() error = %v", err) t.Fatalf("Load() error = %v", err)
} }
if cfg.MQTT.Port != 1884 { if cfg.MQTT.Port != 1884 {
t.Fatalf("port = %d, want 1884", cfg.MQTT.Port) t.Fatalf("port = %d, want 1884", cfg.MQTT.Port)
@@ -98,7 +98,7 @@ func TestLoadConfigFillsMissingFields(t *testing.T) {
} }
func TestLoadConfigPreservesExplicitFalse(t *testing.T) { func TestLoadConfigPreservesExplicitFalse(t *testing.T) {
path := filepath.Join(t.TempDir(), "mesh_mqtt_go", configFileName) path := filepath.Join(t.TempDir(), "mesh_mqtt_go", FileName)
if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil { if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -107,9 +107,9 @@ func TestLoadConfigPreservesExplicitFalse(t *testing.T) {
t.Fatal(err) t.Fatal(err)
} }
cfg, err := loadConfig(path) cfg, err := Load(path)
if err != nil { if err != nil {
t.Fatalf("loadConfig() error = %v", err) t.Fatalf("Load() error = %v", err)
} }
if cfg.MQTT.TLS.Enabled { if cfg.MQTT.TLS.Enabled {
t.Fatalf("tls enabled = true, want explicit false") t.Fatalf("tls enabled = true, want explicit false")
@@ -120,7 +120,7 @@ func TestLoadConfigPreservesExplicitFalse(t *testing.T) {
} }
func TestLoadConfigPreservesExplicitWebFalse(t *testing.T) { func TestLoadConfigPreservesExplicitWebFalse(t *testing.T) {
path := filepath.Join(t.TempDir(), "mesh_mqtt_go", configFileName) path := filepath.Join(t.TempDir(), "mesh_mqtt_go", FileName)
if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil { if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -129,9 +129,9 @@ func TestLoadConfigPreservesExplicitWebFalse(t *testing.T) {
t.Fatal(err) t.Fatal(err)
} }
cfg, err := loadConfig(path) cfg, err := Load(path)
if err != nil { if err != nil {
t.Fatalf("loadConfig() error = %v", err) t.Fatalf("Load() error = %v", err)
} }
if cfg.Web.Enabled { if cfg.Web.Enabled {
t.Fatalf("web enabled = true, want explicit false") t.Fatalf("web enabled = true, want explicit false")
@@ -145,7 +145,7 @@ func TestLoadConfigPreservesExplicitWebFalse(t *testing.T) {
} }
func TestLoadConfigMalformedYAMLDoesNotOverwrite(t *testing.T) { func TestLoadConfigMalformedYAMLDoesNotOverwrite(t *testing.T) {
path := filepath.Join(t.TempDir(), "mesh_mqtt_go", configFileName) path := filepath.Join(t.TempDir(), "mesh_mqtt_go", FileName)
if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil { if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -154,9 +154,9 @@ func TestLoadConfigMalformedYAMLDoesNotOverwrite(t *testing.T) {
t.Fatal(err) t.Fatal(err)
} }
_, err := loadConfig(path) _, err := Load(path)
if err == nil { if err == nil {
t.Fatalf("loadConfig() error = nil, want parse error") t.Fatalf("Load() error = nil, want parse error")
} }
data, readErr := os.ReadFile(path) data, readErr := os.ReadFile(path)
if readErr != nil { if readErr != nil {
@@ -218,10 +218,10 @@ func TestDefaultWebSocketPathForGOOS(t *testing.T) {
} }
func TestClearWebSocketPathOnUnsupportedGOOS(t *testing.T) { func TestClearWebSocketPathOnUnsupportedGOOS(t *testing.T) {
cfg := defaultConfig() cfg := Default()
cfg.Web.SocketPath = filepath.Join(".", "win", "opt", "mesh_mqtt_go", "web.sock") cfg.Web.SocketPath = filepath.Join(".", "win", "opt", "mesh_mqtt_go", "web.sock")
if !clearWebSocketPathOnUnsupportedGOOS(cfg, "windows") { if !ClearWebSocketPathOnUnsupportedGOOS(cfg, "windows") {
t.Fatalf("clearWebSocketPathOnUnsupportedGOOS() = false, want true") t.Fatalf("ClearWebSocketPathOnUnsupportedGOOS() = false, want true")
} }
if cfg.Web.SocketPath != "" { if cfg.Web.SocketPath != "" {
t.Fatalf("windows web socket path = %q, want empty", cfg.Web.SocketPath) t.Fatalf("windows web socket path = %q, want empty", cfg.Web.SocketPath)
@@ -231,8 +231,8 @@ func TestClearWebSocketPathOnUnsupportedGOOS(t *testing.T) {
} }
cfg.Web.SocketPath = "/opt/mesh_mqtt_go/web.sock" cfg.Web.SocketPath = "/opt/mesh_mqtt_go/web.sock"
if clearWebSocketPathOnUnsupportedGOOS(cfg, "linux") { if ClearWebSocketPathOnUnsupportedGOOS(cfg, "linux") {
t.Fatalf("linux clearWebSocketPathOnUnsupportedGOOS() = true, want false") t.Fatalf("linux ClearWebSocketPathOnUnsupportedGOOS() = true, want false")
} }
if cfg.Web.SocketPath == "" { if cfg.Web.SocketPath == "" {
t.Fatalf("linux web socket path was cleared") t.Fatalf("linux web socket path was cleared")
@@ -256,102 +256,102 @@ func TestDefaultSQLitePathForGOOS(t *testing.T) {
} }
func TestValidateConfigDatabase(t *testing.T) { func TestValidateConfigDatabase(t *testing.T) {
cfg := defaultConfig() cfg := Default()
cfg.Database.Driver = "postgres" cfg.Database.Driver = "postgres"
if err := validateConfig(cfg); err == nil || !strings.Contains(err.Error(), "database.driver") { if err := Validate(cfg); err == nil || !strings.Contains(err.Error(), "database.driver") {
t.Fatalf("invalid driver error = %v, want database.driver error", err) t.Fatalf("invalid driver error = %v, want database.driver error", err)
} }
cfg = defaultConfig() cfg = Default()
cfg.Database.SQLite.Path = "" cfg.Database.SQLite.Path = ""
if err := validateConfig(cfg); err == nil || !strings.Contains(err.Error(), "database.sqlite.path") { if err := Validate(cfg); err == nil || !strings.Contains(err.Error(), "database.sqlite.path") {
t.Fatalf("missing sqlite path error = %v, want database.sqlite.path error", err) t.Fatalf("missing sqlite path error = %v, want database.sqlite.path error", err)
} }
cfg = defaultConfig() cfg = Default()
cfg.Database.Driver = "mysql" cfg.Database.Driver = "mysql"
cfg.Database.MySQL.DSN = "" cfg.Database.MySQL.DSN = ""
if err := validateConfig(cfg); err == nil || !strings.Contains(err.Error(), "database.mysql.dsn") { if err := Validate(cfg); err == nil || !strings.Contains(err.Error(), "database.mysql.dsn") {
t.Fatalf("missing mysql dsn error = %v, want database.mysql.dsn error", err) t.Fatalf("missing mysql dsn error = %v, want database.mysql.dsn error", err)
} }
} }
func TestValidateConfigWeb(t *testing.T) { func TestValidateConfigWeb(t *testing.T) {
cfg := defaultConfig() cfg := Default()
cfg.Web.Port = 0 cfg.Web.Port = 0
if err := validateConfig(cfg); err == nil || !strings.Contains(err.Error(), "web port") { if err := Validate(cfg); err == nil || !strings.Contains(err.Error(), "web port") {
t.Fatalf("invalid web port error = %v, want web port error", err) t.Fatalf("invalid web port error = %v, want web port error", err)
} }
cfg = defaultConfig() cfg = Default()
cfg.Web.PortEnabled = false cfg.Web.PortEnabled = false
cfg.Web.Port = 0 cfg.Web.Port = 0
if err := validateConfig(cfg); err != nil { if err := Validate(cfg); err != nil {
t.Fatalf("disabled web port with invalid port error = %v, want nil", err) t.Fatalf("disabled web port with invalid port error = %v, want nil", err)
} }
cfg = defaultConfig() cfg = Default()
cfg.Web.SocketEnabled = false cfg.Web.SocketEnabled = false
cfg.Web.SocketPath = "" cfg.Web.SocketPath = ""
if err := validateConfig(cfg); err != nil { if err := Validate(cfg); err != nil {
t.Fatalf("disabled web socket with empty path error = %v, want nil", err) t.Fatalf("disabled web socket with empty path error = %v, want nil", err)
} }
cfg = defaultConfig() cfg = Default()
cfg.Web.PortEnabled = false cfg.Web.PortEnabled = false
cfg.Web.SocketEnabled = true cfg.Web.SocketEnabled = true
cfg.Web.SocketPath = "" cfg.Web.SocketPath = ""
if err := validateConfig(cfg); err == nil || !strings.Contains(err.Error(), "web.socket_path") { if err := Validate(cfg); err == nil || !strings.Contains(err.Error(), "web.socket_path") {
t.Fatalf("missing web socket path error = %v, want web.socket_path error", err) t.Fatalf("missing web socket path error = %v, want web.socket_path error", err)
} }
cfg = defaultConfig() cfg = Default()
cfg.Web.PortEnabled = false cfg.Web.PortEnabled = false
cfg.Web.SocketEnabled = false cfg.Web.SocketEnabled = false
if err := validateConfig(cfg); err == nil || !strings.Contains(err.Error(), "web.port_enabled") { if err := Validate(cfg); err == nil || !strings.Contains(err.Error(), "web.port_enabled") {
t.Fatalf("disabled web listeners error = %v, want web.port_enabled error", err) t.Fatalf("disabled web listeners error = %v, want web.port_enabled error", err)
} }
cfg = defaultConfig() cfg = Default()
cfg.Web.StaticDir = "" cfg.Web.StaticDir = ""
if err := validateConfig(cfg); err == nil || !strings.Contains(err.Error(), "web.static_dir") { if err := Validate(cfg); err == nil || !strings.Contains(err.Error(), "web.static_dir") {
t.Fatalf("missing web static dir error = %v, want web.static_dir error", err) t.Fatalf("missing web static dir error = %v, want web.static_dir error", err)
} }
cfg = defaultConfig() cfg = Default()
cfg.Web.MapTileCacheDir = "" cfg.Web.MapTileCacheDir = ""
if err := validateConfig(cfg); err == nil || !strings.Contains(err.Error(), "web.map_tile_cache_dir") { if err := Validate(cfg); err == nil || !strings.Contains(err.Error(), "web.map_tile_cache_dir") {
t.Fatalf("missing map tile cache dir error = %v, want web.map_tile_cache_dir error", err) t.Fatalf("missing map tile cache dir error = %v, want web.map_tile_cache_dir error", err)
} }
cfg = defaultConfig() cfg = Default()
cfg.Web.Enabled = false cfg.Web.Enabled = false
cfg.Web.PortEnabled = false cfg.Web.PortEnabled = false
cfg.Web.SocketEnabled = false cfg.Web.SocketEnabled = false
cfg.Web.Port = 0 cfg.Web.Port = 0
cfg.Web.StaticDir = "" cfg.Web.StaticDir = ""
if err := validateConfig(cfg); err != nil { if err := Validate(cfg); err != nil {
t.Fatalf("disabled web validate error = %v, want nil", err) t.Fatalf("disabled web validate error = %v, want nil", err)
} }
} }
func TestBuildTLSConfigDisabled(t *testing.T) { func TestBuildTLSConfigDisabled(t *testing.T) {
cfg, err := buildTLSConfig(tlsConfig{}) cfg, err := BuildTLS(TLSConfig{})
if err != nil { if err != nil {
t.Fatalf("buildTLSConfig() error = %v", err) t.Fatalf("BuildTLS() error = %v", err)
} }
if cfg != nil { if cfg != nil {
t.Fatalf("buildTLSConfig() = %#v, want nil", cfg) t.Fatalf("BuildTLS() = %#v, want nil", cfg)
} }
} }
func TestBuildTLSConfigRequiresCertAndKey(t *testing.T) { func TestBuildTLSConfigRequiresCertAndKey(t *testing.T) {
_, err := buildTLSConfig(tlsConfig{Enabled: true}) _, err := BuildTLS(TLSConfig{Enabled: true})
if err == nil || !strings.Contains(err.Error(), "cert_file") { if err == nil || !strings.Contains(err.Error(), "cert_file") {
t.Fatalf("missing cert error = %v, want cert_file error", err) t.Fatalf("missing cert error = %v, want cert_file error", err)
} }
_, err = buildTLSConfig(tlsConfig{Enabled: true, CertFile: "cert.pem"}) _, err = BuildTLS(TLSConfig{Enabled: true, CertFile: "cert.pem"})
if err == nil || !strings.Contains(err.Error(), "key_file") { if err == nil || !strings.Contains(err.Error(), "key_file") {
t.Fatalf("missing key error = %v, want key_file error", err) t.Fatalf("missing key error = %v, want key_file error", err)
} }
@@ -1,4 +1,4 @@
package main package store
import ( import (
"errors" "errors"
@@ -10,14 +10,14 @@ import (
"gorm.io/gorm" "gorm.io/gorm"
) )
const forbiddenWordMatchContains = "contains" const ForbiddenWordMatchContains = "contains"
var errBlockingAlreadyExists = errors.New("blocking rule already exists") var ErrBlockingAlreadyExists = errors.New("blocking rule already exists")
func (s *store) ListNodeBlocking(opts listOptions) ([]nodeBlockingRecord, error) { func (s *Store) ListNodeBlocking(opts ListOptions) ([]NodeBlockingRecord, error) {
opts = normalizeListOptions(opts) opts = NormalizeListOptions(opts)
var rows []nodeBlockingRecord var rows []NodeBlockingRecord
q := s.db.Model(&nodeBlockingRecord{}). q := s.db.Model(&NodeBlockingRecord{}).
Order("updated_at DESC"). Order("updated_at DESC").
Order("id DESC"). Order("id DESC").
Limit(opts.Limit). Limit(opts.Limit).
@@ -25,17 +25,17 @@ func (s *store) ListNodeBlocking(opts listOptions) ([]nodeBlockingRecord, error)
return rows, q.Find(&rows).Error return rows, q.Find(&rows).Error
} }
func (s *store) CountNodeBlocking(opts listOptions) (int64, error) { func (s *Store) CountNodeBlocking(opts ListOptions) (int64, error) {
var total int64 var total int64
return total, s.db.Model(&nodeBlockingRecord{}).Count(&total).Error return total, s.db.Model(&NodeBlockingRecord{}).Count(&total).Error
} }
func (s *store) ListEnabledNodeBlocking() ([]nodeBlockingRecord, error) { func (s *Store) ListEnabledNodeBlocking() ([]NodeBlockingRecord, error) {
var rows []nodeBlockingRecord var rows []NodeBlockingRecord
return rows, s.db.Where("enabled = ?", true).Find(&rows).Error return rows, s.db.Where("enabled = ?", true).Find(&rows).Error
} }
func (s *store) CreateNodeBlocking(nodeID string, nodeNum *int64, reason string, enabled bool) (*nodeBlockingRecord, error) { func (s *Store) CreateNodeBlocking(nodeID string, nodeNum *int64, reason string, enabled bool) (*NodeBlockingRecord, error) {
nodeID = strings.TrimSpace(nodeID) nodeID = strings.TrimSpace(nodeID)
if nodeID == "" { if nodeID == "" {
return nil, fmt.Errorf("node id is required") return nil, fmt.Errorf("node id is required")
@@ -43,14 +43,14 @@ func (s *store) CreateNodeBlocking(nodeID string, nodeNum *int64, reason string,
if err := s.ensureNodeBlockingUnique(0, nodeID); err != nil { if err := s.ensureNodeBlockingUnique(0, nodeID); err != nil {
return nil, err return nil, err
} }
row := nodeBlockingRecord{NodeID: nodeID, NodeNum: nodeNum, Reason: strings.TrimSpace(reason), Enabled: enabled} row := NodeBlockingRecord{NodeID: nodeID, NodeNum: nodeNum, Reason: strings.TrimSpace(reason), Enabled: enabled}
if err := s.db.Create(&row).Error; err != nil { if err := s.db.Create(&row).Error; err != nil {
return nil, err return nil, err
} }
return &row, nil return &row, nil
} }
func (s *store) UpdateNodeBlocking(id uint64, nodeID string, nodeNum *int64, reason string, enabled bool) (*nodeBlockingRecord, error) { func (s *Store) UpdateNodeBlocking(id uint64, nodeID string, nodeNum *int64, reason string, enabled bool) (*NodeBlockingRecord, error) {
if id == 0 { if id == 0 {
return nil, fmt.Errorf("blocking rule id is required") return nil, fmt.Errorf("blocking rule id is required")
} }
@@ -65,14 +65,14 @@ func (s *store) UpdateNodeBlocking(id uint64, nodeID string, nodeNum *int64, rea
return nil, err return nil, err
} }
updates := map[string]any{"node_id": nodeID, "node_num": nodeNum, "reason": strings.TrimSpace(reason), "enabled": enabled, "updated_at": time.Now()} updates := map[string]any{"node_id": nodeID, "node_num": nodeNum, "reason": strings.TrimSpace(reason), "enabled": enabled, "updated_at": time.Now()}
if err := s.db.Model(&nodeBlockingRecord{}).Where("id = ?", id).Updates(updates).Error; err != nil { if err := s.db.Model(&NodeBlockingRecord{}).Where("id = ?", id).Updates(updates).Error; err != nil {
return nil, err return nil, err
} }
return s.getNodeBlockingByID(id) return s.getNodeBlockingByID(id)
} }
func (s *store) DeleteNodeBlocking(id uint64) error { func (s *Store) DeleteNodeBlocking(id uint64) error {
result := s.db.Where("id = ?", id).Delete(&nodeBlockingRecord{}) result := s.db.Where("id = ?", id).Delete(&NodeBlockingRecord{})
if result.Error != nil { if result.Error != nil {
return result.Error return result.Error
} }
@@ -82,10 +82,10 @@ func (s *store) DeleteNodeBlocking(id uint64) error {
return nil return nil
} }
func (s *store) ListIPBlocking(opts listOptions) ([]ipBlockingRecord, error) { func (s *Store) ListIPBlocking(opts ListOptions) ([]IPBlockingRecord, error) {
opts = normalizeListOptions(opts) opts = NormalizeListOptions(opts)
var rows []ipBlockingRecord var rows []IPBlockingRecord
q := s.db.Model(&ipBlockingRecord{}). q := s.db.Model(&IPBlockingRecord{}).
Order("updated_at DESC"). Order("updated_at DESC").
Order("id DESC"). Order("id DESC").
Limit(opts.Limit). Limit(opts.Limit).
@@ -93,17 +93,17 @@ func (s *store) ListIPBlocking(opts listOptions) ([]ipBlockingRecord, error) {
return rows, q.Find(&rows).Error return rows, q.Find(&rows).Error
} }
func (s *store) CountIPBlocking(opts listOptions) (int64, error) { func (s *Store) CountIPBlocking(opts ListOptions) (int64, error) {
var total int64 var total int64
return total, s.db.Model(&ipBlockingRecord{}).Count(&total).Error return total, s.db.Model(&IPBlockingRecord{}).Count(&total).Error
} }
func (s *store) ListEnabledIPBlocking() ([]ipBlockingRecord, error) { func (s *Store) ListEnabledIPBlocking() ([]IPBlockingRecord, error) {
var rows []ipBlockingRecord var rows []IPBlockingRecord
return rows, s.db.Where("enabled = ?", true).Find(&rows).Error return rows, s.db.Where("enabled = ?", true).Find(&rows).Error
} }
func (s *store) CreateIPBlocking(ipValue string, reason string, enabled bool) (*ipBlockingRecord, error) { func (s *Store) CreateIPBlocking(ipValue string, reason string, enabled bool) (*IPBlockingRecord, error) {
value, err := normalizeIPBlockingValue(ipValue) value, err := normalizeIPBlockingValue(ipValue)
if err != nil { if err != nil {
return nil, err return nil, err
@@ -111,14 +111,14 @@ func (s *store) CreateIPBlocking(ipValue string, reason string, enabled bool) (*
if err := s.ensureIPBlockingUnique(0, value); err != nil { if err := s.ensureIPBlockingUnique(0, value); err != nil {
return nil, err return nil, err
} }
row := ipBlockingRecord{IPValue: value, Reason: strings.TrimSpace(reason), Enabled: enabled} row := IPBlockingRecord{IPValue: value, Reason: strings.TrimSpace(reason), Enabled: enabled}
if err := s.db.Create(&row).Error; err != nil { if err := s.db.Create(&row).Error; err != nil {
return nil, err return nil, err
} }
return &row, nil return &row, nil
} }
func (s *store) UpdateIPBlocking(id uint64, ipValue string, reason string, enabled bool) (*ipBlockingRecord, error) { func (s *Store) UpdateIPBlocking(id uint64, ipValue string, reason string, enabled bool) (*IPBlockingRecord, error) {
if id == 0 { if id == 0 {
return nil, fmt.Errorf("blocking rule id is required") return nil, fmt.Errorf("blocking rule id is required")
} }
@@ -133,14 +133,14 @@ func (s *store) UpdateIPBlocking(id uint64, ipValue string, reason string, enabl
return nil, err return nil, err
} }
updates := map[string]any{"ip_value": value, "reason": strings.TrimSpace(reason), "enabled": enabled, "updated_at": time.Now()} updates := map[string]any{"ip_value": value, "reason": strings.TrimSpace(reason), "enabled": enabled, "updated_at": time.Now()}
if err := s.db.Model(&ipBlockingRecord{}).Where("id = ?", id).Updates(updates).Error; err != nil { if err := s.db.Model(&IPBlockingRecord{}).Where("id = ?", id).Updates(updates).Error; err != nil {
return nil, err return nil, err
} }
return s.getIPBlockingByID(id) return s.getIPBlockingByID(id)
} }
func (s *store) DeleteIPBlocking(id uint64) error { func (s *Store) DeleteIPBlocking(id uint64) error {
result := s.db.Where("id = ?", id).Delete(&ipBlockingRecord{}) result := s.db.Where("id = ?", id).Delete(&IPBlockingRecord{})
if result.Error != nil { if result.Error != nil {
return result.Error return result.Error
} }
@@ -150,10 +150,10 @@ func (s *store) DeleteIPBlocking(id uint64) error {
return nil return nil
} }
func (s *store) ListForbiddenWordBlocking(opts listOptions) ([]forbiddenWordBlockingRecord, error) { func (s *Store) ListForbiddenWordBlocking(opts ListOptions) ([]ForbiddenWordBlockingRecord, error) {
opts = normalizeListOptions(opts) opts = NormalizeListOptions(opts)
var rows []forbiddenWordBlockingRecord var rows []ForbiddenWordBlockingRecord
q := s.db.Model(&forbiddenWordBlockingRecord{}). q := s.db.Model(&ForbiddenWordBlockingRecord{}).
Order("updated_at DESC"). Order("updated_at DESC").
Order("id DESC"). Order("id DESC").
Limit(opts.Limit). Limit(opts.Limit).
@@ -161,17 +161,17 @@ func (s *store) ListForbiddenWordBlocking(opts listOptions) ([]forbiddenWordBloc
return rows, q.Find(&rows).Error return rows, q.Find(&rows).Error
} }
func (s *store) CountForbiddenWordBlocking(opts listOptions) (int64, error) { func (s *Store) CountForbiddenWordBlocking(opts ListOptions) (int64, error) {
var total int64 var total int64
return total, s.db.Model(&forbiddenWordBlockingRecord{}).Count(&total).Error return total, s.db.Model(&ForbiddenWordBlockingRecord{}).Count(&total).Error
} }
func (s *store) ListEnabledForbiddenWordBlocking() ([]forbiddenWordBlockingRecord, error) { func (s *Store) ListEnabledForbiddenWordBlocking() ([]ForbiddenWordBlockingRecord, error) {
var rows []forbiddenWordBlockingRecord var rows []ForbiddenWordBlockingRecord
return rows, s.db.Where("enabled = ?", true).Find(&rows).Error return rows, s.db.Where("enabled = ?", true).Find(&rows).Error
} }
func (s *store) CreateForbiddenWordBlocking(word, matchType string, caseSensitive bool, reason string, enabled bool) (*forbiddenWordBlockingRecord, error) { func (s *Store) CreateForbiddenWordBlocking(word, matchType string, caseSensitive bool, reason string, enabled bool) (*ForbiddenWordBlockingRecord, error) {
word = strings.TrimSpace(word) word = strings.TrimSpace(word)
if word == "" { if word == "" {
return nil, fmt.Errorf("forbidden word is required") return nil, fmt.Errorf("forbidden word is required")
@@ -183,14 +183,14 @@ func (s *store) CreateForbiddenWordBlocking(word, matchType string, caseSensitiv
if err := s.ensureForbiddenWordBlockingUnique(0, word); err != nil { if err := s.ensureForbiddenWordBlockingUnique(0, word); err != nil {
return nil, err return nil, err
} }
row := forbiddenWordBlockingRecord{Word: word, MatchType: matchType, CaseSensitive: caseSensitive, Reason: strings.TrimSpace(reason), Enabled: enabled} row := ForbiddenWordBlockingRecord{Word: word, MatchType: matchType, CaseSensitive: caseSensitive, Reason: strings.TrimSpace(reason), Enabled: enabled}
if err := s.db.Create(&row).Error; err != nil { if err := s.db.Create(&row).Error; err != nil {
return nil, err return nil, err
} }
return &row, nil return &row, nil
} }
func (s *store) UpdateForbiddenWordBlocking(id uint64, word, matchType string, caseSensitive bool, reason string, enabled bool) (*forbiddenWordBlockingRecord, error) { func (s *Store) UpdateForbiddenWordBlocking(id uint64, word, matchType string, caseSensitive bool, reason string, enabled bool) (*ForbiddenWordBlockingRecord, error) {
if id == 0 { if id == 0 {
return nil, fmt.Errorf("blocking rule id is required") return nil, fmt.Errorf("blocking rule id is required")
} }
@@ -209,14 +209,14 @@ func (s *store) UpdateForbiddenWordBlocking(id uint64, word, matchType string, c
return nil, err return nil, err
} }
updates := map[string]any{"word": word, "match_type": matchType, "case_sensitive": caseSensitive, "reason": strings.TrimSpace(reason), "enabled": enabled, "updated_at": time.Now()} updates := map[string]any{"word": word, "match_type": matchType, "case_sensitive": caseSensitive, "reason": strings.TrimSpace(reason), "enabled": enabled, "updated_at": time.Now()}
if err := s.db.Model(&forbiddenWordBlockingRecord{}).Where("id = ?", id).Updates(updates).Error; err != nil { if err := s.db.Model(&ForbiddenWordBlockingRecord{}).Where("id = ?", id).Updates(updates).Error; err != nil {
return nil, err return nil, err
} }
return s.getForbiddenWordBlockingByID(id) return s.getForbiddenWordBlockingByID(id)
} }
func (s *store) DeleteForbiddenWordBlocking(id uint64) error { func (s *Store) DeleteForbiddenWordBlocking(id uint64) error {
result := s.db.Where("id = ?", id).Delete(&forbiddenWordBlockingRecord{}) result := s.db.Where("id = ?", id).Delete(&ForbiddenWordBlockingRecord{})
if result.Error != nil { if result.Error != nil {
return result.Error return result.Error
} }
@@ -226,39 +226,39 @@ func (s *store) DeleteForbiddenWordBlocking(id uint64) error {
return nil return nil
} }
func (s *store) getNodeBlockingByID(id uint64) (*nodeBlockingRecord, error) { func (s *Store) getNodeBlockingByID(id uint64) (*NodeBlockingRecord, error) {
var row nodeBlockingRecord var row NodeBlockingRecord
if err := s.db.Where("id = ?", id).Take(&row).Error; err != nil { if err := s.db.Where("id = ?", id).Take(&row).Error; err != nil {
return nil, err return nil, err
} }
return &row, nil return &row, nil
} }
func (s *store) getIPBlockingByID(id uint64) (*ipBlockingRecord, error) { func (s *Store) getIPBlockingByID(id uint64) (*IPBlockingRecord, error) {
var row ipBlockingRecord var row IPBlockingRecord
if err := s.db.Where("id = ?", id).Take(&row).Error; err != nil { if err := s.db.Where("id = ?", id).Take(&row).Error; err != nil {
return nil, err return nil, err
} }
return &row, nil return &row, nil
} }
func (s *store) getForbiddenWordBlockingByID(id uint64) (*forbiddenWordBlockingRecord, error) { func (s *Store) getForbiddenWordBlockingByID(id uint64) (*ForbiddenWordBlockingRecord, error) {
var row forbiddenWordBlockingRecord var row ForbiddenWordBlockingRecord
if err := s.db.Where("id = ?", id).Take(&row).Error; err != nil { if err := s.db.Where("id = ?", id).Take(&row).Error; err != nil {
return nil, err return nil, err
} }
return &row, nil return &row, nil
} }
func (s *store) ensureNodeBlockingUnique(id uint64, nodeID string) error { func (s *Store) ensureNodeBlockingUnique(id uint64, nodeID string) error {
var existing nodeBlockingRecord var existing NodeBlockingRecord
q := s.db.Where("node_id = ?", nodeID) q := s.db.Where("node_id = ?", nodeID)
if id != 0 { if id != 0 {
q = q.Where("id <> ?", id) q = q.Where("id <> ?", id)
} }
err := q.Take(&existing).Error err := q.Take(&existing).Error
if err == nil { if err == nil {
return errBlockingAlreadyExists return ErrBlockingAlreadyExists
} }
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
return nil return nil
@@ -266,15 +266,15 @@ func (s *store) ensureNodeBlockingUnique(id uint64, nodeID string) error {
return err return err
} }
func (s *store) ensureIPBlockingUnique(id uint64, ipValue string) error { func (s *Store) ensureIPBlockingUnique(id uint64, ipValue string) error {
var existing ipBlockingRecord var existing IPBlockingRecord
q := s.db.Where("ip_value = ?", ipValue) q := s.db.Where("ip_value = ?", ipValue)
if id != 0 { if id != 0 {
q = q.Where("id <> ?", id) q = q.Where("id <> ?", id)
} }
err := q.Take(&existing).Error err := q.Take(&existing).Error
if err == nil { if err == nil {
return errBlockingAlreadyExists return ErrBlockingAlreadyExists
} }
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
return nil return nil
@@ -282,15 +282,15 @@ func (s *store) ensureIPBlockingUnique(id uint64, ipValue string) error {
return err return err
} }
func (s *store) ensureForbiddenWordBlockingUnique(id uint64, word string) error { func (s *Store) ensureForbiddenWordBlockingUnique(id uint64, word string) error {
var existing forbiddenWordBlockingRecord var existing ForbiddenWordBlockingRecord
q := s.db.Where("word = ?", word) q := s.db.Where("word = ?", word)
if id != 0 { if id != 0 {
q = q.Where("id <> ?", id) q = q.Where("id <> ?", id)
} }
err := q.Take(&existing).Error err := q.Take(&existing).Error
if err == nil { if err == nil {
return errBlockingAlreadyExists return ErrBlockingAlreadyExists
} }
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
return nil return nil
@@ -316,9 +316,9 @@ func normalizeIPBlockingValue(value string) (string, error) {
func normalizeForbiddenWordMatchType(matchType string) (string, error) { func normalizeForbiddenWordMatchType(matchType string) (string, error) {
matchType = strings.TrimSpace(matchType) matchType = strings.TrimSpace(matchType)
if matchType == "" { if matchType == "" {
return forbiddenWordMatchContains, nil return ForbiddenWordMatchContains, nil
} }
if matchType != forbiddenWordMatchContains { if matchType != ForbiddenWordMatchContains {
return "", fmt.Errorf("unsupported forbidden word match type") return "", fmt.Errorf("unsupported forbidden word match type")
} }
return matchType, nil return matchType, nil
@@ -1,4 +1,4 @@
package main package store
import ( import (
"errors" "errors"
@@ -20,8 +20,8 @@ func TestNodeBlockingCRUD(t *testing.T) {
t.Fatalf("created node rule = %+v, want normalized fields", rule) t.Fatalf("created node rule = %+v, want normalized fields", rule)
} }
if _, err := st.CreateNodeBlocking("!12345678", nil, "duplicate", true); !errors.Is(err, errBlockingAlreadyExists) { if _, err := st.CreateNodeBlocking("!12345678", nil, "duplicate", true); !errors.Is(err, ErrBlockingAlreadyExists) {
t.Fatalf("duplicate CreateNodeBlocking() error = %v, want errBlockingAlreadyExists", err) t.Fatalf("duplicate CreateNodeBlocking() error = %v, want ErrBlockingAlreadyExists", err)
} }
updatedNum := int64(7) updatedNum := int64(7)
@@ -33,14 +33,14 @@ func TestNodeBlockingCRUD(t *testing.T) {
t.Fatalf("updated node rule = %+v, want updated fields", updated) t.Fatalf("updated node rule = %+v, want updated fields", updated)
} }
rows, err := st.ListNodeBlocking(listOptions{}) rows, err := st.ListNodeBlocking(ListOptions{})
if err != nil { if err != nil {
t.Fatalf("ListNodeBlocking() error = %v", err) t.Fatalf("ListNodeBlocking() error = %v", err)
} }
if len(rows) != 1 || rows[0].ID != rule.ID { if len(rows) != 1 || rows[0].ID != rule.ID {
t.Fatalf("ListNodeBlocking() = %+v, want one updated rule", rows) t.Fatalf("ListNodeBlocking() = %+v, want one updated rule", rows)
} }
total, err := st.CountNodeBlocking(listOptions{}) total, err := st.CountNodeBlocking(ListOptions{})
if err != nil || total != 1 { if err != nil || total != 1 {
t.Fatalf("CountNodeBlocking() = %d, %v, want 1, nil", total, err) t.Fatalf("CountNodeBlocking() = %d, %v, want 1, nil", total, err)
} }
@@ -85,8 +85,8 @@ func TestIPBlockingCRUDAndValidation(t *testing.T) {
t.Fatalf("cidr IPValue = %q, want 192.168.1.0/24", cidr.IPValue) t.Fatalf("cidr IPValue = %q, want 192.168.1.0/24", cidr.IPValue)
} }
if _, err := st.CreateIPBlocking("127.0.0.1", "duplicate", true); !errors.Is(err, errBlockingAlreadyExists) { if _, err := st.CreateIPBlocking("127.0.0.1", "duplicate", true); !errors.Is(err, ErrBlockingAlreadyExists) {
t.Fatalf("duplicate CreateIPBlocking() error = %v, want errBlockingAlreadyExists", err) t.Fatalf("duplicate CreateIPBlocking() error = %v, want ErrBlockingAlreadyExists", err)
} }
if _, err := st.CreateIPBlocking("not-an-ip", "invalid", true); err == nil { if _, err := st.CreateIPBlocking("not-an-ip", "invalid", true); err == nil {
t.Fatal("CreateIPBlocking(invalid) error = nil, want error") t.Fatal("CreateIPBlocking(invalid) error = nil, want error")
@@ -100,14 +100,14 @@ func TestIPBlockingCRUDAndValidation(t *testing.T) {
t.Fatalf("updated ip rule = %+v, want updated fields", updated) t.Fatalf("updated ip rule = %+v, want updated fields", updated)
} }
rows, err := st.ListIPBlocking(listOptions{}) rows, err := st.ListIPBlocking(ListOptions{})
if err != nil { if err != nil {
t.Fatalf("ListIPBlocking() error = %v", err) t.Fatalf("ListIPBlocking() error = %v", err)
} }
if len(rows) != 2 { if len(rows) != 2 {
t.Fatalf("ListIPBlocking() length = %d, want 2", len(rows)) t.Fatalf("ListIPBlocking() length = %d, want 2", len(rows))
} }
total, err := st.CountIPBlocking(listOptions{}) total, err := st.CountIPBlocking(ListOptions{})
if err != nil || total != 2 { if err != nil || total != 2 {
t.Fatalf("CountIPBlocking() = %d, %v, want 2, nil", total, err) t.Fatalf("CountIPBlocking() = %d, %v, want 2, nil", total, err)
} }
@@ -164,12 +164,12 @@ func TestForbiddenWordBlockingCRUDAndValidation(t *testing.T) {
if err != nil { if err != nil {
t.Fatalf("CreateForbiddenWordBlocking() error = %v", err) t.Fatalf("CreateForbiddenWordBlocking() error = %v", err)
} }
if rule.Word != "spam" || rule.MatchType != forbiddenWordMatchContains || rule.CaseSensitive || rule.Reason != "junk" || !rule.Enabled { if rule.Word != "spam" || rule.MatchType != ForbiddenWordMatchContains || rule.CaseSensitive || rule.Reason != "junk" || !rule.Enabled {
t.Fatalf("created word rule = %+v, want normalized fields", rule) t.Fatalf("created word rule = %+v, want normalized fields", rule)
} }
if _, err := st.CreateForbiddenWordBlocking("spam", "contains", false, "duplicate", true); !errors.Is(err, errBlockingAlreadyExists) { if _, err := st.CreateForbiddenWordBlocking("spam", "contains", false, "duplicate", true); !errors.Is(err, ErrBlockingAlreadyExists) {
t.Fatalf("duplicate CreateForbiddenWordBlocking() error = %v, want errBlockingAlreadyExists", err) t.Fatalf("duplicate CreateForbiddenWordBlocking() error = %v, want ErrBlockingAlreadyExists", err)
} }
if _, err := st.CreateForbiddenWordBlocking(" ", "contains", false, "empty", true); err == nil { if _, err := st.CreateForbiddenWordBlocking(" ", "contains", false, "empty", true); err == nil {
t.Fatal("CreateForbiddenWordBlocking(empty) error = nil, want error") t.Fatal("CreateForbiddenWordBlocking(empty) error = nil, want error")
@@ -186,14 +186,14 @@ func TestForbiddenWordBlockingCRUDAndValidation(t *testing.T) {
t.Fatalf("updated word rule = %+v, want updated fields", updated) t.Fatalf("updated word rule = %+v, want updated fields", updated)
} }
rows, err := st.ListForbiddenWordBlocking(listOptions{}) rows, err := st.ListForbiddenWordBlocking(ListOptions{})
if err != nil { if err != nil {
t.Fatalf("ListForbiddenWordBlocking() error = %v", err) t.Fatalf("ListForbiddenWordBlocking() error = %v", err)
} }
if len(rows) != 1 || rows[0].ID != rule.ID { if len(rows) != 1 || rows[0].ID != rule.ID {
t.Fatalf("ListForbiddenWordBlocking() = %+v, want one updated rule", rows) t.Fatalf("ListForbiddenWordBlocking() = %+v, want one updated rule", rows)
} }
total, err := st.CountForbiddenWordBlocking(listOptions{}) total, err := st.CountForbiddenWordBlocking(ListOptions{})
if err != nil || total != 1 { if err != nil || total != 1 {
t.Fatalf("CountForbiddenWordBlocking() = %d, %v, want 1, nil", total, err) t.Fatalf("CountForbiddenWordBlocking() = %d, %v, want 1, nil", total, err)
} }
@@ -1,4 +1,4 @@
package main package store
import ( import (
"encoding/json" "encoding/json"
@@ -10,17 +10,17 @@ import (
"gorm.io/gorm" "gorm.io/gorm"
) )
type botDirectMessageListOptions struct { type BotDirectMessageListOptions struct {
listOptions ListOptions
BotID uint64 BotID uint64
PeerNodeNum int64 PeerNodeNum int64
Direction string Direction string
} }
// InsertBotDirectMessage 把一条机器人 DM(出向或入向)写入 bot_direct_messages 表。 // InsertBotDirectMessage 把一条机器人 DM(出向或入向)写入 bot_direct_messages 表。
func (s *store) InsertBotDirectMessage(row *botDirectMessageRecord) error { func (s *Store) InsertBotDirectMessage(row *BotDirectMessageRecord) error {
if s == nil || s.db == nil { if s == nil || s.db == nil {
return fmt.Errorf("store is not configured") return fmt.Errorf("Store is not configured")
} }
if row == nil { if row == nil {
return fmt.Errorf("bot direct message is required") return fmt.Errorf("bot direct message is required")
@@ -32,9 +32,9 @@ func (s *store) InsertBotDirectMessage(row *botDirectMessageRecord) error {
} }
// UpdateBotDirectMessageStatus 更新一条出向 DM 的发送状态(pending → published/failed)。 // UpdateBotDirectMessageStatus 更新一条出向 DM 的发送状态(pending → published/failed)。
func (s *store) UpdateBotDirectMessageStatus(id uint64, status, errText string, publishedAt *time.Time) error { func (s *Store) UpdateBotDirectMessageStatus(id uint64, status, errText string, publishedAt *time.Time) error {
if s == nil || s.db == nil { if s == nil || s.db == nil {
return fmt.Errorf("store is not configured") return fmt.Errorf("Store is not configured")
} }
if id == 0 { if id == 0 {
return fmt.Errorf("bot direct message id is required") return fmt.Errorf("bot direct message id is required")
@@ -44,7 +44,7 @@ func (s *store) UpdateBotDirectMessageStatus(id uint64, status, errText string,
"error": strings.TrimSpace(errText), "error": strings.TrimSpace(errText),
"published_at": publishedAt, "published_at": publishedAt,
} }
result := s.db.Model(&botDirectMessageRecord{}).Where("id = ?", id).Updates(updates) result := s.db.Model(&BotDirectMessageRecord{}).Where("id = ?", id).Updates(updates)
if result.Error != nil { if result.Error != nil {
return result.Error return result.Error
} }
@@ -55,9 +55,9 @@ func (s *store) UpdateBotDirectMessageStatus(id uint64, status, errText string,
} }
// ListBotDirectMessagesByConversation 按 (bot, peer) 反序拉取 DM 历史,给 /admin/bot/direct 页面。 // ListBotDirectMessagesByConversation 按 (bot, peer) 反序拉取 DM 历史,给 /admin/bot/direct 页面。
func (s *store) ListBotDirectMessagesByConversation(opts botDirectMessageListOptions) ([]botDirectMessageRecord, error) { func (s *Store) ListBotDirectMessagesByConversation(opts BotDirectMessageListOptions) ([]BotDirectMessageRecord, error) {
if s == nil || s.db == nil { if s == nil || s.db == nil {
return nil, fmt.Errorf("store is not configured") return nil, fmt.Errorf("Store is not configured")
} }
if opts.BotID == 0 { if opts.BotID == 0 {
return nil, fmt.Errorf("bot id is required") return nil, fmt.Errorf("bot id is required")
@@ -65,9 +65,9 @@ func (s *store) ListBotDirectMessagesByConversation(opts botDirectMessageListOpt
if opts.PeerNodeNum == 0 { if opts.PeerNodeNum == 0 {
return nil, fmt.Errorf("peer node num is required") return nil, fmt.Errorf("peer node num is required")
} }
opts.listOptions = normalizeListOptions(opts.listOptions) opts.ListOptions = NormalizeListOptions(opts.ListOptions)
var rows []botDirectMessageRecord var rows []BotDirectMessageRecord
q := s.db.Model(&botDirectMessageRecord{}). q := s.db.Model(&BotDirectMessageRecord{}).
Where("bot_id = ? AND peer_node_num = ?", opts.BotID, opts.PeerNodeNum). Where("bot_id = ? AND peer_node_num = ?", opts.BotID, opts.PeerNodeNum).
Order("created_at DESC"). Order("created_at DESC").
Order("id DESC"). Order("id DESC").
@@ -86,15 +86,15 @@ func (s *store) ListBotDirectMessagesByConversation(opts botDirectMessageListOpt
} }
// CountBotDirectMessagesByConversation 返回会话总条数(前端无限滚动可用,可选)。 // CountBotDirectMessagesByConversation 返回会话总条数(前端无限滚动可用,可选)。
func (s *store) CountBotDirectMessagesByConversation(opts botDirectMessageListOptions) (int64, error) { func (s *Store) CountBotDirectMessagesByConversation(opts BotDirectMessageListOptions) (int64, error) {
if s == nil || s.db == nil { if s == nil || s.db == nil {
return 0, fmt.Errorf("store is not configured") return 0, fmt.Errorf("Store is not configured")
} }
if opts.BotID == 0 || opts.PeerNodeNum == 0 { if opts.BotID == 0 || opts.PeerNodeNum == 0 {
return 0, fmt.Errorf("bot id and peer node num are required") return 0, fmt.Errorf("bot id and peer node num are required")
} }
var total int64 var total int64
q := s.db.Model(&botDirectMessageRecord{}). q := s.db.Model(&BotDirectMessageRecord{}).
Where("bot_id = ? AND peer_node_num = ?", opts.BotID, opts.PeerNodeNum) Where("bot_id = ? AND peer_node_num = ?", opts.BotID, opts.PeerNodeNum)
if opts.Direction != "" { if opts.Direction != "" {
q = q.Where("direction = ?", opts.Direction) q = q.Where("direction = ?", opts.Direction)
@@ -110,9 +110,9 @@ func (s *store) CountBotDirectMessagesByConversation(opts botDirectMessageListOp
// FindBotForIncomingPKIPacket 在 bot_direct_messages 写入路径上判断接收方是否为受管 bot。 // FindBotForIncomingPKIPacket 在 bot_direct_messages 写入路径上判断接收方是否为受管 bot。
// 返回的 bot 用于填充 BotID/BotNodeID/BotNodeNum;不命中时返回 ErrRecordNotFound。 // 返回的 bot 用于填充 BotID/BotNodeID/BotNodeNum;不命中时返回 ErrRecordNotFound。
func (s *store) FindBotForIncomingPKIPacket(toNodeNum int64) (*botNodeRecord, error) { func (s *Store) FindBotForIncomingPKIPacket(toNodeNum int64) (*BotNodeRecord, error) {
if s == nil || s.db == nil { if s == nil || s.db == nil {
return nil, fmt.Errorf("store is not configured") return nil, fmt.Errorf("Store is not configured")
} }
bot, err := s.GetBotNodeByNodeNum(toNodeNum) bot, err := s.GetBotNodeByNodeNum(toNodeNum)
if err != nil { if err != nil {
@@ -124,10 +124,10 @@ func (s *store) FindBotForIncomingPKIPacket(toNodeNum int64) (*botNodeRecord, er
return bot, nil return bot, nil
} }
// botDirectConversation 是 /admin/bot/direct 侧边栏需要的会话摘要。 // BotDirectConversation 是 /admin/bot/direct 侧边栏需要的会话摘要。
// LastMessageAt / LastText / LastDirection 描述会话最后一条消息,便于按时间排序与预览; // LastMessageAt / LastText / LastDirection 描述会话最后一条消息,便于按时间排序与预览;
// UnreadCount 仅对 inbound 计数(即未读消息数)。 // UnreadCount 仅对 inbound 计数(即未读消息数)。
type botDirectConversation struct { type BotDirectConversation struct {
BotID uint64 `gorm:"column:bot_id"` BotID uint64 `gorm:"column:bot_id"`
PeerNodeID string `gorm:"column:peer_node_id"` PeerNodeID string `gorm:"column:peer_node_id"`
PeerNodeNum int64 `gorm:"column:peer_node_num"` PeerNodeNum int64 `gorm:"column:peer_node_num"`
@@ -139,21 +139,21 @@ type botDirectConversation struct {
} }
// ListBotDirectConversations 聚合给定 bot 下的所有 (peer) 会话,返回最后一条消息及未读数。 // ListBotDirectConversations 聚合给定 bot 下的所有 (peer) 会话,返回最后一条消息及未读数。
// 按最后一条消息时间倒序(最新会话排前面)。limit/offset 走 listOptions。 // 按最后一条消息时间倒序(最新会话排前面)。limit/offset 走 ListOptions。
func (s *store) ListBotDirectConversations(botID uint64, opts listOptions) ([]botDirectConversation, error) { func (s *Store) ListBotDirectConversations(botID uint64, opts ListOptions) ([]BotDirectConversation, error) {
if s == nil || s.db == nil { if s == nil || s.db == nil {
return nil, fmt.Errorf("store is not configured") return nil, fmt.Errorf("Store is not configured")
} }
if botID == 0 { if botID == 0 {
return nil, fmt.Errorf("bot id is required") return nil, fmt.Errorf("bot id is required")
} }
opts = normalizeListOptions(opts) opts = NormalizeListOptions(opts)
var rows []botDirectConversation var rows []BotDirectConversation
// 先把每对会话的最后一条消息 ID 取出来,再把这条消息的元数据 join 回去; // 先把每对会话的最后一条消息 ID 取出来,再把这条消息的元数据 join 回去;
// 同时聚合 unread_countinbound 且 read_at IS NULL)和 total_count。 // 同时聚合 unread_countinbound 且 read_at IS NULL)和 total_count。
// 这样的两步 join 避免在 GROUP BY 后引用非聚合列(MySQL 严格模式 / SQLite 兼容)。 // 这样的两步 join 避免在 GROUP BY 后引用非聚合列(MySQL 严格模式 / SQLite 兼容)。
subLast := s.db.Model(&botDirectMessageRecord{}). subLast := s.db.Model(&BotDirectMessageRecord{}).
Select("bot_id, peer_node_id, peer_node_num, MAX(id) AS last_id, COUNT(*) AS total_count, SUM(CASE WHEN direction = ? AND read_at IS NULL THEN 1 ELSE 0 END) AS unread_count", botDirectMessageDirectionInbound). Select("bot_id, peer_node_id, peer_node_num, MAX(id) AS last_id, COUNT(*) AS total_count, SUM(CASE WHEN direction = ? AND read_at IS NULL THEN 1 ELSE 0 END) AS unread_count", BotDirectMessageDirectionInbound).
Where("bot_id = ?", botID). Where("bot_id = ?", botID).
Group("bot_id, peer_node_id, peer_node_num") Group("bot_id, peer_node_id, peer_node_num")
q := s.db.Table("(?) AS agg", subLast). q := s.db.Table("(?) AS agg", subLast).
@@ -167,16 +167,16 @@ func (s *store) ListBotDirectConversations(botID uint64, opts listOptions) ([]bo
} }
// MarkBotDirectMessagesRead 把 (bot, peer) 下未读的 inbound 消息全部标记为已读,返回更新行数。 // MarkBotDirectMessagesRead 把 (bot, peer) 下未读的 inbound 消息全部标记为已读,返回更新行数。
func (s *store) MarkBotDirectMessagesRead(botID uint64, peerNodeNum int64) (int64, error) { func (s *Store) MarkBotDirectMessagesRead(botID uint64, peerNodeNum int64) (int64, error) {
if s == nil || s.db == nil { if s == nil || s.db == nil {
return 0, fmt.Errorf("store is not configured") return 0, fmt.Errorf("Store is not configured")
} }
if botID == 0 || peerNodeNum == 0 { if botID == 0 || peerNodeNum == 0 {
return 0, fmt.Errorf("bot id and peer node num are required") return 0, fmt.Errorf("bot id and peer node num are required")
} }
now := time.Now() now := time.Now()
result := s.db.Model(&botDirectMessageRecord{}). result := s.db.Model(&BotDirectMessageRecord{}).
Where("bot_id = ? AND peer_node_num = ? AND direction = ? AND read_at IS NULL", botID, peerNodeNum, botDirectMessageDirectionInbound). Where("bot_id = ? AND peer_node_num = ? AND direction = ? AND read_at IS NULL", botID, peerNodeNum, BotDirectMessageDirectionInbound).
Update("read_at", &now) Update("read_at", &now)
if result.Error != nil { if result.Error != nil {
return 0, result.Error return 0, result.Error
@@ -185,16 +185,16 @@ func (s *store) MarkBotDirectMessagesRead(botID uint64, peerNodeNum int64) (int6
} }
// CountBotDirectUnread 返回某个 bot 全部未读 inbound 消息总数(用于头部小红点)。 // CountBotDirectUnread 返回某个 bot 全部未读 inbound 消息总数(用于头部小红点)。
func (s *store) CountBotDirectUnread(botID uint64) (int64, error) { func (s *Store) CountBotDirectUnread(botID uint64) (int64, error) {
if s == nil || s.db == nil { if s == nil || s.db == nil {
return 0, fmt.Errorf("store is not configured") return 0, fmt.Errorf("Store is not configured")
} }
if botID == 0 { if botID == 0 {
return 0, fmt.Errorf("bot id is required") return 0, fmt.Errorf("bot id is required")
} }
var total int64 var total int64
err := s.db.Model(&botDirectMessageRecord{}). err := s.db.Model(&BotDirectMessageRecord{}).
Where("bot_id = ? AND direction = ? AND read_at IS NULL", botID, botDirectMessageDirectionInbound). Where("bot_id = ? AND direction = ? AND read_at IS NULL", botID, BotDirectMessageDirectionInbound).
Count(&total).Error Count(&total).Error
return total, err return total, err
} }
@@ -202,7 +202,7 @@ func (s *store) CountBotDirectUnread(botID uint64) (int64, error) {
// isInboundBotDirectMessage 判断 record 是否是“PKI 加密、发往受管 bot”的入向 DM。 // isInboundBotDirectMessage 判断 record 是否是“PKI 加密、发往受管 bot”的入向 DM。
// 仅在 type=text_message、pki_encrypted=true、packet_to_num 命中受管 bot 时返回 true。 // 仅在 type=text_message、pki_encrypted=true、packet_to_num 命中受管 bot 时返回 true。
// 任何步骤失败都返回 false,让记录回落到 text_message 表(与之前行为兼容)。 // 任何步骤失败都返回 false,让记录回落到 text_message 表(与之前行为兼容)。
func isInboundBotDirectMessage(s *store, record map[string]any) bool { func isInboundBotDirectMessage(s *Store, record map[string]any) bool {
if s == nil || record == nil { if s == nil || record == nil {
return false return false
} }
@@ -221,10 +221,10 @@ func isInboundBotDirectMessage(s *store, record map[string]any) bool {
} }
// insertInboundBotDirectMessage 把一条入向 PKI DM 转写入 bot_direct_messages 表。 // insertInboundBotDirectMessage 把一条入向 PKI DM 转写入 bot_direct_messages 表。
// 失败时返回错误,由 dbWriteQueue 统一打印 db_error 事件。 // 失败时返回错误,由 WriteQueue 统一打印 db_error 事件。
func insertInboundBotDirectMessage(s *store, record map[string]any, clientInfo mqttClientInfo) error { func insertInboundBotDirectMessage(s *Store, record map[string]any, clientInfo MQTTClientInfo) error {
if s == nil { if s == nil {
return fmt.Errorf("store is not configured") return fmt.Errorf("Store is not configured")
} }
if record == nil { if record == nil {
return fmt.Errorf("record is required") return fmt.Errorf("record is required")
@@ -267,13 +267,13 @@ func insertInboundBotDirectMessage(s *store, record map[string]any, clientInfo m
contentPtr = &s contentPtr = &s
} }
now := time.Now() now := time.Now()
dm := &botDirectMessageRecord{ dm := &BotDirectMessageRecord{
BotID: bot.ID, BotID: bot.ID,
BotNodeID: bot.NodeID, BotNodeID: bot.NodeID,
BotNodeNum: bot.NodeNum, BotNodeNum: bot.NodeNum,
PeerNodeID: peerNodeID, PeerNodeID: peerNodeID,
PeerNodeNum: int64(peerNum), PeerNodeNum: int64(peerNum),
Direction: botDirectMessageDirectionInbound, Direction: BotDirectMessageDirectionInbound,
Topic: topic, Topic: topic,
PacketID: int64(packetID), PacketID: int64(packetID),
Text: text, Text: text,
@@ -281,7 +281,7 @@ func insertInboundBotDirectMessage(s *store, record map[string]any, clientInfo m
PKIEncrypted: true, PKIEncrypted: true,
WantAck: wantAck, WantAck: wantAck,
GatewayID: gatewayPtr, GatewayID: gatewayPtr,
Status: botMessageStatusPublished, Status: BotMessageStatusPublished,
ReceivedAt: &now, ReceivedAt: &now,
ContentJSON: contentPtr, ContentJSON: contentPtr,
} }
@@ -292,9 +292,9 @@ func insertInboundBotDirectMessage(s *store, record map[string]any, clientInfo m
// 同时将消息添加到 LLM 队列(忽略机器人自己发送的消息) // 同时将消息添加到 LLM 队列(忽略机器人自己发送的消息)
if peerNodeID != bot.NodeID { if peerNodeID != bot.NodeID {
longName := nullableString(record["long_name"]) longName := NullableString(record["long_name"])
shortName := nullableString(record["short_name"]) shortName := NullableString(record["short_name"])
channelID := nullableString(record["channel_id"]) channelID := NullableString(record["channel_id"])
_, err = s.EnqueueLLMMessage(LLMMessageQueueInput{ _, err = s.EnqueueLLMMessage(LLMMessageQueueInput{
BotID: bot.ID, BotID: bot.ID,
BotNodeID: bot.NodeID, BotNodeID: bot.NodeID,
+48
View File
@@ -0,0 +1,48 @@
package store
import (
"encoding/base64"
"encoding/hex"
"errors"
"strings"
"gorm.io/gorm"
)
// GetBotNodeByNodeNum 按节点号查找受管 bot 节点;用于 PKI 解密时把 to 字段映射回本地私钥。
func (s *Store) GetBotNodeByNodeNum(nodeNum int64) (*BotNodeRecord, error) {
if s == nil || s.db == nil {
return nil, errors.New("store not configured")
}
var row BotNodeRecord
if err := s.db.Where("node_num = ?", nodeNum).Take(&row).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, err
}
return nil, err
}
return &row, nil
}
// LookupNodeInfoPublicKey 在 nodeinfo 表中按 node_num 查 X25519 公钥,
// 兼容 hex 与 base64 两种历史存储格式。
func (s *Store) LookupNodeInfoPublicKey(nodeNum uint32) ([]byte, bool) {
var row NodeInfoRecord
if err := s.db.Where("node_num = ?", int64(nodeNum)).Take(&row).Error; err != nil {
return nil, false
}
if row.PublicKey == nil {
return nil, false
}
value := strings.TrimSpace(*row.PublicKey)
if value == "" {
return nil, false
}
if decoded, err := hex.DecodeString(value); err == nil && len(decoded) == 32 {
return decoded, true
}
if decoded, err := base64.StdEncoding.DecodeString(value); err == nil && len(decoded) == 32 {
return decoded, true
}
return nil, false
}
+59 -59
View File
@@ -1,4 +1,4 @@
package main package store
import ( import (
"crypto/ecdh" "crypto/ecdh"
@@ -17,19 +17,19 @@ import (
) )
const ( const (
botDefaultTopicPrefix = "msh/CN" BotDefaultTopicPrefix = "msh/CN"
botDefaultPSK = "AQ==" BotDefaultPSK = "AQ=="
botDefaultNodeInfoBroadcastSeconds = int64(3600) BotDefaultNodeInfoBroadcastSeconds = int64(3600)
botMessageTypeChannel = "channel" BotMessageTypeChannel = "channel"
botMessageTypeDirect = "direct" BotMessageTypeDirect = "direct"
botMessageStatusPending = "pending" BotMessageStatusPending = "pending"
botMessageStatusPublished = "published" BotMessageStatusPublished = "published"
botMessageStatusFailed = "failed" BotMessageStatusFailed = "failed"
) )
var errBotNodeAlreadyExists = errors.New("bot node already exists") var ErrBotNodeAlreadyExists = errors.New("bot node already exists")
type botNodeInput struct { type BotNodeInput struct {
NodeNum *int64 NodeNum *int64
LongName string LongName string
ShortName string ShortName string
@@ -43,17 +43,17 @@ type botNodeInput struct {
LLMIncludeChannelMessages bool LLMIncludeChannelMessages bool
} }
type botMessageListOptions struct { type BotMessageListOptions struct {
listOptions ListOptions
BotID uint64 BotID uint64
MessageType string MessageType string
ChannelID string ChannelID string
} }
func (s *store) ListBotNodes(opts listOptions) ([]botNodeRecord, error) { func (s *Store) ListBotNodes(opts ListOptions) ([]BotNodeRecord, error) {
opts = normalizeListOptions(opts) opts = NormalizeListOptions(opts)
var rows []botNodeRecord var rows []BotNodeRecord
q := s.db.Model(&botNodeRecord{}). q := s.db.Model(&BotNodeRecord{}).
Order("updated_at DESC"). Order("updated_at DESC").
Order("id DESC"). Order("id DESC").
Limit(opts.Limit). Limit(opts.Limit).
@@ -61,20 +61,20 @@ func (s *store) ListBotNodes(opts listOptions) ([]botNodeRecord, error) {
return rows, q.Find(&rows).Error return rows, q.Find(&rows).Error
} }
func (s *store) CountBotNodes(opts listOptions) (int64, error) { func (s *Store) CountBotNodes(opts ListOptions) (int64, error) {
var total int64 var total int64
return total, s.db.Model(&botNodeRecord{}).Count(&total).Error return total, s.db.Model(&BotNodeRecord{}).Count(&total).Error
} }
func (s *store) GetBotNode(id uint64) (*botNodeRecord, error) { func (s *Store) GetBotNode(id uint64) (*BotNodeRecord, error) {
var row botNodeRecord var row BotNodeRecord
if err := s.db.Where("id = ?", id).Take(&row).Error; err != nil { if err := s.db.Where("id = ?", id).Take(&row).Error; err != nil {
return nil, err return nil, err
} }
return &row, nil return &row, nil
} }
func (s *store) CreateBotNode(input botNodeInput) (*botNodeRecord, error) { func (s *Store) CreateBotNode(input BotNodeInput) (*BotNodeRecord, error) {
row, err := s.normalizedBotNodeRecord(input) row, err := s.normalizedBotNodeRecord(input)
if err != nil { if err != nil {
return nil, err return nil, err
@@ -94,7 +94,7 @@ func (s *store) CreateBotNode(input botNodeInput) (*botNodeRecord, error) {
return row, nil return row, nil
} }
func (s *store) UpdateBotNode(id uint64, input botNodeInput) (*botNodeRecord, error) { func (s *Store) UpdateBotNode(id uint64, input BotNodeInput) (*BotNodeRecord, error) {
if id == 0 { if id == 0 {
return nil, fmt.Errorf("bot node id is required") return nil, fmt.Errorf("bot node id is required")
} }
@@ -132,14 +132,14 @@ func (s *store) UpdateBotNode(id uint64, input botNodeInput) (*botNodeRecord, er
"llm_include_channel_messages": row.LLMIncludeChannelMessages, "llm_include_channel_messages": row.LLMIncludeChannelMessages,
"updated_at": time.Now(), "updated_at": time.Now(),
} }
if err := s.db.Model(&botNodeRecord{}).Where("id = ?", id).Updates(updates).Error; err != nil { if err := s.db.Model(&BotNodeRecord{}).Where("id = ?", id).Updates(updates).Error; err != nil {
return nil, err return nil, err
} }
return s.GetBotNode(id) return s.GetBotNode(id)
} }
func (s *store) DeleteBotNode(id uint64) error { func (s *Store) DeleteBotNode(id uint64) error {
result := s.db.Where("id = ?", id).Delete(&botNodeRecord{}) result := s.db.Where("id = ?", id).Delete(&BotNodeRecord{})
if result.Error != nil { if result.Error != nil {
return result.Error return result.Error
} }
@@ -149,13 +149,13 @@ func (s *store) DeleteBotNode(id uint64) error {
return nil return nil
} }
func (s *store) InsertBotMessage(row *botMessageRecord) error { func (s *Store) InsertBotMessage(row *BotMessageRecord) error {
return s.db.Create(row).Error return s.db.Create(row).Error
} }
func (s *store) UpdateBotMessageStatus(id uint64, status, errText string, publishedAt *time.Time) error { func (s *Store) UpdateBotMessageStatus(id uint64, status, errText string, publishedAt *time.Time) error {
updates := map[string]any{"status": status, "error": strings.TrimSpace(errText), "published_at": publishedAt} updates := map[string]any{"status": status, "error": strings.TrimSpace(errText), "published_at": publishedAt}
result := s.db.Model(&botMessageRecord{}).Where("id = ?", id).Updates(updates) result := s.db.Model(&BotMessageRecord{}).Where("id = ?", id).Updates(updates)
if result.Error != nil { if result.Error != nil {
return result.Error return result.Error
} }
@@ -165,8 +165,8 @@ func (s *store) UpdateBotMessageStatus(id uint64, status, errText string, publis
return nil return nil
} }
func (s *store) UpdateBotNodeInfoBroadcastAt(id uint64, t time.Time) error { func (s *Store) UpdateBotNodeInfoBroadcastAt(id uint64, t time.Time) error {
result := s.db.Model(&botNodeRecord{}).Where("id = ?", id).Updates(map[string]any{"last_nodeinfo_broadcast_at": &t, "updated_at": time.Now()}) result := s.db.Model(&BotNodeRecord{}).Where("id = ?", id).Updates(map[string]any{"last_nodeinfo_broadcast_at": &t, "updated_at": time.Now()})
if result.Error != nil { if result.Error != nil {
return result.Error return result.Error
} }
@@ -176,7 +176,7 @@ func (s *store) UpdateBotNodeInfoBroadcastAt(id uint64, t time.Time) error {
return nil return nil
} }
func (s *store) RegenerateBotNodeKeys(id uint64) (*botNodeRecord, error) { func (s *Store) RegenerateBotNodeKeys(id uint64) (*BotNodeRecord, error) {
if id == 0 { if id == 0 {
return nil, fmt.Errorf("bot node id is required") return nil, fmt.Errorf("bot node id is required")
} }
@@ -188,16 +188,16 @@ func (s *store) RegenerateBotNodeKeys(id uint64) (*botNodeRecord, error) {
return nil, err return nil, err
} }
updates := map[string]any{"public_key": row.PublicKey, "private_key": row.PrivateKey, "updated_at": time.Now()} updates := map[string]any{"public_key": row.PublicKey, "private_key": row.PrivateKey, "updated_at": time.Now()}
if err := s.db.Model(&botNodeRecord{}).Where("id = ?", id).Updates(updates).Error; err != nil { if err := s.db.Model(&BotNodeRecord{}).Where("id = ?", id).Updates(updates).Error; err != nil {
return nil, err return nil, err
} }
return s.GetBotNode(id) return s.GetBotNode(id)
} }
func (s *store) ListBotMessages(opts botMessageListOptions) ([]botMessageRecord, error) { func (s *Store) ListBotMessages(opts BotMessageListOptions) ([]BotMessageRecord, error) {
opts.listOptions = normalizeListOptions(opts.listOptions) opts.ListOptions = NormalizeListOptions(opts.ListOptions)
var rows []botMessageRecord var rows []BotMessageRecord
q := applyBotMessageFilters(s.db.Model(&botMessageRecord{}), opts). q := applyBotMessageFilters(s.db.Model(&BotMessageRecord{}), opts).
Order("created_at DESC"). Order("created_at DESC").
Order("id DESC"). Order("id DESC").
Limit(opts.Limit). Limit(opts.Limit).
@@ -205,13 +205,13 @@ func (s *store) ListBotMessages(opts botMessageListOptions) ([]botMessageRecord,
return rows, q.Find(&rows).Error return rows, q.Find(&rows).Error
} }
func (s *store) CountBotMessages(opts botMessageListOptions) (int64, error) { func (s *Store) CountBotMessages(opts BotMessageListOptions) (int64, error) {
var total int64 var total int64
q := applyBotMessageFilters(s.db.Model(&botMessageRecord{}), opts) q := applyBotMessageFilters(s.db.Model(&BotMessageRecord{}), opts)
return total, q.Count(&total).Error return total, q.Count(&total).Error
} }
func applyBotMessageFilters(q *gorm.DB, opts botMessageListOptions) *gorm.DB { func applyBotMessageFilters(q *gorm.DB, opts BotMessageListOptions) *gorm.DB {
if opts.BotID != 0 { if opts.BotID != 0 {
q = q.Where("bot_id = ?", opts.BotID) q = q.Where("bot_id = ?", opts.BotID)
} }
@@ -230,20 +230,20 @@ func applyBotMessageFilters(q *gorm.DB, opts botMessageListOptions) *gorm.DB {
return q return q
} }
func (s *store) normalizedBotNodeRecord(input botNodeInput) (*botNodeRecord, error) { func (s *Store) normalizedBotNodeRecord(input BotNodeInput) (*BotNodeRecord, error) {
longName := strings.TrimSpace(input.LongName) longName := strings.TrimSpace(input.LongName)
shortName := strings.TrimSpace(input.ShortName) shortName := strings.TrimSpace(input.ShortName)
channelID := strings.TrimSpace(input.DefaultChannelID) channelID := strings.TrimSpace(input.DefaultChannelID)
psk := strings.TrimSpace(input.PSK) psk := strings.TrimSpace(input.PSK)
if psk == "" { if psk == "" {
psk = botDefaultPSK psk = BotDefaultPSK
} }
if _, err := mqtpp.ExpandPSK(psk); err != nil { if _, err := mqtpp.ExpandPSK(psk); err != nil {
return nil, err return nil, err
} }
topicPrefix := strings.Trim(strings.TrimSpace(input.TopicPrefix), "/") topicPrefix := strings.Trim(strings.TrimSpace(input.TopicPrefix), "/")
if topicPrefix == "" { if topicPrefix == "" {
topicPrefix = botDefaultTopicPrefix topicPrefix = BotDefaultTopicPrefix
} }
if longName == "" { if longName == "" {
return nil, fmt.Errorf("long name is required") return nil, fmt.Errorf("long name is required")
@@ -262,7 +262,7 @@ func (s *store) normalizedBotNodeRecord(input botNodeInput) (*botNodeRecord, err
} }
interval := input.NodeInfoBroadcastIntervalSeconds interval := input.NodeInfoBroadcastIntervalSeconds
if interval <= 0 { if interval <= 0 {
interval = botDefaultNodeInfoBroadcastSeconds interval = BotDefaultNodeInfoBroadcastSeconds
} }
if interval < 60 { if interval < 60 {
return nil, fmt.Errorf("nodeinfo broadcast interval must be at least 60 seconds") return nil, fmt.Errorf("nodeinfo broadcast interval must be at least 60 seconds")
@@ -277,13 +277,13 @@ func (s *store) normalizedBotNodeRecord(input botNodeInput) (*botNodeRecord, err
} else { } else {
nodeNum = *input.NodeNum nodeNum = *input.NodeNum
} }
if err := validateBotNodeNum(nodeNum); err != nil { if err := ValidateBotNodeNum(nodeNum); err != nil {
return nil, err return nil, err
} }
return &botNodeRecord{NodeID: mqtpp.NodeNumToID(uint32(nodeNum)), NodeNum: nodeNum, LongName: longName, ShortName: shortName, Enabled: input.Enabled, DefaultChannelID: channelID, TopicPrefix: topicPrefix, PSK: psk, NodeInfoBroadcastEnabled: input.NodeInfoBroadcastEnabled, NodeInfoBroadcastIntervalSeconds: interval, LLMQueueEnabled: input.LLMQueueEnabled, LLMIncludeChannelMessages: input.LLMIncludeChannelMessages}, nil return &BotNodeRecord{NodeID: mqtpp.NodeNumToID(uint32(nodeNum)), NodeNum: nodeNum, LongName: longName, ShortName: shortName, Enabled: input.Enabled, DefaultChannelID: channelID, TopicPrefix: topicPrefix, PSK: psk, NodeInfoBroadcastEnabled: input.NodeInfoBroadcastEnabled, NodeInfoBroadcastIntervalSeconds: interval, LLMQueueEnabled: input.LLMQueueEnabled, LLMIncludeChannelMessages: input.LLMIncludeChannelMessages}, nil
} }
func populateBotNodeKeys(row *botNodeRecord) error { func populateBotNodeKeys(row *BotNodeRecord) error {
privateKey, err := ecdh.X25519().GenerateKey(rand.Reader) privateKey, err := ecdh.X25519().GenerateKey(rand.Reader)
if err != nil { if err != nil {
return err return err
@@ -293,7 +293,7 @@ func populateBotNodeKeys(row *botNodeRecord) error {
return nil return nil
} }
func decodeBotPublicKey(row botNodeRecord) ([]byte, error) { func DecodeBotPublicKey(row BotNodeRecord) ([]byte, error) {
if strings.TrimSpace(row.PublicKey) == "" { if strings.TrimSpace(row.PublicKey) == "" {
return nil, nil return nil, nil
} }
@@ -304,31 +304,31 @@ func decodeBotPublicKey(row botNodeRecord) ([]byte, error) {
return key, nil return key, nil
} }
func validateBotNodeNum(nodeNum int64) error { func ValidateBotNodeNum(nodeNum int64) error {
if nodeNum <= 0 || nodeNum >= int64(mqtpp.NodeNumBroadcast) { if nodeNum <= 0 || nodeNum >= int64(mqtpp.NodeNumBroadcast) {
return fmt.Errorf("node num must be between 1 and 4294967294") return fmt.Errorf("node num must be between 1 and 4294967294")
} }
return nil return nil
} }
func (s *store) generateBotNodeNum() (int64, error) { func (s *Store) generateBotNodeNum() (int64, error) {
for i := 0; i < 32; i++ { for i := 0; i < 32; i++ {
var buf [4]byte var buf [4]byte
if _, err := rand.Read(buf[:]); err != nil { if _, err := rand.Read(buf[:]); err != nil {
return 0, err return 0, err
} }
nodeNum := int64(binary.LittleEndian.Uint32(buf[:]) & 0x7fffffff) nodeNum := int64(binary.LittleEndian.Uint32(buf[:]) & 0x7fffffff)
if err := validateBotNodeNum(nodeNum); err != nil { if err := ValidateBotNodeNum(nodeNum); err != nil {
continue continue
} }
if err := s.ensureBotNodeUnique(0, mqtpp.NodeNumToID(uint32(nodeNum)), nodeNum); err != nil { if err := s.ensureBotNodeUnique(0, mqtpp.NodeNumToID(uint32(nodeNum)), nodeNum); err != nil {
if errors.Is(err, errBotNodeAlreadyExists) { if errors.Is(err, ErrBotNodeAlreadyExists) {
continue continue
} }
return 0, err return 0, err
} }
if err := s.ensureBotNodeDoesNotConflictWithNodeInfo(nodeNum, mqtpp.NodeNumToID(uint32(nodeNum))); err != nil { if err := s.ensureBotNodeDoesNotConflictWithNodeInfo(nodeNum, mqtpp.NodeNumToID(uint32(nodeNum))); err != nil {
if errors.Is(err, errBotNodeAlreadyExists) { if errors.Is(err, ErrBotNodeAlreadyExists) {
continue continue
} }
return 0, err return 0, err
@@ -338,15 +338,15 @@ func (s *store) generateBotNodeNum() (int64, error) {
return 0, fmt.Errorf("generate bot node num failed") return 0, fmt.Errorf("generate bot node num failed")
} }
func (s *store) ensureBotNodeUnique(id uint64, nodeID string, nodeNum int64) error { func (s *Store) ensureBotNodeUnique(id uint64, nodeID string, nodeNum int64) error {
var existing botNodeRecord var existing BotNodeRecord
q := s.db.Where("node_id = ? OR node_num = ?", nodeID, nodeNum) q := s.db.Where("node_id = ? OR node_num = ?", nodeID, nodeNum)
if id != 0 { if id != 0 {
q = q.Where("id <> ?", id) q = q.Where("id <> ?", id)
} }
err := q.Take(&existing).Error err := q.Take(&existing).Error
if err == nil { if err == nil {
return errBotNodeAlreadyExists return ErrBotNodeAlreadyExists
} }
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
return nil return nil
@@ -354,8 +354,8 @@ func (s *store) ensureBotNodeUnique(id uint64, nodeID string, nodeNum int64) err
return err return err
} }
func (s *store) ensureBotNodeDoesNotConflictWithNodeInfo(nodeNum int64, selfNodeID string) error { func (s *Store) ensureBotNodeDoesNotConflictWithNodeInfo(nodeNum int64, selfNodeID string) error {
var existing nodeInfoRecord var existing NodeInfoRecord
q := s.db.Where("node_num = ?", nodeNum) q := s.db.Where("node_num = ?", nodeNum)
if selfNodeID != "" { if selfNodeID != "" {
// 机器人自己广播 NodeInfo 后会以同样的 node_id/node_num 回写 nodeinfo // 机器人自己广播 NodeInfo 后会以同样的 node_id/node_num 回写 nodeinfo
@@ -364,7 +364,7 @@ func (s *store) ensureBotNodeDoesNotConflictWithNodeInfo(nodeNum int64, selfNode
} }
err := q.Take(&existing).Error err := q.Take(&existing).Error
if err == nil { if err == nil {
return errBotNodeAlreadyExists return ErrBotNodeAlreadyExists
} }
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
return nil return nil
File diff suppressed because it is too large Load Diff
+47 -45
View File
@@ -1,4 +1,4 @@
package main package store
import ( import (
"database/sql" "database/sql"
@@ -10,6 +10,8 @@ import (
"time" "time"
"gorm.io/gorm" "gorm.io/gorm"
"meshtastic_mqtt_server/internal/config"
) )
func TestOpenStoreCreatesTables(t *testing.T) { func TestOpenStoreCreatesTables(t *testing.T) {
@@ -49,7 +51,7 @@ func TestCountSignsByDayFormatsDateString(t *testing.T) {
t.Fatalf("CreateSign() error = %v", err) t.Fatalf("CreateSign() error = %v", err)
} }
rows, err := st.CountSignsByDay(listOptions{}) rows, err := st.CountSignsByDay(ListOptions{})
if err != nil { if err != nil {
t.Fatalf("CountSignsByDay() error = %v", err) t.Fatalf("CountSignsByDay() error = %v", err)
} }
@@ -164,7 +166,7 @@ func TestListMapReportsFiltersByBounds(t *testing.T) {
minLat, maxLat := 10.0, 11.0 minLat, maxLat := 10.0, 11.0
minLng, maxLng := 20.0, 21.0 minLng, maxLng := 20.0, 21.0
opts := listOptions{Limit: 100, MinLat: &minLat, MaxLat: &maxLat, MinLng: &minLng, MaxLng: &maxLng} opts := ListOptions{Limit: 100, MinLat: &minLat, MaxLat: &maxLat, MinLng: &minLng, MaxLng: &maxLng}
rows, err := st.ListMapReports(opts) rows, err := st.ListMapReports(opts)
if err != nil { if err != nil {
t.Fatalf("ListMapReports() error = %v", err) t.Fatalf("ListMapReports() error = %v", err)
@@ -206,7 +208,7 @@ func TestListMapReportsFiltersAcrossAntimeridian(t *testing.T) {
minLat, maxLat := -10.0, 10.0 minLat, maxLat := -10.0, 10.0
minLng, maxLng := 170.0, -170.0 minLng, maxLng := 170.0, -170.0
rows, err := st.ListMapReports(listOptions{Limit: 100, MinLat: &minLat, MaxLat: &maxLat, MinLng: &minLng, MaxLng: &maxLng}) rows, err := st.ListMapReports(ListOptions{Limit: 100, MinLat: &minLat, MaxLat: &maxLat, MinLng: &minLng, MaxLng: &maxLng})
if err != nil { if err != nil {
t.Fatalf("ListMapReports() error = %v", err) t.Fatalf("ListMapReports() error = %v", err)
} }
@@ -239,8 +241,8 @@ func TestListMapReportViewportReturnsPointsBelowThreshold(t *testing.T) {
minLat, maxLat := -1.0, 5.0 minLat, maxLat := -1.0, 5.0
minLng, maxLng := -1.0, 5.0 minLng, maxLng := -1.0, 5.0
result, err := st.ListMapReportViewport(mapReportViewportOptions{ result, err := st.ListMapReportViewport(MapReportViewportOptions{
ListOptions: listOptions{MinLat: &minLat, MaxLat: &maxLat, MinLng: &minLng, MaxLng: &maxLng}, ListOptions: ListOptions{MinLat: &minLat, MaxLat: &maxLat, MinLng: &minLng, MaxLng: &maxLng},
Zoom: 8, Zoom: 8,
Limit: 1000, Limit: 1000,
ClusterThreshold: 10, ClusterThreshold: 10,
@@ -271,8 +273,8 @@ func TestListMapReportViewportReturnsClustersAboveThreshold(t *testing.T) {
minLat, maxLat := 9.0, 11.0 minLat, maxLat := 9.0, 11.0
minLng, maxLng := 19.0, 21.0 minLng, maxLng := 19.0, 21.0
result, err := st.ListMapReportViewport(mapReportViewportOptions{ result, err := st.ListMapReportViewport(MapReportViewportOptions{
ListOptions: listOptions{MinLat: &minLat, MaxLat: &maxLat, MinLng: &minLng, MaxLng: &maxLng}, ListOptions: ListOptions{MinLat: &minLat, MaxLat: &maxLat, MinLng: &minLng, MaxLng: &maxLng},
Zoom: 4, Zoom: 4,
Limit: 1000, Limit: 1000,
ClusterThreshold: 2, ClusterThreshold: 2,
@@ -498,7 +500,7 @@ func TestEnsureDefaultAdminCreatesAdminUser(t *testing.T) {
if err != nil { if err != nil {
t.Fatalf("GetUserByUsername() error = %v", err) t.Fatalf("GetUserByUsername() error = %v", err)
} }
if user.Role != adminRole { if user.Role != AdminRole {
t.Fatalf("role = %q, want admin", user.Role) t.Fatalf("role = %q, want admin", user.Role)
} }
if user.PasswordHash == "admin" || user.PasswordHash == "" { if user.PasswordHash == "admin" || user.PasswordHash == "" {
@@ -536,7 +538,7 @@ func TestCreateAdminUserCreatesHashedAdmin(t *testing.T) {
if err != nil { if err != nil {
t.Fatalf("CreateAdminUser() error = %v", err) t.Fatalf("CreateAdminUser() error = %v", err)
} }
if user.Username != "new-admin" || user.Role != adminRole { if user.Username != "new-admin" || user.Role != AdminRole {
t.Fatalf("user = %#v, want new-admin admin", user) t.Fatalf("user = %#v, want new-admin admin", user)
} }
if user.PasswordHash == "secret" || !verifyPassword(user.PasswordHash, "secret") { if user.PasswordHash == "secret" || !verifyPassword(user.PasswordHash, "secret") {
@@ -551,8 +553,8 @@ func TestCreateAdminUserRejectsDuplicateUsername(t *testing.T) {
if _, err := st.CreateAdminUser("new-admin", "secret"); err != nil { if _, err := st.CreateAdminUser("new-admin", "secret"); err != nil {
t.Fatalf("first CreateAdminUser() error = %v", err) t.Fatalf("first CreateAdminUser() error = %v", err)
} }
if _, err := st.CreateAdminUser("new-admin", "secret"); !errors.Is(err, errUserAlreadyExists) { if _, err := st.CreateAdminUser("new-admin", "secret"); !errors.Is(err, ErrUserAlreadyExists) {
t.Fatalf("duplicate CreateAdminUser() error = %v, want errUserAlreadyExists", err) t.Fatalf("duplicate CreateAdminUser() error = %v, want ErrUserAlreadyExists", err)
} }
} }
@@ -591,14 +593,14 @@ func TestInsertAndListLoginLogs(t *testing.T) {
defer st.Close() defer st.Close()
userID := uint64(1) userID := uint64(1)
if err := st.InsertLoginLog(loginLogRecord{Username: "admin", UserID: &userID, Success: true, Reason: "success", RemoteAddr: "127.0.0.1:1234", RemoteHost: "127.0.0.1", UserAgent: "test-agent"}); err != nil { if err := st.InsertLoginLog(LoginLogRecord{Username: "admin", UserID: &userID, Success: true, Reason: "success", RemoteAddr: "127.0.0.1:1234", RemoteHost: "127.0.0.1", UserAgent: "test-agent"}); err != nil {
t.Fatalf("InsertLoginLog(success) error = %v", err) t.Fatalf("InsertLoginLog(success) error = %v", err)
} }
if err := st.InsertLoginLog(loginLogRecord{Username: "admin", Success: false, Reason: "invalid username or password", RemoteAddr: "127.0.0.1:1235", RemoteHost: "127.0.0.1", UserAgent: "test-agent"}); err != nil { if err := st.InsertLoginLog(LoginLogRecord{Username: "admin", Success: false, Reason: "invalid username or password", RemoteAddr: "127.0.0.1:1235", RemoteHost: "127.0.0.1", UserAgent: "test-agent"}); err != nil {
t.Fatalf("InsertLoginLog(failure) error = %v", err) t.Fatalf("InsertLoginLog(failure) error = %v", err)
} }
logs, err := st.ListLoginLogs(listOptions{Limit: 10}) logs, err := st.ListLoginLogs(ListOptions{Limit: 10})
if err != nil { if err != nil {
t.Fatalf("ListLoginLogs() error = %v", err) t.Fatalf("ListLoginLogs() error = %v", err)
} }
@@ -621,7 +623,7 @@ func TestInsertDiscardDetailsStoresRawBase64AndClientInfo(t *testing.T) {
defer st.Close() defer st.Close()
raw := []byte{0xff, 0x00, 0x01} raw := []byte{0xff, 0x00, 0x01}
clientInfo := mqttClientInfo{ClientID: "client-1", Username: "user-1", Listener: "tcp", RemoteAddr: "127.0.0.1:54321", RemoteHost: "127.0.0.1", RemotePort: "54321"} clientInfo := MQTTClientInfo{ClientID: "client-1", Username: "user-1", Listener: "tcp", RemoteAddr: "127.0.0.1:54321", RemoteHost: "127.0.0.1", RemotePort: "54321"}
record := map[string]any{"topic": "msh/US/test", "error": "protobuf decode failed", "payload_len": len(raw)} record := map[string]any{"topic": "msh/US/test", "error": "protobuf decode failed", "payload_len": len(raw)}
if err := st.InsertDiscardDetails(record, raw, clientInfo); err != nil { if err := st.InsertDiscardDetails(record, raw, clientInfo); err != nil {
t.Fatalf("InsertDiscardDetails() error = %v", err) t.Fatalf("InsertDiscardDetails() error = %v", err)
@@ -647,13 +649,13 @@ func TestListDiscardDetailsOrdersNewestFirst(t *testing.T) {
st := openTestStore(t) st := openTestStore(t)
defer st.Close() defer st.Close()
if err := st.InsertDiscardDetails(map[string]any{"topic": "first", "error": "first"}, []byte{1}, mqttClientInfo{}); err != nil { if err := st.InsertDiscardDetails(map[string]any{"topic": "first", "error": "first"}, []byte{1}, MQTTClientInfo{}); err != nil {
t.Fatalf("first InsertDiscardDetails() error = %v", err) t.Fatalf("first InsertDiscardDetails() error = %v", err)
} }
if err := st.InsertDiscardDetails(map[string]any{"topic": "second", "error": "second"}, []byte{2}, mqttClientInfo{}); err != nil { if err := st.InsertDiscardDetails(map[string]any{"topic": "second", "error": "second"}, []byte{2}, MQTTClientInfo{}); err != nil {
t.Fatalf("second InsertDiscardDetails() error = %v", err) t.Fatalf("second InsertDiscardDetails() error = %v", err)
} }
rows, err := st.ListDiscardDetails(listOptions{Limit: 10}) rows, err := st.ListDiscardDetails(ListOptions{Limit: 10})
if err != nil { if err != nil {
t.Fatalf("ListDiscardDetails() error = %v", err) t.Fatalf("ListDiscardDetails() error = %v", err)
} }
@@ -669,7 +671,7 @@ func TestInsertTextMessageAppendsRows(t *testing.T) {
st := openTestStore(t) st := openTestStore(t)
defer st.Close() defer st.Close()
clientInfo := mqttClientInfo{ClientID: "client-1", Username: "user-1", Listener: "tcp", RemoteAddr: "127.0.0.1:54321", RemoteHost: "127.0.0.1", RemotePort: "54321"} clientInfo := MQTTClientInfo{ClientID: "client-1", Username: "user-1", Listener: "tcp", RemoteAddr: "127.0.0.1:54321", RemoteHost: "127.0.0.1", RemotePort: "54321"}
if err := st.InsertTextMessage(textMessageTestRecord("hello"), clientInfo); err != nil { if err := st.InsertTextMessage(textMessageTestRecord("hello"), clientInfo); err != nil {
t.Fatalf("first InsertTextMessage() error = %v", err) t.Fatalf("first InsertTextMessage() error = %v", err)
} }
@@ -710,7 +712,7 @@ func TestDeleteTextMessageDeletesRows(t *testing.T) {
st := openTestStore(t) st := openTestStore(t)
defer st.Close() defer st.Close()
if err := st.InsertTextMessage(textMessageTestRecord("hello"), mqttClientInfo{}); err != nil { if err := st.InsertTextMessage(textMessageTestRecord("hello"), MQTTClientInfo{}); err != nil {
t.Fatalf("InsertTextMessage() error = %v", err) t.Fatalf("InsertTextMessage() error = %v", err)
} }
var id uint64 var id uint64
@@ -736,7 +738,7 @@ func TestInsertTextMessageStoresClientInfo(t *testing.T) {
st := openTestStore(t) st := openTestStore(t)
defer st.Close() defer st.Close()
clientInfo := mqttClientInfo{ClientID: "client-1", Username: "user-1", Listener: "tcp", RemoteAddr: "127.0.0.1:54321", RemoteHost: "127.0.0.1", RemotePort: "54321"} clientInfo := MQTTClientInfo{ClientID: "client-1", Username: "user-1", Listener: "tcp", RemoteAddr: "127.0.0.1:54321", RemoteHost: "127.0.0.1", RemotePort: "54321"}
if err := st.InsertTextMessage(textMessageTestRecord("hello"), clientInfo); err != nil { if err := st.InsertTextMessage(textMessageTestRecord("hello"), clientInfo); err != nil {
t.Fatalf("InsertTextMessage() error = %v", err) t.Fatalf("InsertTextMessage() error = %v", err)
} }
@@ -756,7 +758,7 @@ func TestInsertTextMessageStoresPayloadHex(t *testing.T) {
record := textMessageTestRecord(nil) record := textMessageTestRecord(nil)
record["payload_hex"] = "fffefd" record["payload_hex"] = "fffefd"
if err := st.InsertTextMessage(record, mqttClientInfo{}); err != nil { if err := st.InsertTextMessage(record, MQTTClientInfo{}); err != nil {
t.Fatalf("InsertTextMessage() error = %v", err) t.Fatalf("InsertTextMessage() error = %v", err)
} }
@@ -777,16 +779,16 @@ func TestInsertTextMessageRequiresFields(t *testing.T) {
st := openTestStore(t) st := openTestStore(t)
defer st.Close() defer st.Close()
if err := st.InsertTextMessage(map[string]any{"type": "nodeinfo"}, mqttClientInfo{}); err == nil || !strings.Contains(err.Error(), "text_message") { if err := st.InsertTextMessage(map[string]any{"type": "nodeinfo"}, MQTTClientInfo{}); err == nil || !strings.Contains(err.Error(), "text_message") {
t.Fatalf("wrong type error = %v, want text_message error", err) t.Fatalf("wrong type error = %v, want text_message error", err)
} }
if err := st.InsertTextMessage(map[string]any{"type": "text_message", "from_num": 1, "topic": "msh/test"}, mqttClientInfo{}); err == nil || !strings.Contains(err.Error(), "from") { if err := st.InsertTextMessage(map[string]any{"type": "text_message", "from_num": 1, "topic": "msh/test"}, MQTTClientInfo{}); err == nil || !strings.Contains(err.Error(), "from") {
t.Fatalf("missing from error = %v, want from error", err) t.Fatalf("missing from error = %v, want from error", err)
} }
if err := st.InsertTextMessage(map[string]any{"type": "text_message", "from": "!00000001", "topic": "msh/test"}, mqttClientInfo{}); err == nil || !strings.Contains(err.Error(), "from_num") { if err := st.InsertTextMessage(map[string]any{"type": "text_message", "from": "!00000001", "topic": "msh/test"}, MQTTClientInfo{}); err == nil || !strings.Contains(err.Error(), "from_num") {
t.Fatalf("missing from_num error = %v, want from_num error", err) t.Fatalf("missing from_num error = %v, want from_num error", err)
} }
if err := st.InsertTextMessage(map[string]any{"type": "text_message", "from": "!00000001", "from_num": 1}, mqttClientInfo{}); err == nil || !strings.Contains(err.Error(), "topic") { if err := st.InsertTextMessage(map[string]any{"type": "text_message", "from": "!00000001", "from_num": 1}, MQTTClientInfo{}); err == nil || !strings.Contains(err.Error(), "topic") {
t.Fatalf("missing topic error = %v, want topic error", err) t.Fatalf("missing topic error = %v, want topic error", err)
} }
} }
@@ -795,7 +797,7 @@ func TestInsertPositionAppendsRows(t *testing.T) {
st := openTestStore(t) st := openTestStore(t)
defer st.Close() defer st.Close()
clientInfo := mqttClientInfo{ClientID: "client-1", RemoteAddr: "127.0.0.1:54321", RemoteHost: "127.0.0.1", RemotePort: "54321"} clientInfo := MQTTClientInfo{ClientID: "client-1", RemoteAddr: "127.0.0.1:54321", RemoteHost: "127.0.0.1", RemotePort: "54321"}
if err := st.InsertPosition(positionTestRecord(), clientInfo); err != nil { if err := st.InsertPosition(positionTestRecord(), clientInfo); err != nil {
t.Fatalf("first InsertPosition() error = %v", err) t.Fatalf("first InsertPosition() error = %v", err)
} }
@@ -826,7 +828,7 @@ func TestInsertPositionCreatesMapReportWhenMissing(t *testing.T) {
st := openTestStore(t) st := openTestStore(t)
defer st.Close() defer st.Close()
if err := st.InsertPosition(positionTestRecord(), mqttClientInfo{}); err != nil { if err := st.InsertPosition(positionTestRecord(), MQTTClientInfo{}); err != nil {
t.Fatalf("InsertPosition() error = %v", err) t.Fatalf("InsertPosition() error = %v", err)
} }
@@ -854,7 +856,7 @@ func TestInsertPositionUpdatesExistingMapReportCoordinates(t *testing.T) {
position["longitude"] = 120.75 position["longitude"] = 120.75
position["altitude"] = int32(88) position["altitude"] = int32(88)
position["precision_bits"] = uint32(10) position["precision_bits"] = uint32(10)
if err := st.InsertPosition(position, mqttClientInfo{}); err != nil { if err := st.InsertPosition(position, MQTTClientInfo{}); err != nil {
t.Fatalf("InsertPosition() error = %v", err) t.Fatalf("InsertPosition() error = %v", err)
} }
@@ -873,7 +875,7 @@ func TestInsertTelemetryAppendsRowsAndStoresMetricsJSON(t *testing.T) {
st := openTestStore(t) st := openTestStore(t)
defer st.Close() defer st.Close()
if err := st.InsertTelemetry(telemetryTestRecord(), mqttClientInfo{}); err != nil { if err := st.InsertTelemetry(telemetryTestRecord(), MQTTClientInfo{}); err != nil {
t.Fatalf("InsertTelemetry() error = %v", err) t.Fatalf("InsertTelemetry() error = %v", err)
} }
@@ -896,16 +898,16 @@ func TestInsertRoutingAndTracerouteAppendRows(t *testing.T) {
st := openTestStore(t) st := openTestStore(t)
defer st.Close() defer st.Close()
if err := st.InsertRouting(routingTestRecord(), mqttClientInfo{}); err != nil { if err := st.InsertRouting(routingTestRecord(), MQTTClientInfo{}); err != nil {
t.Fatalf("first InsertRouting() error = %v", err) t.Fatalf("first InsertRouting() error = %v", err)
} }
if err := st.InsertRouting(routingTestRecord(), mqttClientInfo{}); err != nil { if err := st.InsertRouting(routingTestRecord(), MQTTClientInfo{}); err != nil {
t.Fatalf("second InsertRouting() error = %v", err) t.Fatalf("second InsertRouting() error = %v", err)
} }
if err := st.InsertTraceroute(tracerouteTestRecord(), mqttClientInfo{}); err != nil { if err := st.InsertTraceroute(tracerouteTestRecord(), MQTTClientInfo{}); err != nil {
t.Fatalf("first InsertTraceroute() error = %v", err) t.Fatalf("first InsertTraceroute() error = %v", err)
} }
if err := st.InsertTraceroute(tracerouteTestRecord(), mqttClientInfo{}); err != nil { if err := st.InsertTraceroute(tracerouteTestRecord(), MQTTClientInfo{}); err != nil {
t.Fatalf("second InsertTraceroute() error = %v", err) t.Fatalf("second InsertTraceroute() error = %v", err)
} }
@@ -938,10 +940,10 @@ func TestInsertPacketTablesRequireFields(t *testing.T) {
insert func(map[string]any) error insert func(map[string]any) error
record map[string]any record map[string]any
}{ }{
{name: "position", insert: func(r map[string]any) error { return st.InsertPosition(r, mqttClientInfo{}) }, record: positionTestRecord()}, {name: "position", insert: func(r map[string]any) error { return st.InsertPosition(r, MQTTClientInfo{}) }, record: positionTestRecord()},
{name: "telemetry", insert: func(r map[string]any) error { return st.InsertTelemetry(r, mqttClientInfo{}) }, record: telemetryTestRecord()}, {name: "telemetry", insert: func(r map[string]any) error { return st.InsertTelemetry(r, MQTTClientInfo{}) }, record: telemetryTestRecord()},
{name: "routing", insert: func(r map[string]any) error { return st.InsertRouting(r, mqttClientInfo{}) }, record: routingTestRecord()}, {name: "routing", insert: func(r map[string]any) error { return st.InsertRouting(r, MQTTClientInfo{}) }, record: routingTestRecord()},
{name: "traceroute", insert: func(r map[string]any) error { return st.InsertTraceroute(r, mqttClientInfo{}) }, record: tracerouteTestRecord()}, {name: "traceroute", insert: func(r map[string]any) error { return st.InsertTraceroute(r, MQTTClientInfo{}) }, record: tracerouteTestRecord()},
} }
for _, tt := range tests { for _, tt := range tests {
@@ -971,19 +973,19 @@ func TestInsertPacketTablesRequireFields(t *testing.T) {
} }
} }
func openTestStore(t *testing.T) *store { func openTestStore(t *testing.T) *Store {
t.Helper() t.Helper()
st, err := openStore(databaseConfig{ st, err := OpenStore(config.DatabaseConfig{
Driver: databaseDriverSQLite, Driver: config.DriverSQLite,
SQLite: sqliteConfig{Path: filepath.Join(t.TempDir(), "mesh_mqtt_go.db")}, SQLite: config.SQLiteConfig{Path: filepath.Join(t.TempDir(), "mesh_mqtt_go.db")},
}) })
if err != nil { if err != nil {
t.Fatalf("openStore() error = %v", err) t.Fatalf("OpenStore() error = %v", err)
} }
return st return st
} }
func rawTestDB(t *testing.T, st *store) *sql.DB { func rawTestDB(t *testing.T, st *Store) *sql.DB {
t.Helper() t.Helper()
db, err := st.db.DB() db, err := st.db.DB()
if err != nil { if err != nil {
@@ -1,94 +1,94 @@
package main package store
import "sync" import "sync"
type dbWriteQueue struct { type WriteQueue struct {
store *store store *Store
jobs chan dbWriteJob jobs chan writeJob
wg sync.WaitGroup wg sync.WaitGroup
} }
type dbWriteJob struct { type writeJob struct {
typeName string typeName string
from any from any
run func() error run func() error
errorEvent map[string]any errorEvent map[string]any
} }
func newDBWriteQueue(store *store) *dbWriteQueue { func NewWriteQueue(s *Store) *WriteQueue {
if store == nil { if s == nil {
return nil return nil
} }
q := &dbWriteQueue{ q := &WriteQueue{
store: store, store: s,
jobs: make(chan dbWriteJob, 1024), jobs: make(chan writeJob, 1024),
} }
q.wg.Add(1) q.wg.Add(1)
go q.run() go q.run()
return q return q
} }
func (q *dbWriteQueue) EnqueueRecord(record map[string]any, clientInfo mqttClientInfo) { func (q *WriteQueue) EnqueueRecord(record map[string]any, clientInfo MQTTClientInfo) {
if q == nil { if q == nil {
return return
} }
record = cloneDBWriteRecord(record) record = cloneDBWriteRecord(record)
switch record["type"] { switch record["type"] {
case "nodeinfo": case "nodeinfo":
q.enqueue(dbWriteJob{typeName: "nodeinfo", from: record["from"], run: func() error { q.enqueue(writeJob{typeName: "nodeinfo", from: record["from"], run: func() error {
return q.store.UpsertNodeInfo(record) return q.store.UpsertNodeInfo(record)
}}) }})
case "map_report": case "map_report":
q.enqueue(dbWriteJob{typeName: "map_report", from: record["from"], run: func() error { q.enqueue(writeJob{typeName: "map_report", from: record["from"], run: func() error {
return q.store.UpsertMapReport(record) return q.store.UpsertMapReport(record)
}}) }})
case "text_message": case "text_message":
// 私聊(PKI 加密、发往受管 bot)单独走 bot_direct_messages 表, // 私聊(PKI 加密、发往受管 bot)单独走 bot_direct_messages 表,
// 不再写入 text_message 以避免和频道消息混在一起。 // 不再写入 text_message 以避免和频道消息混在一起。
if isInboundBotDirectMessage(q.store, record) { if isInboundBotDirectMessage(q.store, record) {
q.enqueue(dbWriteJob{typeName: "bot_direct_message_inbound", from: record["from"], run: func() error { q.enqueue(writeJob{typeName: "bot_direct_message_inbound", from: record["from"], run: func() error {
return insertInboundBotDirectMessage(q.store, record, clientInfo) return insertInboundBotDirectMessage(q.store, record, clientInfo)
}}) }})
return return
} }
// 频道消息同时也写入 LLM 队列(如果启用的话) // 频道消息同时也写入 LLM 队列(如果启用的话)
q.enqueue(dbWriteJob{typeName: "llm_channel_message", from: record["from"], run: func() error { q.enqueue(writeJob{typeName: "llm_channel_message", from: record["from"], run: func() error {
return enqueueChannelMessageToLLM(q.store, record) return enqueueChannelMessageToLLM(q.store, record)
}}) }})
q.enqueue(dbWriteJob{typeName: "text_message", from: record["from"], run: func() error { q.enqueue(writeJob{typeName: "text_message", from: record["from"], run: func() error {
return q.store.InsertTextMessage(record, clientInfo) return q.store.InsertTextMessage(record, clientInfo)
}}) }})
case "position": case "position":
q.enqueue(dbWriteJob{typeName: "position", from: record["from"], run: func() error { q.enqueue(writeJob{typeName: "position", from: record["from"], run: func() error {
return q.store.InsertPosition(record, clientInfo) return q.store.InsertPosition(record, clientInfo)
}}) }})
case "telemetry": case "telemetry":
q.enqueue(dbWriteJob{typeName: "telemetry", from: record["from"], run: func() error { q.enqueue(writeJob{typeName: "telemetry", from: record["from"], run: func() error {
return q.store.InsertTelemetry(record, clientInfo) return q.store.InsertTelemetry(record, clientInfo)
}}) }})
case "routing": case "routing":
q.enqueue(dbWriteJob{typeName: "routing", from: record["from"], run: func() error { q.enqueue(writeJob{typeName: "routing", from: record["from"], run: func() error {
return q.store.InsertRouting(record, clientInfo) return q.store.InsertRouting(record, clientInfo)
}}) }})
case "traceroute": case "traceroute":
q.enqueue(dbWriteJob{typeName: "traceroute", from: record["from"], run: func() error { q.enqueue(writeJob{typeName: "traceroute", from: record["from"], run: func() error {
return q.store.InsertTraceroute(record, clientInfo) return q.store.InsertTraceroute(record, clientInfo)
}}) }})
} }
} }
func (q *dbWriteQueue) EnqueueDiscard(record map[string]any, raw []byte, clientInfo mqttClientInfo) { func (q *WriteQueue) EnqueueDiscard(record map[string]any, raw []byte, clientInfo MQTTClientInfo) {
if q == nil { if q == nil {
return return
} }
record = cloneDBWriteRecord(record) record = cloneDBWriteRecord(record)
raw = append([]byte(nil), raw...) raw = append([]byte(nil), raw...)
q.enqueue(dbWriteJob{typeName: "discard_details", from: record["from"], errorEvent: map[string]any{"event": "db_error", "type": "discard_details", "topic": record["topic"]}, run: func() error { q.enqueue(writeJob{typeName: "discard_details", from: record["from"], errorEvent: map[string]any{"event": "db_error", "type": "discard_details", "topic": record["topic"]}, run: func() error {
return q.store.InsertDiscardDetails(record, raw, clientInfo) return q.store.InsertDiscardDetails(record, raw, clientInfo)
}}) }})
} }
func (q *dbWriteQueue) Close() { func (q *WriteQueue) Close() {
if q == nil { if q == nil {
return return
} }
@@ -96,18 +96,18 @@ func (q *dbWriteQueue) Close() {
q.wg.Wait() q.wg.Wait()
} }
func (q *dbWriteQueue) Len() int { func (q *WriteQueue) Len() int {
if q == nil { if q == nil {
return 0 return 0
} }
return len(q.jobs) return len(q.jobs)
} }
func (q *dbWriteQueue) enqueue(job dbWriteJob) { func (q *WriteQueue) enqueue(job writeJob) {
q.jobs <- job q.jobs <- job
} }
func (q *dbWriteQueue) run() { func (q *WriteQueue) run() {
defer q.wg.Done() defer q.wg.Done()
for job := range q.jobs { for job := range q.jobs {
if err := job.run(); err != nil { if err := job.run(); err != nil {
@@ -1,4 +1,4 @@
package main package store
import ( import (
"database/sql" "database/sql"
@@ -11,7 +11,7 @@ func TestDBWriteQueueWritesRecordsAsync(t *testing.T) {
queue := newDBWriteQueue(st) queue := newDBWriteQueue(st)
record := textMessageTestRecord("queued") record := textMessageTestRecord("queued")
queue.EnqueueRecord(record, mqttClientInfo{ClientID: "client-1"}) queue.EnqueueRecord(record, MQTTClientInfo{ClientID: "client-1"})
record["text"] = "mutated after enqueue" record["text"] = "mutated after enqueue"
queue.Close() queue.Close()
@@ -30,7 +30,7 @@ func TestDBWriteQueueWritesDiscardAsync(t *testing.T) {
queue := newDBWriteQueue(st) queue := newDBWriteQueue(st)
record := map[string]any{"topic": "msh/test", "error": "bad packet"} record := map[string]any{"topic": "msh/test", "error": "bad packet"}
queue.EnqueueDiscard(record, []byte{1, 2, 3}, mqttClientInfo{RemoteAddr: "127.0.0.1:1883"}) queue.EnqueueDiscard(record, []byte{1, 2, 3}, MQTTClientInfo{RemoteAddr: "127.0.0.1:1883"})
record["error"] = "mutated after enqueue" record["error"] = "mutated after enqueue"
queue.Close() queue.Close()
@@ -44,8 +44,8 @@ func TestDBWriteQueueWritesDiscardAsync(t *testing.T) {
} }
func TestDBWriteQueueLen(t *testing.T) { func TestDBWriteQueueLen(t *testing.T) {
queue := &dbWriteQueue{jobs: make(chan dbWriteJob, 1)} queue := &WriteQueue{jobs: make(chan writeJob, 1)}
queue.enqueue(dbWriteJob{run: func() error { return nil }}) queue.enqueue(writeJob{run: func() error { return nil }})
if queue.Len() != 1 { if queue.Len() != 1 {
t.Fatalf("queue.Len() = %d, want 1", queue.Len()) t.Fatalf("queue.Len() = %d, want 1", queue.Len())
} }
@@ -56,7 +56,7 @@ func TestDBWriteQueueIgnoresUnsupportedRecordType(t *testing.T) {
defer st.Close() defer st.Close()
queue := newDBWriteQueue(st) queue := newDBWriteQueue(st)
queue.EnqueueRecord(map[string]any{"type": "empty_packet", "from": "!12345678"}, mqttClientInfo{}) queue.EnqueueRecord(map[string]any{"type": "empty_packet", "from": "!12345678"}, MQTTClientInfo{})
queue.Close() queue.Close()
var count int var count int
@@ -72,9 +72,9 @@ func TestDBWriteQueueNilStore(t *testing.T) {
if queue := newDBWriteQueue(nil); queue != nil { if queue := newDBWriteQueue(nil); queue != nil {
t.Fatalf("newDBWriteQueue(nil) = %#v, want nil", queue) t.Fatalf("newDBWriteQueue(nil) = %#v, want nil", queue)
} }
var queue *dbWriteQueue var queue *WriteQueue
queue.EnqueueRecord(textMessageTestRecord("ignored"), mqttClientInfo{}) queue.EnqueueRecord(textMessageTestRecord("ignored"), MQTTClientInfo{})
queue.EnqueueDiscard(map[string]any{"topic": "ignored"}, []byte{1}, mqttClientInfo{}) queue.EnqueueDiscard(map[string]any{"topic": "ignored"}, []byte{1}, MQTTClientInfo{})
queue.Close() queue.Close()
} }
@@ -85,8 +85,8 @@ func TestDBWriteQueueRecordValidationErrorDoesNotStopWorker(t *testing.T) {
queue := newDBWriteQueue(st) queue := newDBWriteQueue(st)
badRecord := textMessageTestRecord("bad") badRecord := textMessageTestRecord("bad")
delete(badRecord, "from") delete(badRecord, "from")
queue.EnqueueRecord(badRecord, mqttClientInfo{}) queue.EnqueueRecord(badRecord, MQTTClientInfo{})
queue.EnqueueRecord(textMessageTestRecord("good"), mqttClientInfo{}) queue.EnqueueRecord(textMessageTestRecord("good"), MQTTClientInfo{})
queue.Close() queue.Close()
var text string var text string
@@ -1,4 +1,4 @@
package main package store
import ( import (
"encoding/base64" "encoding/base64"
@@ -6,7 +6,7 @@ import (
"fmt" "fmt"
) )
func (s *store) InsertDiscardDetails(record map[string]any, raw []byte, clientInfo mqttClientInfo) error { func (s *Store) InsertDiscardDetails(record map[string]any, raw []byte, clientInfo MQTTClientInfo) error {
details, err := discardDetailsFromRecord(record, raw, clientInfo) details, err := discardDetailsFromRecord(record, raw, clientInfo)
if err != nil { if err != nil {
return err return err
@@ -17,28 +17,28 @@ func (s *store) InsertDiscardDetails(record map[string]any, raw []byte, clientIn
return nil return nil
} }
func discardDetailsFromRecord(record map[string]any, raw []byte, clientInfo mqttClientInfo) (*discardDetailsRecord, error) { func discardDetailsFromRecord(record map[string]any, raw []byte, clientInfo MQTTClientInfo) (*DiscardDetailsRecord, error) {
contentJSON, err := json.Marshal(record) contentJSON, err := json.Marshal(record)
if err != nil { if err != nil {
return nil, fmt.Errorf("encode discard_details content_json: %w", err) return nil, fmt.Errorf("encode discard_details content_json: %w", err)
} }
return &discardDetailsRecord{ return &DiscardDetailsRecord{
Topic: stringValue(record["topic"]), Topic: stringValue(record["topic"]),
Error: stringValue(record["error"]), Error: stringValue(record["error"]),
PayloadLen: int64(len(raw)), PayloadLen: int64(len(raw)),
RawBase64: base64.StdEncoding.EncodeToString(raw), RawBase64: base64.StdEncoding.EncodeToString(raw),
ContentJSON: string(contentJSON), ContentJSON: string(contentJSON),
MQTTClientID: nullableStringValue(clientInfo.ClientID), MQTTClientID: NullableStringValue(clientInfo.ClientID),
MQTTUsername: nullableStringValue(clientInfo.Username), MQTTUsername: NullableStringValue(clientInfo.Username),
MQTTListener: nullableStringValue(clientInfo.Listener), MQTTListener: NullableStringValue(clientInfo.Listener),
MQTTRemoteAddr: nullableStringValue(clientInfo.RemoteAddr), MQTTRemoteAddr: NullableStringValue(clientInfo.RemoteAddr),
MQTTRemoteHost: nullableStringValue(clientInfo.RemoteHost), MQTTRemoteHost: NullableStringValue(clientInfo.RemoteHost),
MQTTRemotePort: nullableStringValue(clientInfo.RemotePort), MQTTRemotePort: NullableStringValue(clientInfo.RemotePort),
}, nil }, nil
} }
func stringValue(value any) string { func stringValue(value any) string {
if s := nullableStringValue(value); s != nil { if s := NullableStringValue(value); s != nil {
return *s return *s
} }
return "" return ""
@@ -1,4 +1,4 @@
package main package store
import ( import (
"fmt" "fmt"
@@ -7,7 +7,7 @@ import (
const maxHelpMarkdownBytes = 200 * 1024 const maxHelpMarkdownBytes = 200 * 1024
const defaultHelpMarkdown = `## 连接地址 const DefaultHelpMarkdown = `## 连接地址
Meshtastic 设备连接到本服务提供的 MQTT broker Meshtastic 设备连接到本服务提供的 MQTT broker
@@ -29,15 +29,15 @@ const defaultHelpMarkdown = `## 连接地址
如果遇到 bug请在 GitHub [提交 issue](https://github.com/wuwenfengmi1998/meshtastic_mqtt_server),或联系邮箱 [kevin@lmve.net](mailto:kevin@lmve.net)。` 如果遇到 bug请在 GitHub [提交 issue](https://github.com/wuwenfengmi1998/meshtastic_mqtt_server),或联系邮箱 [kevin@lmve.net](mailto:kevin@lmve.net)。`
func (s *store) GetLatestHelpContent() (*helpContentRecord, error) { func (s *Store) GetLatestHelpContent() (*HelpContentRecord, error) {
var row helpContentRecord var row HelpContentRecord
if err := s.db.Order("id DESC").Take(&row).Error; err != nil { if err := s.db.Order("id DESC").Take(&row).Error; err != nil {
return nil, err return nil, err
} }
return &row, nil return &row, nil
} }
func (s *store) InsertHelpContent(markdown, createdBy string) (*helpContentRecord, error) { func (s *Store) InsertHelpContent(markdown, createdBy string) (*HelpContentRecord, error) {
markdown = strings.TrimSpace(markdown) markdown = strings.TrimSpace(markdown)
createdBy = strings.TrimSpace(createdBy) createdBy = strings.TrimSpace(createdBy)
if markdown == "" { if markdown == "" {
@@ -46,7 +46,7 @@ func (s *store) InsertHelpContent(markdown, createdBy string) (*helpContentRecor
if len([]byte(markdown)) > maxHelpMarkdownBytes { if len([]byte(markdown)) > maxHelpMarkdownBytes {
return nil, fmt.Errorf("markdown exceeds %d bytes", maxHelpMarkdownBytes) return nil, fmt.Errorf("markdown exceeds %d bytes", maxHelpMarkdownBytes)
} }
row := helpContentRecord{Markdown: markdown, CreatedBy: createdBy} row := HelpContentRecord{Markdown: markdown, CreatedBy: createdBy}
if err := s.db.Create(&row).Error; err != nil { if err := s.db.Create(&row).Error; err != nil {
return nil, err return nil, err
} }
+45
View File
@@ -0,0 +1,45 @@
package store
import "golang.org/x/crypto/bcrypt"
// AdminRole 是管理员账号在用户表里的角色字符串。其它包通过这个常量与
// `users.role` 字段对齐,避免硬编码。
const AdminRole = "admin"
// printJSON 是 store 包内部的诊断输出钩子。当前实现为 noop——保持与
// 重构前 main.go 的行为一致;如需启用调试,可在调用方替换。
func printJSON(record map[string]any) {
_ = record
}
// hashPassword 与 auth.go 中的散列实现保持一致(bcrypt 默认 cost)。
func hashPassword(password string) (string, error) {
hash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
if err != nil {
return "", err
}
return string(hash), nil
}
// uint32FromRecord 把 map[string]any 中的整型字段安全转换为 uint32。
func uint32FromRecord(value any) (uint32, bool) {
switch v := value.(type) {
case uint32:
return v, true
case int:
if v >= 0 {
return uint32(v), true
}
case int64:
if v >= 0 {
return uint32(v), true
}
case uint64:
return uint32(v), true
case float64:
if v >= 0 {
return uint32(v), true
}
}
return 0, false
}
+53 -53
View File
@@ -1,4 +1,4 @@
package main package store
import ( import (
"encoding/json" "encoding/json"
@@ -14,9 +14,9 @@ import (
// ============================================ // ============================================
// ListLLMProviders 列出所有 LLM Provider // ListLLMProviders 列出所有 LLM Provider
func (s *store) ListLLMProviders(includeInactive bool) ([]llmProviderRecord, error) { func (s *Store) ListLLMProviders(includeInactive bool) ([]LLMProviderRecord, error) {
var rows []llmProviderRecord var rows []LLMProviderRecord
query := s.db.Model(&llmProviderRecord{}) query := s.db.Model(&LLMProviderRecord{})
if !includeInactive { if !includeInactive {
query = query.Where("active = ?", true) query = query.Where("active = ?", true)
} }
@@ -27,8 +27,8 @@ func (s *store) ListLLMProviders(includeInactive bool) ([]llmProviderRecord, err
} }
// GetLLMProvider 获取单个 LLM Provider // GetLLMProvider 获取单个 LLM Provider
func (s *store) GetLLMProvider(name string) (*llmProviderRecord, error) { func (s *Store) GetLLMProvider(name string) (*LLMProviderRecord, error) {
var record llmProviderRecord var record LLMProviderRecord
if err := s.db.Where("name = ?", name).Take(&record).Error; err != nil { if err := s.db.Where("name = ?", name).Take(&record).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, err return nil, err
@@ -39,7 +39,7 @@ func (s *store) GetLLMProvider(name string) (*llmProviderRecord, error) {
} }
// CreateLLMProvider 创建 LLM Provider // CreateLLMProvider 创建 LLM Provider
func (s *store) CreateLLMProvider(record *llmProviderRecord) error { func (s *Store) CreateLLMProvider(record *LLMProviderRecord) error {
if err := s.db.Create(record).Error; err != nil { if err := s.db.Create(record).Error; err != nil {
return fmt.Errorf("create llm provider %s: %w", record.Name, err) return fmt.Errorf("create llm provider %s: %w", record.Name, err)
} }
@@ -47,16 +47,16 @@ func (s *store) CreateLLMProvider(record *llmProviderRecord) error {
} }
// UpdateLLMProvider 更新 LLM Provider // UpdateLLMProvider 更新 LLM Provider
func (s *store) UpdateLLMProvider(name string, updates map[string]any) error { func (s *Store) UpdateLLMProvider(name string, updates map[string]any) error {
if err := s.db.Model(&llmProviderRecord{}).Where("name = ?", name).Updates(updates).Error; err != nil { if err := s.db.Model(&LLMProviderRecord{}).Where("name = ?", name).Updates(updates).Error; err != nil {
return fmt.Errorf("update llm provider %s: %w", name, err) return fmt.Errorf("update llm provider %s: %w", name, err)
} }
return nil return nil
} }
// DeleteLLMProvider 删除 LLM Provider // DeleteLLMProvider 删除 LLM Provider
func (s *store) DeleteLLMProvider(name string) error { func (s *Store) DeleteLLMProvider(name string) error {
if err := s.db.Where("name = ?", name).Delete(&llmProviderRecord{}).Error; err != nil { if err := s.db.Where("name = ?", name).Delete(&LLMProviderRecord{}).Error; err != nil {
return fmt.Errorf("delete llm provider %s: %w", name, err) return fmt.Errorf("delete llm provider %s: %w", name, err)
} }
return nil return nil
@@ -64,7 +64,7 @@ func (s *store) DeleteLLMProvider(name string) error {
// EnsureDefaultLLMProvider 确保存在默认 LLM Provider 配置 // EnsureDefaultLLMProvider 确保存在默认 LLM Provider 配置
// 只有当数据库中完全没有任何 provider 配置时,才创建默认配置 // 只有当数据库中完全没有任何 provider 配置时,才创建默认配置
func (s *store) EnsureDefaultLLMProvider() error { func (s *Store) EnsureDefaultLLMProvider() error {
// 先检查是否已经有任何 provider 配置 // 先检查是否已经有任何 provider 配置
providers, err := s.ListLLMProviders(true) providers, err := s.ListLLMProviders(true)
if err != nil { if err != nil {
@@ -74,7 +74,7 @@ func (s *store) EnsureDefaultLLMProvider() error {
return nil // 已有配置,不创建默认 return nil // 已有配置,不创建默认
} }
// 创建默认配置 // 创建默认配置
defaultConfig := &llmProviderRecord{ defaultConfig := &LLMProviderRecord{
Name: "default", Name: "default",
Active: true, Active: true,
APIKey: "", APIKey: "",
@@ -91,8 +91,8 @@ func (s *store) EnsureDefaultLLMProvider() error {
// ============================================ // ============================================
// GetLLMToolRouter 获取当前激活的 Tool Router 配置 // GetLLMToolRouter 获取当前激活的 Tool Router 配置
func (s *store) GetLLMToolRouter() (*llmToolRouterRecord, error) { func (s *Store) GetLLMToolRouter() (*LLMToolRouterRecord, error) {
var record llmToolRouterRecord var record LLMToolRouterRecord
// 默认取第一条记录(ID 最小的),因为通常只需要一个配置 // 默认取第一条记录(ID 最小的),因为通常只需要一个配置
if err := s.db.Order("id ASC").First(&record).Error; err != nil { if err := s.db.Order("id ASC").First(&record).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
@@ -104,7 +104,7 @@ func (s *store) GetLLMToolRouter() (*llmToolRouterRecord, error) {
} }
// CreateLLMToolRouter 创建 Tool Router 配置 // CreateLLMToolRouter 创建 Tool Router 配置
func (s *store) CreateLLMToolRouter(record *llmToolRouterRecord) error { func (s *Store) CreateLLMToolRouter(record *LLMToolRouterRecord) error {
if err := s.db.Create(record).Error; err != nil { if err := s.db.Create(record).Error; err != nil {
return fmt.Errorf("create llm tool router: %w", err) return fmt.Errorf("create llm tool router: %w", err)
} }
@@ -112,15 +112,15 @@ func (s *store) CreateLLMToolRouter(record *llmToolRouterRecord) error {
} }
// UpdateLLMToolRouter 更新 Tool Router 配置 // UpdateLLMToolRouter 更新 Tool Router 配置
func (s *store) UpdateLLMToolRouter(id uint64, updates map[string]any) error { func (s *Store) UpdateLLMToolRouter(id uint64, updates map[string]any) error {
if err := s.db.Model(&llmToolRouterRecord{}).Where("id = ?", id).Updates(updates).Error; err != nil { if err := s.db.Model(&LLMToolRouterRecord{}).Where("id = ?", id).Updates(updates).Error; err != nil {
return fmt.Errorf("update llm tool router %d: %w", id, err) return fmt.Errorf("update llm tool router %d: %w", id, err)
} }
return nil return nil
} }
// EnsureDefaultLLMToolRouter 确保存在默认 Tool Router 配置 // EnsureDefaultLLMToolRouter 确保存在默认 Tool Router 配置
func (s *store) EnsureDefaultLLMToolRouter() error { func (s *Store) EnsureDefaultLLMToolRouter() error {
_, err := s.GetLLMToolRouter() _, err := s.GetLLMToolRouter()
if err == nil { if err == nil {
return nil // 已存在 return nil // 已存在
@@ -129,7 +129,7 @@ func (s *store) EnsureDefaultLLMToolRouter() error {
return err return err
} }
// 创建默认配置 // 创建默认配置
defaultConfig := &llmToolRouterRecord{ defaultConfig := &LLMToolRouterRecord{
Enabled: true, Enabled: true,
OpenAIName: "", OpenAIName: "",
Timeout: 30, Timeout: 30,
@@ -144,8 +144,8 @@ func (s *store) EnsureDefaultLLMToolRouter() error {
// ============================================ // ============================================
// GetLLMPrimaryConfig 获取当前激活的主 AI 回复配置 // GetLLMPrimaryConfig 获取当前激活的主 AI 回复配置
func (s *store) GetLLMPrimaryConfig() (*llmPrimaryConfigRecord, error) { func (s *Store) GetLLMPrimaryConfig() (*LLMPrimaryConfigRecord, error) {
var record llmPrimaryConfigRecord var record LLMPrimaryConfigRecord
// 默认取第一条记录(ID 最小的),因为通常只需要一个配置 // 默认取第一条记录(ID 最小的),因为通常只需要一个配置
if err := s.db.Order("id ASC").First(&record).Error; err != nil { if err := s.db.Order("id ASC").First(&record).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
@@ -158,7 +158,7 @@ func (s *store) GetLLMPrimaryConfig() (*llmPrimaryConfigRecord, error) {
// GetLLMPrimaryConfigSystemPrompt 获取主 AI 回复配置中的系统提示词 // GetLLMPrimaryConfigSystemPrompt 获取主 AI 回复配置中的系统提示词
// 如果没有配置或出错,返回空字符串(autoreply service 会处理这种情况) // 如果没有配置或出错,返回空字符串(autoreply service 会处理这种情况)
func (s *store) GetLLMPrimaryConfigSystemPrompt() (string, error) { func (s *Store) GetLLMPrimaryConfigSystemPrompt() (string, error) {
record, err := s.GetLLMPrimaryConfig() record, err := s.GetLLMPrimaryConfig()
if err != nil { if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
@@ -171,7 +171,7 @@ func (s *store) GetLLMPrimaryConfigSystemPrompt() (string, error) {
} }
// GetLLMPrimaryConfigEnableTool 获取是否启用工具调用 // GetLLMPrimaryConfigEnableTool 获取是否启用工具调用
func (s *store) GetLLMPrimaryConfigEnableTool() (bool, error) { func (s *Store) GetLLMPrimaryConfigEnableTool() (bool, error) {
record, err := s.GetLLMPrimaryConfig() record, err := s.GetLLMPrimaryConfig()
if err != nil { if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
@@ -183,7 +183,7 @@ func (s *store) GetLLMPrimaryConfigEnableTool() (bool, error) {
} }
// CreateLLMPrimaryConfig 创建主 AI 回复配置 // CreateLLMPrimaryConfig 创建主 AI 回复配置
func (s *store) CreateLLMPrimaryConfig(record *llmPrimaryConfigRecord) error { func (s *Store) CreateLLMPrimaryConfig(record *LLMPrimaryConfigRecord) error {
if err := s.db.Create(record).Error; err != nil { if err := s.db.Create(record).Error; err != nil {
return fmt.Errorf("create llm primary config: %w", err) return fmt.Errorf("create llm primary config: %w", err)
} }
@@ -191,15 +191,15 @@ func (s *store) CreateLLMPrimaryConfig(record *llmPrimaryConfigRecord) error {
} }
// UpdateLLMPrimaryConfig 更新主 AI 回复配置 // UpdateLLMPrimaryConfig 更新主 AI 回复配置
func (s *store) UpdateLLMPrimaryConfig(id uint64, updates map[string]any) error { func (s *Store) UpdateLLMPrimaryConfig(id uint64, updates map[string]any) error {
if err := s.db.Model(&llmPrimaryConfigRecord{}).Where("id = ?", id).Updates(updates).Error; err != nil { if err := s.db.Model(&LLMPrimaryConfigRecord{}).Where("id = ?", id).Updates(updates).Error; err != nil {
return fmt.Errorf("update llm primary config %d: %w", id, err) return fmt.Errorf("update llm primary config %d: %w", id, err)
} }
return nil return nil
} }
// EnsureDefaultLLMPrimaryConfig 确保存在默认主 AI 回复配置 // EnsureDefaultLLMPrimaryConfig 确保存在默认主 AI 回复配置
func (s *store) EnsureDefaultLLMPrimaryConfig() error { func (s *Store) EnsureDefaultLLMPrimaryConfig() error {
_, err := s.GetLLMPrimaryConfig() _, err := s.GetLLMPrimaryConfig()
if err == nil { if err == nil {
return nil // 已存在 return nil // 已存在
@@ -208,7 +208,7 @@ func (s *store) EnsureDefaultLLMPrimaryConfig() error {
return err return err
} }
// 创建默认配置 // 创建默认配置
defaultConfig := &llmPrimaryConfigRecord{ defaultConfig := &LLMPrimaryConfigRecord{
Enabled: false, Enabled: false,
ProviderName: "", ProviderName: "",
Timeout: 120, Timeout: 120,
@@ -237,7 +237,7 @@ type LLMMessageQueueInput struct {
} }
// EnqueueLLMMessage 将消息添加到 LLM 队列 // EnqueueLLMMessage 将消息添加到 LLM 队列
func (s *store) EnqueueLLMMessage(input LLMMessageQueueInput) (*llmMessageQueueRecord, error) { func (s *Store) EnqueueLLMMessage(input LLMMessageQueueInput) (*LLMMessageQueueRecord, error) {
var err error var err error
if input.BotID == 0 { if input.BotID == 0 {
@@ -269,16 +269,16 @@ func (s *store) EnqueueLLMMessage(input LLMMessageQueueInput) (*llmMessageQueueR
// packet_id > 0: 用 bot_id + packet_id 去重(频道消息) // packet_id > 0: 用 bot_id + packet_id 去重(频道消息)
// packet_id = 0: 用 bot_id + from_node_id + text 去重(私聊消息,可能没有 packet_id) // packet_id = 0: 用 bot_id + from_node_id + text 去重(私聊消息,可能没有 packet_id)
// 只排除 pending/processing 状态的消息,允许 error 状态的消息重新入队 // 只排除 pending/processing 状态的消息,允许 error 状态的消息重新入队
var existing llmMessageQueueRecord var existing LLMMessageQueueRecord
if input.PacketID > 0 { if input.PacketID > 0 {
// 频道消息:用 bot_id + packet_id 去重 // 频道消息:用 bot_id + packet_id 去重
err = s.db.Where("bot_id = ? AND packet_id = ? AND deleted_at IS NULL AND status IN (?, ?)", err = s.db.Where("bot_id = ? AND packet_id = ? AND deleted_at IS NULL AND status IN (?, ?)",
input.BotID, input.PacketID, llmMessageStatusPending, llmMessageStatusProcessing). input.BotID, input.PacketID, LLMMessageStatusPending, LLMMessageStatusProcessing).
Take(&existing).Error Take(&existing).Error
} else { } else {
// 私聊消息:用 bot_id + from_node_id + text 去重(避免同一人连续发相同内容被拒绝) // 私聊消息:用 bot_id + from_node_id + text 去重(避免同一人连续发相同内容被拒绝)
err = s.db.Where("bot_id = ? AND from_node_id = ? AND text = ? AND deleted_at IS NULL AND status IN (?, ?)", err = s.db.Where("bot_id = ? AND from_node_id = ? AND text = ? AND deleted_at IS NULL AND status IN (?, ?)",
input.BotID, input.FromNodeID, input.Text, llmMessageStatusPending, llmMessageStatusProcessing). input.BotID, input.FromNodeID, input.Text, LLMMessageStatusPending, LLMMessageStatusProcessing).
Take(&existing).Error Take(&existing).Error
} }
if err == nil { if err == nil {
@@ -294,7 +294,7 @@ func (s *store) EnqueueLLMMessage(input LLMMessageQueueInput) (*llmMessageQueueR
if messageType == "" { if messageType == "" {
messageType = "direct" messageType = "direct"
} }
record := &llmMessageQueueRecord{ record := &LLMMessageQueueRecord{
BotID: input.BotID, BotID: input.BotID,
BotNodeID: input.BotNodeID, BotNodeID: input.BotNodeID,
BotNodeNum: input.BotNodeNum, BotNodeNum: input.BotNodeNum,
@@ -307,7 +307,7 @@ func (s *store) EnqueueLLMMessage(input LLMMessageQueueInput) (*llmMessageQueueR
ChannelID: input.ChannelID, ChannelID: input.ChannelID,
Topic: input.Topic, Topic: input.Topic,
MessageType: messageType, MessageType: messageType,
Status: llmMessageStatusPending, Status: LLMMessageStatusPending,
ReceivedAt: now, ReceivedAt: now,
ContentJSON: input.ContentJSON, ContentJSON: input.ContentJSON,
} }
@@ -319,9 +319,9 @@ func (s *store) EnqueueLLMMessage(input LLMMessageQueueInput) (*llmMessageQueueR
} }
// ListLLMMessages 列出 LLM 队列消息 // ListLLMMessages 列出 LLM 队列消息
func (s *store) ListLLMMessages(opts listOptions, botID uint64, includeDeleted bool) ([]llmMessageQueueRecord, int64, error) { func (s *Store) ListLLMMessages(opts ListOptions, botID uint64, includeDeleted bool) ([]LLMMessageQueueRecord, int64, error) {
var rows []llmMessageQueueRecord var rows []LLMMessageQueueRecord
query := s.db.Model(&llmMessageQueueRecord{}) query := s.db.Model(&LLMMessageQueueRecord{})
if botID > 0 { if botID > 0 {
query = query.Where("bot_id = ?", botID) query = query.Where("bot_id = ?", botID)
@@ -352,8 +352,8 @@ func (s *store) ListLLMMessages(opts listOptions, botID uint64, includeDeleted b
} }
// GetLLMMessage 获取单条 LLM 消息 // GetLLMMessage 获取单条 LLM 消息
func (s *store) GetLLMMessage(id uint64) (*llmMessageQueueRecord, error) { func (s *Store) GetLLMMessage(id uint64) (*LLMMessageQueueRecord, error) {
var record llmMessageQueueRecord var record LLMMessageQueueRecord
if err := s.db.Where("id = ?", id).Take(&record).Error; err != nil { if err := s.db.Where("id = ?", id).Take(&record).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, err return nil, err
@@ -364,50 +364,50 @@ func (s *store) GetLLMMessage(id uint64) (*llmMessageQueueRecord, error) {
} }
// UpdateLLMMessageStatus 更新 LLM 消息状态 // UpdateLLMMessageStatus 更新 LLM 消息状态
func (s *store) UpdateLLMMessageStatus(id uint64, status string, errorMsg string) error { func (s *Store) UpdateLLMMessageStatus(id uint64, status string, errorMsg string) error {
updates := map[string]any{ updates := map[string]any{
"status": status, "status": status,
"error": errorMsg, "error": errorMsg,
} }
if status == llmMessageStatusProcessed { if status == LLMMessageStatusProcessed {
now := time.Now() now := time.Now()
updates["processed_at"] = &now updates["processed_at"] = &now
} }
if err := s.db.Model(&llmMessageQueueRecord{}).Where("id = ?", id).Updates(updates).Error; err != nil { if err := s.db.Model(&LLMMessageQueueRecord{}).Where("id = ?", id).Updates(updates).Error; err != nil {
return fmt.Errorf("update llm message status %d: %w", id, err) return fmt.Errorf("update llm message status %d: %w", id, err)
} }
return nil return nil
} }
// SoftDeleteLLMMessage 软删除 LLM 消息 // SoftDeleteLLMMessage 软删除 LLM 消息
func (s *store) SoftDeleteLLMMessage(id uint64) error { func (s *Store) SoftDeleteLLMMessage(id uint64) error {
now := time.Now() now := time.Now()
if err := s.db.Model(&llmMessageQueueRecord{}).Where("id = ?", id).Update("deleted_at", &now).Error; err != nil { if err := s.db.Model(&LLMMessageQueueRecord{}).Where("id = ?", id).Update("deleted_at", &now).Error; err != nil {
return fmt.Errorf("soft delete llm message %d: %w", id, err) return fmt.Errorf("soft delete llm message %d: %w", id, err)
} }
return nil return nil
} }
// SoftDeleteLLMMessagesByBot 软删除指定机器人的所有消息 // SoftDeleteLLMMessagesByBot 软删除指定机器人的所有消息
func (s *store) SoftDeleteLLMMessagesByBot(botID uint64) error { func (s *Store) SoftDeleteLLMMessagesByBot(botID uint64) error {
now := time.Now() now := time.Now()
if err := s.db.Model(&llmMessageQueueRecord{}).Where("bot_id = ? AND deleted_at IS NULL", botID).Update("deleted_at", &now).Error; err != nil { if err := s.db.Model(&LLMMessageQueueRecord{}).Where("bot_id = ? AND deleted_at IS NULL", botID).Update("deleted_at", &now).Error; err != nil {
return fmt.Errorf("soft delete llm messages for bot %d: %w", botID, err) return fmt.Errorf("soft delete llm messages for bot %d: %w", botID, err)
} }
return nil return nil
} }
// CleanupDeletedLLMMessages 清理已软删除超过指定时间的消息 // CleanupDeletedLLMMessages 清理已软删除超过指定时间的消息
func (s *store) CleanupDeletedLLMMessages(before time.Time) (int64, error) { func (s *Store) CleanupDeletedLLMMessages(before time.Time) (int64, error) {
result := s.db.Where("deleted_at IS NOT NULL AND deleted_at < ?", before).Delete(&llmMessageQueueRecord{}) result := s.db.Where("deleted_at IS NOT NULL AND deleted_at < ?", before).Delete(&LLMMessageQueueRecord{})
if result.Error != nil { if result.Error != nil {
return 0, fmt.Errorf("cleanup deleted llm messages: %w", result.Error) return 0, fmt.Errorf("cleanup deleted llm messages: %w", result.Error)
} }
return result.RowsAffected, nil return result.RowsAffected, nil
} }
// llmMessageDTO 将数据库记录转换为 API 响应格式 // LLMMessageDTO 将数据库记录转换为 API 响应格式
func llmMessageDTO(row llmMessageQueueRecord) map[string]any { func LLMMessageDTO(row LLMMessageQueueRecord) map[string]any {
return map[string]any{ return map[string]any{
"id": row.ID, "id": row.ID,
"bot_id": row.BotID, "bot_id": row.BotID,
@@ -432,7 +432,7 @@ func llmMessageDTO(row llmMessageQueueRecord) map[string]any {
// enqueueChannelMessageToLLM 将频道消息添加到 LLM 队列 // enqueueChannelMessageToLLM 将频道消息添加到 LLM 队列
// 为每个启用了「包含频道消息」的机器人都创建一条独立的队列记录 // 为每个启用了「包含频道消息」的机器人都创建一条独立的队列记录
func enqueueChannelMessageToLLM(s *store, record map[string]any) error { func enqueueChannelMessageToLLM(s *Store, record map[string]any) error {
if s == nil { if s == nil {
return nil return nil
} }
@@ -481,7 +481,7 @@ func enqueueChannelMessageToLLM(s *store, record map[string]any) error {
// 查询所有启用了 LLM 队列且包含频道消息的机器人 // 查询所有启用了 LLM 队列且包含频道消息的机器人
// SQLite 中 numeric 布尔值用 1/0 存储,必须用整数查询 // SQLite 中 numeric 布尔值用 1/0 存储,必须用整数查询
var bots []botNodeRecord var bots []BotNodeRecord
err = s.db.Where("llm_queue_enabled = ? AND llm_include_channel_messages = ?", 1, 1).Find(&bots).Error err = s.db.Where("llm_queue_enabled = ? AND llm_include_channel_messages = ?", 1, 1).Find(&bots).Error
if err != nil { if err != nil {
return fmt.Errorf("query bots for channel message enqueue: %w", err) return fmt.Errorf("query bots for channel message enqueue: %w", err)
@@ -1,12 +1,12 @@
package main package store
func (s *store) InsertLoginLog(log loginLogRecord) error { func (s *Store) InsertLoginLog(log LoginLogRecord) error {
return s.db.Create(&log).Error return s.db.Create(&log).Error
} }
func (s *store) ListLoginLogs(opts listOptions) ([]loginLogRecord, error) { func (s *Store) ListLoginLogs(opts ListOptions) ([]LoginLogRecord, error) {
opts = normalizeListOptions(opts) opts = NormalizeListOptions(opts)
var rows []loginLogRecord var rows []LoginLogRecord
q := s.db.Order("created_at DESC").Order("id DESC").Limit(opts.Limit).Offset(opts.Offset) q := s.db.Order("created_at DESC").Order("id DESC").Limit(opts.Limit).Offset(opts.Offset)
if opts.Since != nil { if opts.Since != nil {
q = q.Where("created_at >= ?", *opts.Since) q = q.Where("created_at >= ?", *opts.Since)
@@ -1,4 +1,4 @@
package main package store
import ( import (
"crypto/sha256" "crypto/sha256"
@@ -22,13 +22,13 @@ const (
) )
var ( var (
errMapTileSourceAlreadyExists = errors.New("map source already exists") ErrMapTileSourceAlreadyExists = errors.New("map source already exists")
errMapTileSourceCannotDeleteDefault = errors.New("default map source cannot be deleted") ErrMapTileSourceCannotDeleteDefault = errors.New("default map source cannot be deleted")
errMapTileSourceCannotDisableDefault = errors.New("default map source cannot be disabled") ErrMapTileSourceCannotDisableDefault = errors.New("default map source cannot be disabled")
errMapTileSourceDefaultMustBeEnabled = errors.New("default map source must be enabled") ErrMapTileSourceDefaultMustBeEnabled = errors.New("default map source must be enabled")
) )
type mapTileSourceInput struct { type MapTileSourceInput struct {
Name string Name string
URLTemplate string URLTemplate string
Attribution string Attribution string
@@ -38,10 +38,10 @@ type mapTileSourceInput struct {
ProxyEnabled bool ProxyEnabled bool
} }
func (s *store) ListMapTileSources(opts listOptions) ([]mapTileSourceRecord, error) { func (s *Store) ListMapTileSources(opts ListOptions) ([]MapTileSourceRecord, error) {
opts = normalizeListOptions(opts) opts = NormalizeListOptions(opts)
var rows []mapTileSourceRecord var rows []MapTileSourceRecord
q := s.db.Model(&mapTileSourceRecord{}). q := s.db.Model(&MapTileSourceRecord{}).
Order("is_default DESC"). Order("is_default DESC").
Order("updated_at DESC"). Order("updated_at DESC").
Order("id DESC"). Order("id DESC").
@@ -50,14 +50,14 @@ func (s *store) ListMapTileSources(opts listOptions) ([]mapTileSourceRecord, err
return rows, q.Find(&rows).Error return rows, q.Find(&rows).Error
} }
func (s *store) CountMapTileSources(opts listOptions) (int64, error) { func (s *Store) CountMapTileSources(opts ListOptions) (int64, error) {
var total int64 var total int64
return total, s.db.Model(&mapTileSourceRecord{}).Count(&total).Error return total, s.db.Model(&MapTileSourceRecord{}).Count(&total).Error
} }
func (s *store) ListEnabledMapTileSources() ([]mapTileSourceRecord, error) { func (s *Store) ListEnabledMapTileSources() ([]MapTileSourceRecord, error) {
var rows []mapTileSourceRecord var rows []MapTileSourceRecord
if err := s.db.Model(&mapTileSourceRecord{}). if err := s.db.Model(&MapTileSourceRecord{}).
Where("enabled = ?", true). Where("enabled = ?", true).
Order("is_default DESC"). Order("is_default DESC").
Order("updated_at DESC"). Order("updated_at DESC").
@@ -66,13 +66,13 @@ func (s *store) ListEnabledMapTileSources() ([]mapTileSourceRecord, error) {
return nil, err return nil, err
} }
if len(rows) == 0 { if len(rows) == 0 {
return []mapTileSourceRecord{defaultMapTileSourceRecord()}, nil return []MapTileSourceRecord{defaultMapTileSourceRecord()}, nil
} }
return rows, nil return rows, nil
} }
func (s *store) GetDefaultMapTileSource() (*mapTileSourceRecord, error) { func (s *Store) GetDefaultMapTileSource() (*MapTileSourceRecord, error) {
var row mapTileSourceRecord var row MapTileSourceRecord
err := s.db.Where("enabled = ? AND is_default = ?", true, true).Order("id ASC").Take(&row).Error err := s.db.Where("enabled = ? AND is_default = ?", true, true).Order("id ASC").Take(&row).Error
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
fallback := defaultMapTileSourceRecord() fallback := defaultMapTileSourceRecord()
@@ -84,28 +84,28 @@ func (s *store) GetDefaultMapTileSource() (*mapTileSourceRecord, error) {
return &row, nil return &row, nil
} }
func (s *store) GetEnabledMapTileSourceByHash(hash string) (*mapTileSourceRecord, error) { func (s *Store) GetEnabledMapTileSourceByHash(hash string) (*MapTileSourceRecord, error) {
var row mapTileSourceRecord var row MapTileSourceRecord
if err := s.db.Where("enabled = ? AND proxy_enabled = ? AND url_template_hash = ?", true, true, hash).Take(&row).Error; err != nil { if err := s.db.Where("enabled = ? AND proxy_enabled = ? AND url_template_hash = ?", true, true, hash).Take(&row).Error; err != nil {
return nil, err return nil, err
} }
return &row, nil return &row, nil
} }
func (s *store) CreateMapTileSource(input mapTileSourceInput) (*mapTileSourceRecord, error) { func (s *Store) CreateMapTileSource(input MapTileSourceInput) (*MapTileSourceRecord, error) {
row, err := mapTileSourceFromInput(input) row, err := mapTileSourceFromInput(input)
if err != nil { if err != nil {
return nil, err return nil, err
} }
if row.IsDefault && !row.Enabled { if row.IsDefault && !row.Enabled {
return nil, errMapTileSourceDefaultMustBeEnabled return nil, ErrMapTileSourceDefaultMustBeEnabled
} }
if err := s.ensureMapTileSourceUnique(0, row.Name, row.URLTemplate); err != nil { if err := s.ensureMapTileSourceUnique(0, row.Name, row.URLTemplate); err != nil {
return nil, err return nil, err
} }
if err := s.db.Transaction(func(tx *gorm.DB) error { if err := s.db.Transaction(func(tx *gorm.DB) error {
if row.IsDefault { if row.IsDefault {
if err := tx.Model(&mapTileSourceRecord{}).Where("is_default = ?", true).Update("is_default", false).Error; err != nil { if err := tx.Model(&MapTileSourceRecord{}).Where("is_default = ?", true).Update("is_default", false).Error; err != nil {
return err return err
} }
} }
@@ -116,7 +116,7 @@ func (s *store) CreateMapTileSource(input mapTileSourceInput) (*mapTileSourceRec
return row, nil return row, nil
} }
func (s *store) UpdateMapTileSource(id uint64, input mapTileSourceInput) (*mapTileSourceRecord, error) { func (s *Store) UpdateMapTileSource(id uint64, input MapTileSourceInput) (*MapTileSourceRecord, error) {
if id == 0 { if id == 0 {
return nil, fmt.Errorf("map source id is required") return nil, fmt.Errorf("map source id is required")
} }
@@ -124,17 +124,17 @@ func (s *store) UpdateMapTileSource(id uint64, input mapTileSourceInput) (*mapTi
if err != nil { if err != nil {
return nil, err return nil, err
} }
var updated mapTileSourceRecord var updated MapTileSourceRecord
if err := s.db.Transaction(func(tx *gorm.DB) error { if err := s.db.Transaction(func(tx *gorm.DB) error {
var existing mapTileSourceRecord var existing MapTileSourceRecord
if err := tx.Where("id = ?", id).Take(&existing).Error; err != nil { if err := tx.Where("id = ?", id).Take(&existing).Error; err != nil {
return err return err
} }
if existing.IsDefault && !row.Enabled { if existing.IsDefault && !row.Enabled {
return errMapTileSourceCannotDisableDefault return ErrMapTileSourceCannotDisableDefault
} }
if row.IsDefault && !row.Enabled { if row.IsDefault && !row.Enabled {
return errMapTileSourceDefaultMustBeEnabled return ErrMapTileSourceDefaultMustBeEnabled
} }
if !row.IsDefault && existing.IsDefault { if !row.IsDefault && existing.IsDefault {
row.IsDefault = true row.IsDefault = true
@@ -143,7 +143,7 @@ func (s *store) UpdateMapTileSource(id uint64, input mapTileSourceInput) (*mapTi
return err return err
} }
if row.IsDefault { if row.IsDefault {
if err := tx.Model(&mapTileSourceRecord{}).Where("id <> ? AND is_default = ?", id, true).Update("is_default", false).Error; err != nil { if err := tx.Model(&MapTileSourceRecord{}).Where("id <> ? AND is_default = ?", id, true).Update("is_default", false).Error; err != nil {
return err return err
} }
} }
@@ -158,7 +158,7 @@ func (s *store) UpdateMapTileSource(id uint64, input mapTileSourceInput) (*mapTi
"proxy_enabled": row.ProxyEnabled, "proxy_enabled": row.ProxyEnabled,
"updated_at": time.Now(), "updated_at": time.Now(),
} }
if err := tx.Model(&mapTileSourceRecord{}).Where("id = ?", id).Updates(updates).Error; err != nil { if err := tx.Model(&MapTileSourceRecord{}).Where("id = ?", id).Updates(updates).Error; err != nil {
return err return err
} }
return tx.Where("id = ?", id).Take(&updated).Error return tx.Where("id = ?", id).Take(&updated).Error
@@ -168,19 +168,19 @@ func (s *store) UpdateMapTileSource(id uint64, input mapTileSourceInput) (*mapTi
return &updated, nil return &updated, nil
} }
func (s *store) DeleteMapTileSource(id uint64) error { func (s *Store) DeleteMapTileSource(id uint64) error {
if id == 0 { if id == 0 {
return fmt.Errorf("map source id is required") return fmt.Errorf("map source id is required")
} }
return s.db.Transaction(func(tx *gorm.DB) error { return s.db.Transaction(func(tx *gorm.DB) error {
var row mapTileSourceRecord var row MapTileSourceRecord
if err := tx.Where("id = ?", id).Take(&row).Error; err != nil { if err := tx.Where("id = ?", id).Take(&row).Error; err != nil {
return err return err
} }
if row.IsDefault { if row.IsDefault {
return errMapTileSourceCannotDeleteDefault return ErrMapTileSourceCannotDeleteDefault
} }
result := tx.Where("id = ?", id).Delete(&mapTileSourceRecord{}) result := tx.Where("id = ?", id).Delete(&MapTileSourceRecord{})
if result.Error != nil { if result.Error != nil {
return result.Error return result.Error
} }
@@ -191,22 +191,22 @@ func (s *store) DeleteMapTileSource(id uint64) error {
}) })
} }
func (s *store) SetDefaultMapTileSource(id uint64) (*mapTileSourceRecord, error) { func (s *Store) SetDefaultMapTileSource(id uint64) (*MapTileSourceRecord, error) {
if id == 0 { if id == 0 {
return nil, fmt.Errorf("map source id is required") return nil, fmt.Errorf("map source id is required")
} }
var row mapTileSourceRecord var row MapTileSourceRecord
if err := s.db.Transaction(func(tx *gorm.DB) error { if err := s.db.Transaction(func(tx *gorm.DB) error {
if err := tx.Where("id = ?", id).Take(&row).Error; err != nil { if err := tx.Where("id = ?", id).Take(&row).Error; err != nil {
return err return err
} }
if !row.Enabled { if !row.Enabled {
return errMapTileSourceDefaultMustBeEnabled return ErrMapTileSourceDefaultMustBeEnabled
} }
if err := tx.Model(&mapTileSourceRecord{}).Where("is_default = ?", true).Update("is_default", false).Error; err != nil { if err := tx.Model(&MapTileSourceRecord{}).Where("is_default = ?", true).Update("is_default", false).Error; err != nil {
return err return err
} }
if err := tx.Model(&mapTileSourceRecord{}).Where("id = ?", id).Updates(map[string]any{"is_default": true, "updated_at": time.Now()}).Error; err != nil { if err := tx.Model(&MapTileSourceRecord{}).Where("id = ?", id).Updates(map[string]any{"is_default": true, "updated_at": time.Now()}).Error; err != nil {
return err return err
} }
return tx.Where("id = ?", id).Take(&row).Error return tx.Where("id = ?", id).Take(&row).Error
@@ -216,10 +216,10 @@ func (s *store) SetDefaultMapTileSource(id uint64) (*mapTileSourceRecord, error)
return &row, nil return &row, nil
} }
func (s *store) EnsureDefaultMapTileSource() error { func (s *Store) EnsureDefaultMapTileSource() error {
return s.db.Transaction(func(tx *gorm.DB) error { return s.db.Transaction(func(tx *gorm.DB) error {
var count int64 var count int64
if err := tx.Model(&mapTileSourceRecord{}).Count(&count).Error; err != nil { if err := tx.Model(&MapTileSourceRecord{}).Count(&count).Error; err != nil {
return err return err
} }
if count == 0 { if count == 0 {
@@ -227,28 +227,28 @@ func (s *store) EnsureDefaultMapTileSource() error {
return tx.Create(&row).Error return tx.Create(&row).Error
} }
var defaults []mapTileSourceRecord var defaults []MapTileSourceRecord
if err := tx.Where("enabled = ? AND is_default = ?", true, true).Order("id ASC").Find(&defaults).Error; err != nil { if err := tx.Where("enabled = ? AND is_default = ?", true, true).Order("id ASC").Find(&defaults).Error; err != nil {
return err return err
} }
if len(defaults) > 0 { if len(defaults) > 0 {
return tx.Model(&mapTileSourceRecord{}).Where("id <> ? AND is_default = ?", defaults[0].ID, true).Update("is_default", false).Error return tx.Model(&MapTileSourceRecord{}).Where("id <> ? AND is_default = ?", defaults[0].ID, true).Update("is_default", false).Error
} }
var enabled mapTileSourceRecord var enabled MapTileSourceRecord
err := tx.Where("enabled = ?", true).Order("id ASC").Take(&enabled).Error err := tx.Where("enabled = ?", true).Order("id ASC").Take(&enabled).Error
if err == nil { if err == nil {
return tx.Model(&mapTileSourceRecord{}).Where("id = ?", enabled.ID).Updates(map[string]any{"is_default": true, "updated_at": time.Now()}).Error return tx.Model(&MapTileSourceRecord{}).Where("id = ?", enabled.ID).Updates(map[string]any{"is_default": true, "updated_at": time.Now()}).Error
} }
if !errors.Is(err, gorm.ErrRecordNotFound) { if !errors.Is(err, gorm.ErrRecordNotFound) {
return err return err
} }
row := defaultMapTileSourceRecord() row := defaultMapTileSourceRecord()
var existing mapTileSourceRecord var existing MapTileSourceRecord
err = tx.Where("name = ? OR url_template = ?", row.Name, row.URLTemplate).Order("id ASC").Take(&existing).Error err = tx.Where("name = ? OR url_template = ?", row.Name, row.URLTemplate).Order("id ASC").Take(&existing).Error
if err == nil { if err == nil {
return tx.Model(&mapTileSourceRecord{}).Where("id = ?", existing.ID).Updates(map[string]any{"enabled": true, "is_default": true, "updated_at": time.Now()}).Error return tx.Model(&MapTileSourceRecord{}).Where("id = ?", existing.ID).Updates(map[string]any{"enabled": true, "is_default": true, "updated_at": time.Now()}).Error
} }
if !errors.Is(err, gorm.ErrRecordNotFound) { if !errors.Is(err, gorm.ErrRecordNotFound) {
return err return err
@@ -257,16 +257,16 @@ func (s *store) EnsureDefaultMapTileSource() error {
}) })
} }
func mapTileSourceHash(urlTemplate string) string { func MapTileSourceHash(urlTemplate string) string {
h := sha256.Sum256([]byte(urlTemplate)) h := sha256.Sum256([]byte(urlTemplate))
return hex.EncodeToString(h[:]) return hex.EncodeToString(h[:])
} }
func defaultMapTileSourceRecord() mapTileSourceRecord { func defaultMapTileSourceRecord() MapTileSourceRecord {
return mapTileSourceRecord{ return MapTileSourceRecord{
Name: defaultMapTileSourceName, Name: defaultMapTileSourceName,
URLTemplate: defaultMapTileSourceURLTemplate, URLTemplate: defaultMapTileSourceURLTemplate,
URLTemplateHash: mapTileSourceHash(defaultMapTileSourceURLTemplate), URLTemplateHash: MapTileSourceHash(defaultMapTileSourceURLTemplate),
Attribution: defaultMapTileSourceAttribution, Attribution: defaultMapTileSourceAttribution,
MaxZoom: defaultMapTileSourceMaxZoom, MaxZoom: defaultMapTileSourceMaxZoom,
Enabled: true, Enabled: true,
@@ -275,7 +275,7 @@ func defaultMapTileSourceRecord() mapTileSourceRecord {
} }
} }
func mapTileSourceFromInput(input mapTileSourceInput) (*mapTileSourceRecord, error) { func mapTileSourceFromInput(input MapTileSourceInput) (*MapTileSourceRecord, error) {
name := strings.TrimSpace(input.Name) name := strings.TrimSpace(input.Name)
if name == "" { if name == "" {
return nil, fmt.Errorf("map source name is required") return nil, fmt.Errorf("map source name is required")
@@ -291,10 +291,10 @@ func mapTileSourceFromInput(input mapTileSourceInput) (*mapTileSourceRecord, err
if maxZoom < 1 || maxZoom > 30 { if maxZoom < 1 || maxZoom > 30 {
return nil, fmt.Errorf("max zoom must be between 1 and 30") return nil, fmt.Errorf("max zoom must be between 1 and 30")
} }
return &mapTileSourceRecord{ return &MapTileSourceRecord{
Name: name, Name: name,
URLTemplate: urlTemplate, URLTemplate: urlTemplate,
URLTemplateHash: mapTileSourceHash(urlTemplate), URLTemplateHash: MapTileSourceHash(urlTemplate),
Attribution: strings.TrimSpace(input.Attribution), Attribution: strings.TrimSpace(input.Attribution),
MaxZoom: maxZoom, MaxZoom: maxZoom,
Enabled: input.Enabled, Enabled: input.Enabled,
@@ -337,19 +337,19 @@ func normalizeMapTileSourceURLTemplate(value string) (string, error) {
return value, nil return value, nil
} }
func (s *store) ensureMapTileSourceUnique(id uint64, name, urlTemplate string) error { func (s *Store) ensureMapTileSourceUnique(id uint64, name, urlTemplate string) error {
return ensureMapTileSourceUniqueTx(s.db, id, name, urlTemplate) return ensureMapTileSourceUniqueTx(s.db, id, name, urlTemplate)
} }
func ensureMapTileSourceUniqueTx(tx *gorm.DB, id uint64, name, urlTemplate string) error { func ensureMapTileSourceUniqueTx(tx *gorm.DB, id uint64, name, urlTemplate string) error {
var existing mapTileSourceRecord var existing MapTileSourceRecord
q := tx.Where("name = ? OR url_template = ?", name, urlTemplate) q := tx.Where("name = ? OR url_template = ?", name, urlTemplate)
if id != 0 { if id != 0 {
q = q.Where("id <> ?", id) q = q.Where("id <> ?", id)
} }
err := q.Take(&existing).Error err := q.Take(&existing).Error
if err == nil { if err == nil {
return errMapTileSourceAlreadyExists return ErrMapTileSourceAlreadyExists
} }
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
return nil return nil
@@ -1,4 +1,4 @@
package main package store
import ( import (
"errors" "errors"
@@ -25,13 +25,13 @@ func TestCreateMapTileSourceValidation(t *testing.T) {
st := openTestStore(t) st := openTestStore(t)
defer st.Close() defer st.Close()
if _, err := st.CreateMapTileSource(mapTileSourceInput{Name: "bad", URLTemplate: "https://tiles.example.com/{z}/{x}.png", MaxZoom: 19, Enabled: true, ProxyEnabled: true}); err == nil { if _, err := st.CreateMapTileSource(MapTileSourceInput{Name: "bad", URLTemplate: "https://tiles.example.com/{z}/{x}.png", MaxZoom: 19, Enabled: true, ProxyEnabled: true}); err == nil {
t.Fatal("CreateMapTileSource() missing placeholder error = nil, want error") t.Fatal("CreateMapTileSource() missing placeholder error = nil, want error")
} }
if _, err := st.CreateMapTileSource(mapTileSourceInput{Name: "bad", URLTemplate: "javascript:alert(1)/{z}/{x}/{y}", MaxZoom: 19, Enabled: true, ProxyEnabled: true}); err == nil { if _, err := st.CreateMapTileSource(MapTileSourceInput{Name: "bad", URLTemplate: "javascript:alert(1)/{z}/{x}/{y}", MaxZoom: 19, Enabled: true, ProxyEnabled: true}); err == nil {
t.Fatal("CreateMapTileSource() invalid scheme error = nil, want error") t.Fatal("CreateMapTileSource() invalid scheme error = nil, want error")
} }
if _, err := st.CreateMapTileSource(mapTileSourceInput{Name: "bad", URLTemplate: "https://user:pass@tiles.example.com/{z}/{x}/{y}.png", MaxZoom: 19, Enabled: true, ProxyEnabled: true}); err == nil { if _, err := st.CreateMapTileSource(MapTileSourceInput{Name: "bad", URLTemplate: "https://user:pass@tiles.example.com/{z}/{x}/{y}.png", MaxZoom: 19, Enabled: true, ProxyEnabled: true}); err == nil {
t.Fatal("CreateMapTileSource() credentials error = nil, want error") t.Fatal("CreateMapTileSource() credentials error = nil, want error")
} }
} }
@@ -40,11 +40,11 @@ func TestListEnabledMapTileSources(t *testing.T) {
st := openTestStore(t) st := openTestStore(t)
defer st.Close() defer st.Close()
disabled, err := st.CreateMapTileSource(mapTileSourceInput{Name: "Disabled", URLTemplate: "https://disabled.example.com/{z}/{x}/{y}.png", MaxZoom: 18, Enabled: false}) disabled, err := st.CreateMapTileSource(MapTileSourceInput{Name: "Disabled", URLTemplate: "https://disabled.example.com/{z}/{x}/{y}.png", MaxZoom: 18, Enabled: false})
if err != nil { if err != nil {
t.Fatalf("CreateMapTileSource(disabled) error = %v", err) t.Fatalf("CreateMapTileSource(disabled) error = %v", err)
} }
custom, err := st.CreateMapTileSource(mapTileSourceInput{Name: "Custom", URLTemplate: "https://custom.example.com/{z}/{x}/{y}.png", MaxZoom: 18, Enabled: true, ProxyEnabled: true}) custom, err := st.CreateMapTileSource(MapTileSourceInput{Name: "Custom", URLTemplate: "https://custom.example.com/{z}/{x}/{y}.png", MaxZoom: 18, Enabled: true, ProxyEnabled: true})
if err != nil { if err != nil {
t.Fatalf("CreateMapTileSource(custom) error = %v", err) t.Fatalf("CreateMapTileSource(custom) error = %v", err)
} }
@@ -76,15 +76,15 @@ func TestMapTileSourceDuplicateAndDefaultRules(t *testing.T) {
st := openTestStore(t) st := openTestStore(t)
defer st.Close() defer st.Close()
first, err := st.CreateMapTileSource(mapTileSourceInput{Name: "Custom", URLTemplate: "https://tiles.example.com/{z}/{x}/{y}.png", MaxZoom: 18, Enabled: true, ProxyEnabled: true}) first, err := st.CreateMapTileSource(MapTileSourceInput{Name: "Custom", URLTemplate: "https://tiles.example.com/{z}/{x}/{y}.png", MaxZoom: 18, Enabled: true, ProxyEnabled: true})
if err != nil { if err != nil {
t.Fatalf("CreateMapTileSource() error = %v", err) t.Fatalf("CreateMapTileSource() error = %v", err)
} }
if _, err := st.CreateMapTileSource(mapTileSourceInput{Name: "Custom", URLTemplate: "https://tiles2.example.com/{z}/{x}/{y}.png", MaxZoom: 18, Enabled: true, ProxyEnabled: true}); !errors.Is(err, errMapTileSourceAlreadyExists) { if _, err := st.CreateMapTileSource(MapTileSourceInput{Name: "Custom", URLTemplate: "https://tiles2.example.com/{z}/{x}/{y}.png", MaxZoom: 18, Enabled: true, ProxyEnabled: true}); !errors.Is(err, ErrMapTileSourceAlreadyExists) {
t.Fatalf("duplicate name error = %v, want errMapTileSourceAlreadyExists", err) t.Fatalf("duplicate name error = %v, want ErrMapTileSourceAlreadyExists", err)
} }
if _, err := st.CreateMapTileSource(mapTileSourceInput{Name: "Custom 2", URLTemplate: first.URLTemplate, MaxZoom: 18, Enabled: true, ProxyEnabled: true}); !errors.Is(err, errMapTileSourceAlreadyExists) { if _, err := st.CreateMapTileSource(MapTileSourceInput{Name: "Custom 2", URLTemplate: first.URLTemplate, MaxZoom: 18, Enabled: true, ProxyEnabled: true}); !errors.Is(err, ErrMapTileSourceAlreadyExists) {
t.Fatalf("duplicate url error = %v, want errMapTileSourceAlreadyExists", err) t.Fatalf("duplicate url error = %v, want ErrMapTileSourceAlreadyExists", err)
} }
updated, err := st.SetDefaultMapTileSource(first.ID) updated, err := st.SetDefaultMapTileSource(first.ID)
@@ -102,11 +102,11 @@ func TestMapTileSourceDuplicateAndDefaultRules(t *testing.T) {
if oldDefault.ID != first.ID { if oldDefault.ID != first.ID {
t.Fatalf("default id = %d, want %d", oldDefault.ID, first.ID) t.Fatalf("default id = %d, want %d", oldDefault.ID, first.ID)
} }
if _, err := st.UpdateMapTileSource(first.ID, mapTileSourceInput{Name: first.Name, URLTemplate: first.URLTemplate, Attribution: first.Attribution, MaxZoom: first.MaxZoom, Enabled: false, IsDefault: true}); !errors.Is(err, errMapTileSourceCannotDisableDefault) { if _, err := st.UpdateMapTileSource(first.ID, MapTileSourceInput{Name: first.Name, URLTemplate: first.URLTemplate, Attribution: first.Attribution, MaxZoom: first.MaxZoom, Enabled: false, IsDefault: true}); !errors.Is(err, ErrMapTileSourceCannotDisableDefault) {
t.Fatalf("disable default error = %v, want errMapTileSourceCannotDisableDefault", err) t.Fatalf("disable default error = %v, want ErrMapTileSourceCannotDisableDefault", err)
} }
if err := st.DeleteMapTileSource(first.ID); !errors.Is(err, errMapTileSourceCannotDeleteDefault) { if err := st.DeleteMapTileSource(first.ID); !errors.Is(err, ErrMapTileSourceCannotDeleteDefault) {
t.Fatalf("delete default error = %v, want errMapTileSourceCannotDeleteDefault", err) t.Fatalf("delete default error = %v, want ErrMapTileSourceCannotDeleteDefault", err)
} }
} }
@@ -114,11 +114,11 @@ func TestMapTileSourceHashIsSetOnCreate(t *testing.T) {
st := openTestStore(t) st := openTestStore(t)
defer st.Close() defer st.Close()
row, err := st.CreateMapTileSource(mapTileSourceInput{Name: "Hashed", URLTemplate: "https://test.example.com/{z}/{x}/{y}.png", MaxZoom: 18, Enabled: true, ProxyEnabled: true}) row, err := st.CreateMapTileSource(MapTileSourceInput{Name: "Hashed", URLTemplate: "https://test.example.com/{z}/{x}/{y}.png", MaxZoom: 18, Enabled: true, ProxyEnabled: true})
if err != nil { if err != nil {
t.Fatalf("CreateMapTileSource() error = %v", err) t.Fatalf("CreateMapTileSource() error = %v", err)
} }
want := mapTileSourceHash("https://test.example.com/{z}/{x}/{y}.png") want := MapTileSourceHash("https://test.example.com/{z}/{x}/{y}.png")
if row.URLTemplateHash != want { if row.URLTemplateHash != want {
t.Fatalf("URLTemplateHash = %q, want %q", row.URLTemplateHash, want) t.Fatalf("URLTemplateHash = %q, want %q", row.URLTemplateHash, want)
} }
@@ -135,7 +135,7 @@ func TestMapTileSourceDefaultHasHash(t *testing.T) {
if err != nil { if err != nil {
t.Fatalf("GetDefaultMapTileSource() error = %v", err) t.Fatalf("GetDefaultMapTileSource() error = %v", err)
} }
want := mapTileSourceHash(defaultMapTileSourceURLTemplate) want := MapTileSourceHash(defaultMapTileSourceURLTemplate)
if row.URLTemplateHash != want { if row.URLTemplateHash != want {
t.Fatalf("default URLTemplateHash = %q, want %q", row.URLTemplateHash, want) t.Fatalf("default URLTemplateHash = %q, want %q", row.URLTemplateHash, want)
} }
@@ -145,7 +145,7 @@ func TestGetEnabledMapTileSourceByHash(t *testing.T) {
st := openTestStore(t) st := openTestStore(t)
defer st.Close() defer st.Close()
row, err := st.CreateMapTileSource(mapTileSourceInput{Name: "HashLookup", URLTemplate: "https://lookup.example.com/{z}/{x}/{y}.png", MaxZoom: 18, Enabled: true, ProxyEnabled: true}) row, err := st.CreateMapTileSource(MapTileSourceInput{Name: "HashLookup", URLTemplate: "https://lookup.example.com/{z}/{x}/{y}.png", MaxZoom: 18, Enabled: true, ProxyEnabled: true})
if err != nil { if err != nil {
t.Fatalf("CreateMapTileSource() error = %v", err) t.Fatalf("CreateMapTileSource() error = %v", err)
} }
@@ -163,7 +163,7 @@ func TestGetEnabledMapTileSourceByHashDisabled(t *testing.T) {
st := openTestStore(t) st := openTestStore(t)
defer st.Close() defer st.Close()
row, err := st.CreateMapTileSource(mapTileSourceInput{Name: "DisabledHash", URLTemplate: "https://disabled-hash.example.com/{z}/{x}/{y}.png", MaxZoom: 18, Enabled: false}) row, err := st.CreateMapTileSource(MapTileSourceInput{Name: "DisabledHash", URLTemplate: "https://disabled-hash.example.com/{z}/{x}/{y}.png", MaxZoom: 18, Enabled: false})
if err != nil { if err != nil {
t.Fatalf("CreateMapTileSource() error = %v", err) t.Fatalf("CreateMapTileSource() error = %v", err)
} }
@@ -178,7 +178,7 @@ func TestGetEnabledMapTileSourceByHashProxyDisabled(t *testing.T) {
st := openTestStore(t) st := openTestStore(t)
defer st.Close() defer st.Close()
row, err := st.CreateMapTileSource(mapTileSourceInput{Name: "ProxyDisabledHash", URLTemplate: "https://proxy-disabled.example.com/{z}/{x}/{y}.png", MaxZoom: 18, Enabled: true, ProxyEnabled: false}) row, err := st.CreateMapTileSource(MapTileSourceInput{Name: "ProxyDisabledHash", URLTemplate: "https://proxy-disabled.example.com/{z}/{x}/{y}.png", MaxZoom: 18, Enabled: true, ProxyEnabled: false})
if err != nil { if err != nil {
t.Fatalf("CreateMapTileSource() error = %v", err) t.Fatalf("CreateMapTileSource() error = %v", err)
} }
@@ -203,7 +203,7 @@ func TestPublicMapTileSourceDTOProxyURL(t *testing.T) {
st := openTestStore(t) st := openTestStore(t)
defer st.Close() defer st.Close()
row, err := st.CreateMapTileSource(mapTileSourceInput{Name: "ProxyTest", URLTemplate: "https://proxy.example.com/{z}/{x}/{y}.png", MaxZoom: 18, Enabled: true, ProxyEnabled: true}) row, err := st.CreateMapTileSource(MapTileSourceInput{Name: "ProxyTest", URLTemplate: "https://proxy.example.com/{z}/{x}/{y}.png", MaxZoom: 18, Enabled: true, ProxyEnabled: true})
if err != nil { if err != nil {
t.Fatalf("CreateMapTileSource() error = %v", err) t.Fatalf("CreateMapTileSource() error = %v", err)
} }
@@ -226,7 +226,7 @@ func TestPublicMapTileSourceDTORawURLWhenProxyDisabled(t *testing.T) {
st := openTestStore(t) st := openTestStore(t)
defer st.Close() defer st.Close()
row, err := st.CreateMapTileSource(mapTileSourceInput{Name: "RawTest", URLTemplate: "https://raw.example.com/{z}/{x}/{y}.png", MaxZoom: 18, Enabled: true, ProxyEnabled: false}) row, err := st.CreateMapTileSource(MapTileSourceInput{Name: "RawTest", URLTemplate: "https://raw.example.com/{z}/{x}/{y}.png", MaxZoom: 18, Enabled: true, ProxyEnabled: false})
if err != nil { if err != nil {
t.Fatalf("CreateMapTileSource() error = %v", err) t.Fatalf("CreateMapTileSource() error = %v", err)
} }
@@ -242,9 +242,9 @@ func TestPublicMapTileSourceDTORawURLWhenProxyDisabled(t *testing.T) {
} }
func TestMapTileSourceHashFunction(t *testing.T) { func TestMapTileSourceHashFunction(t *testing.T) {
hash1 := mapTileSourceHash("https://tile.openstreetmap.jp/{z}/{x}/{y}.png") hash1 := MapTileSourceHash("https://tile.openstreetmap.jp/{z}/{x}/{y}.png")
hash2 := mapTileSourceHash("https://tile.openstreetmap.jp/{z}/{x}/{y}.png") hash2 := MapTileSourceHash("https://tile.openstreetmap.jp/{z}/{x}/{y}.png")
hash3 := mapTileSourceHash("https://other.example.com/{z}/{x}/{y}.png") hash3 := MapTileSourceHash("https://other.example.com/{z}/{x}/{y}.png")
if hash1 != hash2 { if hash1 != hash2 {
t.Fatal("hash should be deterministic") t.Fatal("hash should be deterministic")
@@ -1,4 +1,4 @@
package main package store
import ( import (
"errors" "errors"
@@ -10,16 +10,16 @@ import (
) )
const ( const (
mqttForwardDirectionSourceToTarget = "source_to_target" MQTTForwardDirectionSourceToTarget = "source_to_target"
mqttForwardDirectionBidirectional = "bidirectional" MQTTForwardDirectionBidirectional = "bidirectional"
) )
var ( var (
errMQTTForwarderAlreadyExists = errors.New("mqtt forwarder already exists") ErrMQTTForwarderAlreadyExists = errors.New("mqtt forwarder already exists")
errMQTTForwardTopicAlreadyExists = errors.New("mqtt forward topic already exists") ErrMQTTForwardTopicAlreadyExists = errors.New("mqtt forward topic already exists")
) )
type mqttForwarderInput struct { type MQTTForwarderInput struct {
Name string Name string
Enabled bool Enabled bool
SourceHost string SourceHost string
@@ -36,7 +36,7 @@ type mqttForwarderInput struct {
TargetTLS bool TargetTLS bool
} }
type mqttForwardTopicInput struct { type MQTTForwardTopicInput struct {
Topic string Topic string
Enabled bool Enabled bool
Direction string Direction string
@@ -46,15 +46,15 @@ type mqttForwardTopicInput struct {
Retain bool Retain bool
} }
type mqttForwarderConfig struct { type MQTTForwarderConfig struct {
Forwarder mqttForwarderRecord Forwarder MQTTForwarderRecord
Topics []mqttForwardTopicRecord Topics []MQTTForwardTopicRecord
} }
func (s *store) ListMQTTForwarders(opts listOptions) ([]mqttForwarderRecord, error) { func (s *Store) ListMQTTForwarders(opts ListOptions) ([]MQTTForwarderRecord, error) {
opts = normalizeListOptions(opts) opts = NormalizeListOptions(opts)
var rows []mqttForwarderRecord var rows []MQTTForwarderRecord
q := s.db.Model(&mqttForwarderRecord{}). q := s.db.Model(&MQTTForwarderRecord{}).
Order("updated_at DESC"). Order("updated_at DESC").
Order("id DESC"). Order("id DESC").
Limit(opts.Limit). Limit(opts.Limit).
@@ -62,20 +62,20 @@ func (s *store) ListMQTTForwarders(opts listOptions) ([]mqttForwarderRecord, err
return rows, q.Find(&rows).Error return rows, q.Find(&rows).Error
} }
func (s *store) CountMQTTForwarders(opts listOptions) (int64, error) { func (s *Store) CountMQTTForwarders(opts ListOptions) (int64, error) {
var total int64 var total int64
return total, s.db.Model(&mqttForwarderRecord{}).Count(&total).Error return total, s.db.Model(&MQTTForwarderRecord{}).Count(&total).Error
} }
func (s *store) GetMQTTForwarder(id uint64) (*mqttForwarderRecord, error) { func (s *Store) GetMQTTForwarder(id uint64) (*MQTTForwarderRecord, error) {
var row mqttForwarderRecord var row MQTTForwarderRecord
if err := s.db.Where("id = ?", id).Take(&row).Error; err != nil { if err := s.db.Where("id = ?", id).Take(&row).Error; err != nil {
return nil, err return nil, err
} }
return &row, nil return &row, nil
} }
func (s *store) CreateMQTTForwarder(input mqttForwarderInput) (*mqttForwarderRecord, error) { func (s *Store) CreateMQTTForwarder(input MQTTForwarderInput) (*MQTTForwarderRecord, error) {
row, err := mqttForwarderFromInput(input, nil) row, err := mqttForwarderFromInput(input, nil)
if err != nil { if err != nil {
return nil, err return nil, err
@@ -89,7 +89,7 @@ func (s *store) CreateMQTTForwarder(input mqttForwarderInput) (*mqttForwarderRec
return row, nil return row, nil
} }
func (s *store) UpdateMQTTForwarder(id uint64, input mqttForwarderInput) (*mqttForwarderRecord, error) { func (s *Store) UpdateMQTTForwarder(id uint64, input MQTTForwarderInput) (*MQTTForwarderRecord, error) {
if id == 0 { if id == 0 {
return nil, fmt.Errorf("mqtt forwarder id is required") return nil, fmt.Errorf("mqtt forwarder id is required")
} }
@@ -112,21 +112,21 @@ func (s *store) UpdateMQTTForwarder(id uint64, input mqttForwarderInput) (*mqttF
"target_password": row.TargetPassword, "target_client_id": row.TargetClientID, "target_tls": row.TargetTLS, "target_password": row.TargetPassword, "target_client_id": row.TargetClientID, "target_tls": row.TargetTLS,
"updated_at": time.Now(), "updated_at": time.Now(),
} }
if err := s.db.Model(&mqttForwarderRecord{}).Where("id = ?", id).Updates(updates).Error; err != nil { if err := s.db.Model(&MQTTForwarderRecord{}).Where("id = ?", id).Updates(updates).Error; err != nil {
return nil, err return nil, err
} }
return s.GetMQTTForwarder(id) return s.GetMQTTForwarder(id)
} }
func (s *store) DeleteMQTTForwarder(id uint64) error { func (s *Store) DeleteMQTTForwarder(id uint64) error {
if id == 0 { if id == 0 {
return fmt.Errorf("mqtt forwarder id is required") return fmt.Errorf("mqtt forwarder id is required")
} }
return s.db.Transaction(func(tx *gorm.DB) error { return s.db.Transaction(func(tx *gorm.DB) error {
if err := tx.Where("forwarder_id = ?", id).Delete(&mqttForwardTopicRecord{}).Error; err != nil { if err := tx.Where("forwarder_id = ?", id).Delete(&MQTTForwardTopicRecord{}).Error; err != nil {
return err return err
} }
result := tx.Where("id = ?", id).Delete(&mqttForwarderRecord{}) result := tx.Where("id = ?", id).Delete(&MQTTForwarderRecord{})
if result.Error != nil { if result.Error != nil {
return result.Error return result.Error
} }
@@ -137,10 +137,10 @@ func (s *store) DeleteMQTTForwarder(id uint64) error {
}) })
} }
func (s *store) ListMQTTForwardTopics(forwarderID uint64, opts listOptions) ([]mqttForwardTopicRecord, error) { func (s *Store) ListMQTTForwardTopics(forwarderID uint64, opts ListOptions) ([]MQTTForwardTopicRecord, error) {
opts = normalizeListOptions(opts) opts = NormalizeListOptions(opts)
var rows []mqttForwardTopicRecord var rows []MQTTForwardTopicRecord
q := s.db.Model(&mqttForwardTopicRecord{}). q := s.db.Model(&MQTTForwardTopicRecord{}).
Where("forwarder_id = ?", forwarderID). Where("forwarder_id = ?", forwarderID).
Order("updated_at DESC"). Order("updated_at DESC").
Order("id DESC"). Order("id DESC").
@@ -149,20 +149,20 @@ func (s *store) ListMQTTForwardTopics(forwarderID uint64, opts listOptions) ([]m
return rows, q.Find(&rows).Error return rows, q.Find(&rows).Error
} }
func (s *store) CountMQTTForwardTopics(forwarderID uint64) (int64, error) { func (s *Store) CountMQTTForwardTopics(forwarderID uint64) (int64, error) {
var total int64 var total int64
return total, s.db.Model(&mqttForwardTopicRecord{}).Where("forwarder_id = ?", forwarderID).Count(&total).Error return total, s.db.Model(&MQTTForwardTopicRecord{}).Where("forwarder_id = ?", forwarderID).Count(&total).Error
} }
func (s *store) GetMQTTForwardTopic(id uint64) (*mqttForwardTopicRecord, error) { func (s *Store) GetMQTTForwardTopic(id uint64) (*MQTTForwardTopicRecord, error) {
var row mqttForwardTopicRecord var row MQTTForwardTopicRecord
if err := s.db.Where("id = ?", id).Take(&row).Error; err != nil { if err := s.db.Where("id = ?", id).Take(&row).Error; err != nil {
return nil, err return nil, err
} }
return &row, nil return &row, nil
} }
func (s *store) CreateMQTTForwardTopic(forwarderID uint64, input mqttForwardTopicInput) (*mqttForwardTopicRecord, error) { func (s *Store) CreateMQTTForwardTopic(forwarderID uint64, input MQTTForwardTopicInput) (*MQTTForwardTopicRecord, error) {
if _, err := s.GetMQTTForwarder(forwarderID); err != nil { if _, err := s.GetMQTTForwarder(forwarderID); err != nil {
return nil, err return nil, err
} }
@@ -179,7 +179,7 @@ func (s *store) CreateMQTTForwardTopic(forwarderID uint64, input mqttForwardTopi
return row, nil return row, nil
} }
func (s *store) UpdateMQTTForwardTopic(id uint64, input mqttForwardTopicInput) (*mqttForwardTopicRecord, error) { func (s *Store) UpdateMQTTForwardTopic(id uint64, input MQTTForwardTopicInput) (*MQTTForwardTopicRecord, error) {
if id == 0 { if id == 0 {
return nil, fmt.Errorf("mqtt forward topic id is required") return nil, fmt.Errorf("mqtt forward topic id is required")
} }
@@ -199,14 +199,14 @@ func (s *store) UpdateMQTTForwardTopic(id uint64, input mqttForwardTopicInput) (
"source_prefix": row.SourcePrefix, "target_prefix": row.TargetPrefix, "source_prefix": row.SourcePrefix, "target_prefix": row.TargetPrefix,
"qos": row.QoS, "retain": row.Retain, "updated_at": time.Now(), "qos": row.QoS, "retain": row.Retain, "updated_at": time.Now(),
} }
if err := s.db.Model(&mqttForwardTopicRecord{}).Where("id = ?", id).Updates(updates).Error; err != nil { if err := s.db.Model(&MQTTForwardTopicRecord{}).Where("id = ?", id).Updates(updates).Error; err != nil {
return nil, err return nil, err
} }
return s.GetMQTTForwardTopic(id) return s.GetMQTTForwardTopic(id)
} }
func (s *store) DeleteMQTTForwardTopic(id uint64) error { func (s *Store) DeleteMQTTForwardTopic(id uint64) error {
result := s.db.Where("id = ?", id).Delete(&mqttForwardTopicRecord{}) result := s.db.Where("id = ?", id).Delete(&MQTTForwardTopicRecord{})
if result.Error != nil { if result.Error != nil {
return result.Error return result.Error
} }
@@ -216,39 +216,39 @@ func (s *store) DeleteMQTTForwardTopic(id uint64) error {
return nil return nil
} }
func (s *store) GetMQTTForwarderConfig(id uint64) (*mqttForwarderConfig, error) { func (s *Store) GetMQTTForwarderConfig(id uint64) (*MQTTForwarderConfig, error) {
forwarder, err := s.GetMQTTForwarder(id) forwarder, err := s.GetMQTTForwarder(id)
if err != nil { if err != nil {
return nil, err return nil, err
} }
var topics []mqttForwardTopicRecord var topics []MQTTForwardTopicRecord
if err := s.db.Where("forwarder_id = ? AND enabled = ?", id, true).Order("id ASC").Find(&topics).Error; err != nil { if err := s.db.Where("forwarder_id = ? AND enabled = ?", id, true).Order("id ASC").Find(&topics).Error; err != nil {
return nil, err return nil, err
} }
return &mqttForwarderConfig{Forwarder: *forwarder, Topics: topics}, nil return &MQTTForwarderConfig{Forwarder: *forwarder, Topics: topics}, nil
} }
func (s *store) ListEnabledMQTTForwarderConfigs() ([]mqttForwarderConfig, error) { func (s *Store) ListEnabledMQTTForwarderConfigs() ([]MQTTForwarderConfig, error) {
var forwarders []mqttForwarderRecord var forwarders []MQTTForwarderRecord
if err := s.db.Where("enabled = ?", true).Order("id ASC").Find(&forwarders).Error; err != nil { if err := s.db.Where("enabled = ?", true).Order("id ASC").Find(&forwarders).Error; err != nil {
return nil, err return nil, err
} }
configs := make([]mqttForwarderConfig, 0, len(forwarders)) configs := make([]MQTTForwarderConfig, 0, len(forwarders))
for _, forwarder := range forwarders { for _, forwarder := range forwarders {
var topics []mqttForwardTopicRecord var topics []MQTTForwardTopicRecord
if err := s.db.Where("forwarder_id = ? AND enabled = ?", forwarder.ID, true).Order("id ASC").Find(&topics).Error; err != nil { if err := s.db.Where("forwarder_id = ? AND enabled = ?", forwarder.ID, true).Order("id ASC").Find(&topics).Error; err != nil {
return nil, err return nil, err
} }
if len(topics) == 0 { if len(topics) == 0 {
continue continue
} }
configs = append(configs, mqttForwarderConfig{Forwarder: forwarder, Topics: topics}) configs = append(configs, MQTTForwarderConfig{Forwarder: forwarder, Topics: topics})
} }
return configs, nil return configs, nil
} }
func (s *store) ensureMQTTForwarderNameUnique(id uint64, name string) error { func (s *Store) ensureMQTTForwarderNameUnique(id uint64, name string) error {
var existing mqttForwarderRecord var existing MQTTForwarderRecord
q := s.db.Where("name = ?", name) q := s.db.Where("name = ?", name)
if id != 0 { if id != 0 {
q = q.Where("id <> ?", id) q = q.Where("id <> ?", id)
@@ -260,11 +260,11 @@ func (s *store) ensureMQTTForwarderNameUnique(id uint64, name string) error {
if err != nil { if err != nil {
return err return err
} }
return errMQTTForwarderAlreadyExists return ErrMQTTForwarderAlreadyExists
} }
func (s *store) ensureMQTTForwardTopicUnique(id, forwarderID uint64, topic string) error { func (s *Store) ensureMQTTForwardTopicUnique(id, forwarderID uint64, topic string) error {
var existing mqttForwardTopicRecord var existing MQTTForwardTopicRecord
q := s.db.Where("forwarder_id = ? AND topic = ?", forwarderID, topic) q := s.db.Where("forwarder_id = ? AND topic = ?", forwarderID, topic)
if id != 0 { if id != 0 {
q = q.Where("id <> ?", id) q = q.Where("id <> ?", id)
@@ -276,10 +276,10 @@ func (s *store) ensureMQTTForwardTopicUnique(id, forwarderID uint64, topic strin
if err != nil { if err != nil {
return err return err
} }
return errMQTTForwardTopicAlreadyExists return ErrMQTTForwardTopicAlreadyExists
} }
func mqttForwarderFromInput(input mqttForwarderInput, existing *mqttForwarderRecord) (*mqttForwarderRecord, error) { func mqttForwarderFromInput(input MQTTForwarderInput, existing *MQTTForwarderRecord) (*MQTTForwarderRecord, error) {
name := strings.TrimSpace(input.Name) name := strings.TrimSpace(input.Name)
if name == "" { if name == "" {
return nil, fmt.Errorf("mqtt forwarder name is required") return nil, fmt.Errorf("mqtt forwarder name is required")
@@ -298,7 +298,7 @@ func mqttForwarderFromInput(input mqttForwarderInput, existing *mqttForwarderRec
if err := validateMQTTForwardPort(input.TargetPort, "target port"); err != nil { if err := validateMQTTForwardPort(input.TargetPort, "target port"); err != nil {
return nil, err return nil, err
} }
row := &mqttForwarderRecord{ row := &MQTTForwarderRecord{
Name: name, Enabled: input.Enabled, Name: name, Enabled: input.Enabled,
SourceHost: sourceHost, SourcePort: input.SourcePort, SourceUsername: strings.TrimSpace(input.SourceUsername), SourceClientID: strings.TrimSpace(input.SourceClientID), SourceTLS: input.SourceTLS, SourceHost: sourceHost, SourcePort: input.SourcePort, SourceUsername: strings.TrimSpace(input.SourceUsername), SourceClientID: strings.TrimSpace(input.SourceClientID), SourceTLS: input.SourceTLS,
TargetHost: targetHost, TargetPort: input.TargetPort, TargetUsername: strings.TrimSpace(input.TargetUsername), TargetClientID: strings.TrimSpace(input.TargetClientID), TargetTLS: input.TargetTLS, TargetHost: targetHost, TargetPort: input.TargetPort, TargetUsername: strings.TrimSpace(input.TargetUsername), TargetClientID: strings.TrimSpace(input.TargetClientID), TargetTLS: input.TargetTLS,
@@ -316,7 +316,7 @@ func mqttForwarderFromInput(input mqttForwarderInput, existing *mqttForwarderRec
return row, nil return row, nil
} }
func mqttForwardTopicFromInput(forwarderID uint64, input mqttForwardTopicInput) (*mqttForwardTopicRecord, error) { func mqttForwardTopicFromInput(forwarderID uint64, input MQTTForwardTopicInput) (*MQTTForwardTopicRecord, error) {
if forwarderID == 0 { if forwarderID == 0 {
return nil, fmt.Errorf("mqtt forwarder id is required") return nil, fmt.Errorf("mqtt forwarder id is required")
} }
@@ -331,7 +331,7 @@ func mqttForwardTopicFromInput(forwarderID uint64, input mqttForwardTopicInput)
if input.QoS < 0 || input.QoS > 2 { if input.QoS < 0 || input.QoS > 2 {
return nil, fmt.Errorf("qos must be 0, 1, or 2") return nil, fmt.Errorf("qos must be 0, 1, or 2")
} }
return &mqttForwardTopicRecord{ return &MQTTForwardTopicRecord{
ForwarderID: forwarderID, Topic: topic, Enabled: input.Enabled, Direction: direction, ForwarderID: forwarderID, Topic: topic, Enabled: input.Enabled, Direction: direction,
SourcePrefix: strings.Trim(strings.TrimSpace(input.SourcePrefix), "/"), SourcePrefix: strings.Trim(strings.TrimSpace(input.SourcePrefix), "/"),
TargetPrefix: strings.Trim(strings.TrimSpace(input.TargetPrefix), "/"), TargetPrefix: strings.Trim(strings.TrimSpace(input.TargetPrefix), "/"),
@@ -349,10 +349,10 @@ func validateMQTTForwardPort(port int, label string) error {
func normalizeMQTTForwardDirection(direction string) (string, error) { func normalizeMQTTForwardDirection(direction string) (string, error) {
direction = strings.TrimSpace(direction) direction = strings.TrimSpace(direction)
if direction == "" { if direction == "" {
direction = mqttForwardDirectionSourceToTarget direction = MQTTForwardDirectionSourceToTarget
} }
switch direction { switch direction {
case mqttForwardDirectionSourceToTarget, mqttForwardDirectionBidirectional: case MQTTForwardDirectionSourceToTarget, MQTTForwardDirectionBidirectional:
return direction, nil return direction, nil
default: default:
return "", fmt.Errorf("invalid mqtt forward direction") return "", fmt.Errorf("invalid mqtt forward direction")
@@ -1,4 +1,4 @@
package main package store
import ( import (
"errors" "errors"
@@ -12,45 +12,45 @@ import (
) )
const ( const (
runtimeSettingAllowEncryptedForwarding = "mqtt.allow_encrypted_forwarding" RuntimeSettingAllowEncryptedForwarding = "mqtt.allow_encrypted_forwarding"
runtimeSettingLLMQueueEnabled = "llm.queue_enabled" RuntimeSettingLLMQueueEnabled = "llm.queue_enabled"
runtimeSettingLLMQueueIncludeChannel = "llm.include_channel_messages" RuntimeSettingLLMQueueIncludeChannel = "llm.include_channel_messages"
runtimeSettingTypeBool = "bool" runtimeSettingTypeBool = "bool"
) )
type runtimeSettingsSnapshot struct { type RuntimeSettingsSnapshot struct {
AllowEncryptedForwarding bool AllowEncryptedForwarding bool
LLMQueueEnabled bool LLMQueueEnabled bool
LLMIncludeChannel bool LLMIncludeChannel bool
} }
func (s *store) GetRuntimeSettings() (runtimeSettingsSnapshot, error) { func (s *Store) GetRuntimeSettings() (RuntimeSettingsSnapshot, error) {
allowEncrypted, err := s.GetBoolRuntimeSetting(runtimeSettingAllowEncryptedForwarding, false) allowEncrypted, err := s.GetBoolRuntimeSetting(RuntimeSettingAllowEncryptedForwarding, false)
if err != nil { if err != nil {
return runtimeSettingsSnapshot{}, err return RuntimeSettingsSnapshot{}, err
} }
llmQueueEnabled, err := s.GetBoolRuntimeSetting(runtimeSettingLLMQueueEnabled, true) llmQueueEnabled, err := s.GetBoolRuntimeSetting(RuntimeSettingLLMQueueEnabled, true)
if err != nil { if err != nil {
return runtimeSettingsSnapshot{}, err return RuntimeSettingsSnapshot{}, err
} }
llmIncludeChannel, err := s.GetBoolRuntimeSetting(runtimeSettingLLMQueueIncludeChannel, false) llmIncludeChannel, err := s.GetBoolRuntimeSetting(RuntimeSettingLLMQueueIncludeChannel, false)
if err != nil { if err != nil {
return runtimeSettingsSnapshot{}, err return RuntimeSettingsSnapshot{}, err
} }
return runtimeSettingsSnapshot{ return RuntimeSettingsSnapshot{
AllowEncryptedForwarding: allowEncrypted, AllowEncryptedForwarding: allowEncrypted,
LLMQueueEnabled: llmQueueEnabled, LLMQueueEnabled: llmQueueEnabled,
LLMIncludeChannel: llmIncludeChannel, LLMIncludeChannel: llmIncludeChannel,
}, nil }, nil
} }
func (s *store) GetBoolRuntimeSetting(key string, defaultValue bool) (bool, error) { func (s *Store) GetBoolRuntimeSetting(key string, defaultValue bool) (bool, error) {
key = strings.TrimSpace(key) key = strings.TrimSpace(key)
if key == "" { if key == "" {
return false, fmt.Errorf("runtime setting key is required") return false, fmt.Errorf("runtime setting key is required")
} }
var row runtimeSettingRecord var row RuntimeSettingRecord
err := s.db.Where("`key` = ?", key).Take(&row).Error err := s.db.Where("`key` = ?", key).Take(&row).Error
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
return defaultValue, nil return defaultValue, nil
@@ -68,13 +68,13 @@ func (s *store) GetBoolRuntimeSetting(key string, defaultValue bool) (bool, erro
return value, nil return value, nil
} }
func (s *store) SetBoolRuntimeSetting(key string, value bool, label string) (*runtimeSettingRecord, error) { func (s *Store) SetBoolRuntimeSetting(key string, value bool, label string) (*RuntimeSettingRecord, error) {
key = strings.TrimSpace(key) key = strings.TrimSpace(key)
if key == "" { if key == "" {
return nil, fmt.Errorf("runtime setting key is required") return nil, fmt.Errorf("runtime setting key is required")
} }
row := runtimeSettingRecord{ row := RuntimeSettingRecord{
Key: key, Key: key,
Value: strconv.FormatBool(value), Value: strconv.FormatBool(value),
ValueType: runtimeSettingTypeBool, ValueType: runtimeSettingTypeBool,
@@ -1,4 +1,4 @@
package main package store
import "testing" import "testing"
@@ -14,7 +14,7 @@ func TestRuntimeSettingsDefaultAndUpdates(t *testing.T) {
t.Fatalf("AllowEncryptedForwarding = true, want false") t.Fatalf("AllowEncryptedForwarding = true, want false")
} }
if _, err := st.SetBoolRuntimeSetting(runtimeSettingAllowEncryptedForwarding, true, "test setting"); err != nil { if _, err := st.SetBoolRuntimeSetting(RuntimeSettingAllowEncryptedForwarding, true, "test setting"); err != nil {
t.Fatalf("SetBoolRuntimeSetting(true) error = %v", err) t.Fatalf("SetBoolRuntimeSetting(true) error = %v", err)
} }
settings, err = st.GetRuntimeSettings() settings, err = st.GetRuntimeSettings()
@@ -25,7 +25,7 @@ func TestRuntimeSettingsDefaultAndUpdates(t *testing.T) {
t.Fatalf("AllowEncryptedForwarding = false, want true") t.Fatalf("AllowEncryptedForwarding = false, want true")
} }
if _, err := st.SetBoolRuntimeSetting(runtimeSettingAllowEncryptedForwarding, false, "test setting"); err != nil { if _, err := st.SetBoolRuntimeSetting(RuntimeSettingAllowEncryptedForwarding, false, "test setting"); err != nil {
t.Fatalf("SetBoolRuntimeSetting(false) error = %v", err) t.Fatalf("SetBoolRuntimeSetting(false) error = %v", err)
} }
settings, err = st.GetRuntimeSettings() settings, err = st.GetRuntimeSettings()
+23 -21
View File
@@ -1,4 +1,4 @@
package main package store
import ( import (
"fmt" "fmt"
@@ -6,12 +6,14 @@ import (
"time" "time"
"gorm.io/gorm" "gorm.io/gorm"
"meshtastic_mqtt_server/internal/config"
) )
func (s *store) ListSigns(opts listOptions) ([]signRecord, error) { func (s *Store) ListSigns(opts ListOptions) ([]SignRecord, error) {
opts = normalizeListOptions(opts) opts = NormalizeListOptions(opts)
var rows []signRecord var rows []SignRecord
q := applySignFilters(s.db.Model(&signRecord{}), opts). q := applySignFilters(s.db.Model(&SignRecord{}), opts).
Order("sign_time DESC"). Order("sign_time DESC").
Order("id DESC"). Order("id DESC").
Limit(opts.Limit). Limit(opts.Limit).
@@ -19,39 +21,39 @@ func (s *store) ListSigns(opts listOptions) ([]signRecord, error) {
return rows, q.Find(&rows).Error return rows, q.Find(&rows).Error
} }
type signDayCount struct { type SignDayCount struct {
Date string `gorm:"column:sign_date"` Date string `gorm:"column:sign_date"`
Count int64 `gorm:"column:count"` Count int64 `gorm:"column:count"`
} }
func (s *store) CountSigns(opts listOptions) (int64, error) { func (s *Store) CountSigns(opts ListOptions) (int64, error) {
var total int64 var total int64
q := applySignFilters(s.db.Model(&signRecord{}), opts) q := applySignFilters(s.db.Model(&SignRecord{}), opts)
return total, q.Count(&total).Error return total, q.Count(&total).Error
} }
func (s *store) CountSignsByDay(opts listOptions) ([]signDayCount, error) { func (s *Store) CountSignsByDay(opts ListOptions) ([]SignDayCount, error) {
var rows []signDayCount var rows []SignDayCount
dateExpr := "strftime('%Y-%m-%d', sign_time)" dateExpr := "strftime('%Y-%m-%d', sign_time)"
if s.driver == databaseDriverMySQL { if s.driver == config.DriverMySQL {
dateExpr = "DATE_FORMAT(sign_time, '%Y-%m-%d')" dateExpr = "DATE_FORMAT(sign_time, '%Y-%m-%d')"
} }
q := applySignFilters(s.db.Model(&signRecord{}), opts). q := applySignFilters(s.db.Model(&SignRecord{}), opts).
Select(dateExpr + " AS sign_date, COUNT(*) AS count"). Select(dateExpr + " AS sign_date, COUNT(*) AS count").
Group(dateExpr). Group(dateExpr).
Order("sign_date DESC") Order("sign_date DESC")
return rows, q.Scan(&rows).Error return rows, q.Scan(&rows).Error
} }
func (s *store) GetSignByID(id uint64) (*signRecord, error) { func (s *Store) GetSignByID(id uint64) (*SignRecord, error) {
var row signRecord var row SignRecord
if err := s.db.Where("id = ?", id).Take(&row).Error; err != nil { if err := s.db.Where("id = ?", id).Take(&row).Error; err != nil {
return nil, err return nil, err
} }
return &row, nil return &row, nil
} }
func (s *store) CreateSign(nodeID string, longName, shortName *string, signText string, signTime time.Time) (*signRecord, error) { func (s *Store) CreateSign(nodeID string, longName, shortName *string, signText string, signTime time.Time) (*SignRecord, error) {
nodeID = strings.TrimSpace(nodeID) nodeID = strings.TrimSpace(nodeID)
signText = strings.TrimSpace(signText) signText = strings.TrimSpace(signText)
if nodeID == "" { if nodeID == "" {
@@ -63,14 +65,14 @@ func (s *store) CreateSign(nodeID string, longName, shortName *string, signText
if signTime.IsZero() { if signTime.IsZero() {
signTime = time.Now() signTime = time.Now()
} }
row := signRecord{NodeID: nodeID, LongName: trimNullableString(longName), ShortName: trimNullableString(shortName), SignText: signText, SignTime: signTime} row := SignRecord{NodeID: nodeID, LongName: trimNullableString(longName), ShortName: trimNullableString(shortName), SignText: signText, SignTime: signTime}
if err := s.db.Create(&row).Error; err != nil { if err := s.db.Create(&row).Error; err != nil {
return nil, err return nil, err
} }
return &row, nil return &row, nil
} }
func (s *store) UpdateSign(id uint64, nodeID string, longName, shortName *string, signText string, signTime time.Time) (*signRecord, error) { func (s *Store) UpdateSign(id uint64, nodeID string, longName, shortName *string, signText string, signTime time.Time) (*SignRecord, error) {
if id == 0 { if id == 0 {
return nil, fmt.Errorf("sign id is required") return nil, fmt.Errorf("sign id is required")
} }
@@ -89,14 +91,14 @@ func (s *store) UpdateSign(id uint64, nodeID string, longName, shortName *string
return nil, err return nil, err
} }
updates := map[string]any{"node_id": nodeID, "long_name": trimNullableString(longName), "short_name": trimNullableString(shortName), "sign_text": signText, "sign_time": signTime} updates := map[string]any{"node_id": nodeID, "long_name": trimNullableString(longName), "short_name": trimNullableString(shortName), "sign_text": signText, "sign_time": signTime}
if err := s.db.Model(&signRecord{}).Where("id = ?", id).Updates(updates).Error; err != nil { if err := s.db.Model(&SignRecord{}).Where("id = ?", id).Updates(updates).Error; err != nil {
return nil, err return nil, err
} }
return s.GetSignByID(id) return s.GetSignByID(id)
} }
func (s *store) DeleteSign(id uint64) error { func (s *Store) DeleteSign(id uint64) error {
result := s.db.Where("id = ?", id).Delete(&signRecord{}) result := s.db.Where("id = ?", id).Delete(&SignRecord{})
if result.Error != nil { if result.Error != nil {
return result.Error return result.Error
} }
@@ -106,7 +108,7 @@ func (s *store) DeleteSign(id uint64) error {
return nil return nil
} }
func applySignFilters(q *gorm.DB, opts listOptions) *gorm.DB { func applySignFilters(q *gorm.DB, opts ListOptions) *gorm.DB {
if opts.NodeID != "" { if opts.NodeID != "" {
q = q.Where("node_id = ?", opts.NodeID) q = q.Where("node_id = ?", opts.NodeID)
} }
@@ -1,4 +1,4 @@
package main package store
import ( import (
"fmt" "fmt"
@@ -8,7 +8,7 @@ import (
"gorm.io/gorm" "gorm.io/gorm"
) )
type listOptions struct { type ListOptions struct {
Limit int Limit int
Offset int Offset int
NodeID string NodeID string
@@ -21,31 +21,31 @@ type listOptions struct {
MaxLng *float64 MaxLng *float64
} }
type mapReportViewportOptions struct { type MapReportViewportOptions struct {
ListOptions listOptions ListOptions ListOptions
Zoom int Zoom int
Limit int Limit int
ClusterThreshold int ClusterThreshold int
TargetCells int TargetCells int
} }
type mapReportViewportResult struct { type MapReportViewportResult struct {
Mode string Mode string
Total int64 Total int64
Points []mapReportRecord Points []MapReportRecord
Clusters []mapReportClusterRecord Clusters []MapReportClusterRecord
Limit int Limit int
Zoom int Zoom int
} }
type mapReportClusterRecord struct { type MapReportClusterRecord struct {
ClusterID string ClusterID string
Latitude float64 Latitude float64
Longitude float64 Longitude float64
Count int64 Count int64
} }
func (s *store) Ping() error { func (s *Store) Ping() error {
db, err := s.db.DB() db, err := s.db.DB()
if err != nil { if err != nil {
return err return err
@@ -53,7 +53,7 @@ func (s *store) Ping() error {
return db.Ping() return db.Ping()
} }
func normalizeListOptions(opts listOptions) listOptions { func NormalizeListOptions(opts ListOptions) ListOptions {
if opts.Limit <= 0 { if opts.Limit <= 0 {
opts.Limit = 100 opts.Limit = 100
} }
@@ -66,64 +66,64 @@ func normalizeListOptions(opts listOptions) listOptions {
return opts return opts
} }
func (s *store) ListNodeInfo(opts listOptions) ([]nodeInfoRecord, error) { func (s *Store) ListNodeInfo(opts ListOptions) ([]NodeInfoRecord, error) {
opts = normalizeListOptions(opts) opts = NormalizeListOptions(opts)
var rows []nodeInfoRecord var rows []NodeInfoRecord
q := applyNodeFilters(s.db.Model(&nodeInfoRecord{}), opts). q := applyNodeFilters(s.db.Model(&NodeInfoRecord{}), opts).
Order("updated_at DESC"). Order("updated_at DESC").
Limit(opts.Limit). Limit(opts.Limit).
Offset(opts.Offset) Offset(opts.Offset)
return rows, q.Find(&rows).Error return rows, q.Find(&rows).Error
} }
func (s *store) CountNodeInfo(opts listOptions) (int64, error) { func (s *Store) CountNodeInfo(opts ListOptions) (int64, error) {
var total int64 var total int64
q := applyNodeFilters(s.db.Model(&nodeInfoRecord{}), opts) q := applyNodeFilters(s.db.Model(&NodeInfoRecord{}), opts)
return total, q.Count(&total).Error return total, q.Count(&total).Error
} }
func (s *store) GetNodeInfo(nodeID string) (*nodeInfoRecord, error) { func (s *Store) GetNodeInfo(nodeID string) (*NodeInfoRecord, error) {
var row nodeInfoRecord var row NodeInfoRecord
if err := s.db.Where("node_id = ?", nodeID).Take(&row).Error; err != nil { if err := s.db.Where("node_id = ?", nodeID).Take(&row).Error; err != nil {
return nil, err return nil, err
} }
return &row, nil return &row, nil
} }
func (s *store) ListMapReports(opts listOptions) ([]mapReportRecord, error) { func (s *Store) ListMapReports(opts ListOptions) ([]MapReportRecord, error) {
opts = normalizeListOptions(opts) opts = NormalizeListOptions(opts)
var rows []mapReportRecord var rows []MapReportRecord
q := applyMapReportFilters(s.db.Model(&mapReportRecord{}), opts). q := applyMapReportFilters(s.db.Model(&MapReportRecord{}), opts).
Order("updated_at DESC"). Order("updated_at DESC").
Limit(opts.Limit). Limit(opts.Limit).
Offset(opts.Offset) Offset(opts.Offset)
return rows, q.Find(&rows).Error return rows, q.Find(&rows).Error
} }
func (s *store) CountMapReports(opts listOptions) (int64, error) { func (s *Store) CountMapReports(opts ListOptions) (int64, error) {
var total int64 var total int64
q := applyMapReportFilters(s.db.Model(&mapReportRecord{}), opts) q := applyMapReportFilters(s.db.Model(&MapReportRecord{}), opts)
return total, q.Count(&total).Error return total, q.Count(&total).Error
} }
func (s *store) GetMapReport(nodeID string) (*mapReportRecord, error) { func (s *Store) GetMapReport(nodeID string) (*MapReportRecord, error) {
var row mapReportRecord var row MapReportRecord
if err := s.db.Where("node_id = ?", nodeID).Take(&row).Error; err != nil { if err := s.db.Where("node_id = ?", nodeID).Take(&row).Error; err != nil {
return nil, err return nil, err
} }
return &row, nil return &row, nil
} }
func (s *store) ListMapReportViewport(opts mapReportViewportOptions) (*mapReportViewportResult, error) { func (s *Store) ListMapReportViewport(opts MapReportViewportOptions) (*MapReportViewportResult, error) {
opts = normalizeMapReportViewportOptions(opts) opts = NormalizeMapReportViewportOptions(opts)
total, err := s.CountMapReports(opts.ListOptions) total, err := s.CountMapReports(opts.ListOptions)
if err != nil { if err != nil {
return nil, err return nil, err
} }
result := &mapReportViewportResult{Total: total, Limit: opts.Limit, Zoom: opts.Zoom} result := &MapReportViewportResult{Total: total, Limit: opts.Limit, Zoom: opts.Zoom}
if total <= int64(opts.ClusterThreshold) { if total <= int64(opts.ClusterThreshold) {
var points []mapReportRecord var points []MapReportRecord
q := applyMapReportFilters(s.db.Model(&mapReportRecord{}), opts.ListOptions). q := applyMapReportFilters(s.db.Model(&MapReportRecord{}), opts.ListOptions).
Order("updated_at DESC"). Order("updated_at DESC").
Limit(opts.Limit) Limit(opts.Limit)
if err := q.Find(&points).Error; err != nil { if err := q.Find(&points).Error; err != nil {
@@ -142,8 +142,8 @@ func (s *store) ListMapReportViewport(opts mapReportViewportOptions) (*mapReport
return result, nil return result, nil
} }
func (s *store) ListMapReportClusters(opts mapReportViewportOptions) ([]mapReportClusterRecord, error) { func (s *Store) ListMapReportClusters(opts MapReportViewportOptions) ([]MapReportClusterRecord, error) {
opts = normalizeMapReportViewportOptions(opts) opts = NormalizeMapReportViewportOptions(opts)
cellSize := mapReportClusterCellSize(opts.ListOptions, opts.TargetCells) cellSize := mapReportClusterCellSize(opts.ListOptions, opts.TargetCells)
var rows []struct { var rows []struct {
LatBucket int64 LatBucket int64
@@ -152,7 +152,7 @@ func (s *store) ListMapReportClusters(opts mapReportViewportOptions) ([]mapRepor
Longitude float64 Longitude float64
Count int64 Count int64
} }
q := applyMapReportFilters(s.db.Model(&mapReportRecord{}), opts.ListOptions). q := applyMapReportFilters(s.db.Model(&MapReportRecord{}), opts.ListOptions).
Select("CAST((latitude + 90.0) / ? AS INTEGER) AS lat_bucket, CAST((longitude + 180.0) / ? AS INTEGER) AS lng_bucket, AVG(latitude) AS latitude, AVG(longitude) AS longitude, COUNT(*) AS count", cellSize, cellSize). Select("CAST((latitude + 90.0) / ? AS INTEGER) AS lat_bucket, CAST((longitude + 180.0) / ? AS INTEGER) AS lng_bucket, AVG(latitude) AS latitude, AVG(longitude) AS longitude, COUNT(*) AS count", cellSize, cellSize).
Group("lat_bucket, lng_bucket"). Group("lat_bucket, lng_bucket").
Order("count DESC"). Order("count DESC").
@@ -160,9 +160,9 @@ func (s *store) ListMapReportClusters(opts mapReportViewportOptions) ([]mapRepor
if err := q.Scan(&rows).Error; err != nil { if err := q.Scan(&rows).Error; err != nil {
return nil, err return nil, err
} }
clusters := make([]mapReportClusterRecord, 0, len(rows)) clusters := make([]MapReportClusterRecord, 0, len(rows))
for _, row := range rows { for _, row := range rows {
clusters = append(clusters, mapReportClusterRecord{ clusters = append(clusters, MapReportClusterRecord{
ClusterID: fmt.Sprintf("%d:%d", row.LatBucket, row.LngBucket), ClusterID: fmt.Sprintf("%d:%d", row.LatBucket, row.LngBucket),
Latitude: row.Latitude, Latitude: row.Latitude,
Longitude: row.Longitude, Longitude: row.Longitude,
@@ -172,7 +172,7 @@ func (s *store) ListMapReportClusters(opts mapReportViewportOptions) ([]mapRepor
return clusters, nil return clusters, nil
} }
func normalizeMapReportViewportOptions(opts mapReportViewportOptions) mapReportViewportOptions { func NormalizeMapReportViewportOptions(opts MapReportViewportOptions) MapReportViewportOptions {
if opts.Limit <= 0 { if opts.Limit <= 0 {
opts.Limit = 1000 opts.Limit = 1000
} }
@@ -194,7 +194,7 @@ func normalizeMapReportViewportOptions(opts mapReportViewportOptions) mapReportV
return opts return opts
} }
func mapReportClusterCellSize(opts listOptions, targetCells int) float64 { func mapReportClusterCellSize(opts ListOptions, targetCells int) float64 {
latSpan := 180.0 latSpan := 180.0
if opts.MinLat != nil && opts.MaxLat != nil { if opts.MinLat != nil && opts.MaxLat != nil {
latSpan = *opts.MaxLat - *opts.MinLat latSpan = *opts.MaxLat - *opts.MinLat
@@ -215,13 +215,13 @@ func mapReportClusterCellSize(opts listOptions, targetCells int) float64 {
return cellSize return cellSize
} }
func (s *store) DeleteNode(nodeID string) error { func (s *Store) DeleteNode(nodeID string) error {
return s.db.Transaction(func(tx *gorm.DB) error { return s.db.Transaction(func(tx *gorm.DB) error {
nodeResult := tx.Where("node_id = ?", nodeID).Delete(&nodeInfoRecord{}) nodeResult := tx.Where("node_id = ?", nodeID).Delete(&NodeInfoRecord{})
if nodeResult.Error != nil { if nodeResult.Error != nil {
return nodeResult.Error return nodeResult.Error
} }
reportResult := tx.Where("node_id = ?", nodeID).Delete(&mapReportRecord{}) reportResult := tx.Where("node_id = ?", nodeID).Delete(&MapReportRecord{})
if reportResult.Error != nil { if reportResult.Error != nil {
return reportResult.Error return reportResult.Error
} }
@@ -232,7 +232,7 @@ func (s *store) DeleteNode(nodeID string) error {
}) })
} }
func applyNodeFilters(q *gorm.DB, opts listOptions) *gorm.DB { func applyNodeFilters(q *gorm.DB, opts ListOptions) *gorm.DB {
if opts.NodeID != "" { if opts.NodeID != "" {
q = q.Where("node_id = ?", opts.NodeID) q = q.Where("node_id = ?", opts.NodeID)
} }
@@ -245,7 +245,7 @@ func applyNodeFilters(q *gorm.DB, opts listOptions) *gorm.DB {
return q return q
} }
func applyMapReportFilters(q *gorm.DB, opts listOptions) *gorm.DB { func applyMapReportFilters(q *gorm.DB, opts ListOptions) *gorm.DB {
q = applyNodeFilters(q, opts) q = applyNodeFilters(q, opts)
if opts.MinLat != nil && opts.MaxLat != nil { if opts.MinLat != nil && opts.MaxLat != nil {
q = q.Where("latitude IS NOT NULL AND latitude >= ? AND latitude <= ?", *opts.MinLat, *opts.MaxLat) q = q.Where("latitude IS NOT NULL AND latitude >= ? AND latitude <= ?", *opts.MinLat, *opts.MaxLat)
@@ -260,15 +260,15 @@ func applyMapReportFilters(q *gorm.DB, opts listOptions) *gorm.DB {
return q return q
} }
func (s *store) ListTextMessages(opts listOptions) ([]textMessageRecord, error) { func (s *Store) ListTextMessages(opts ListOptions) ([]TextMessageRecord, error) {
var rows []textMessageRecord var rows []TextMessageRecord
return rows, s.listAppendRows(opts, &rows).Error return rows, s.listAppendRows(opts, &rows).Error
} }
func (s *store) ListDiscardDetails(opts listOptions) ([]discardDetailsRecord, error) { func (s *Store) ListDiscardDetails(opts ListOptions) ([]DiscardDetailsRecord, error) {
opts = normalizeListOptions(opts) opts = NormalizeListOptions(opts)
var rows []discardDetailsRecord var rows []DiscardDetailsRecord
q := applyDiscardDetailsFilters(s.db.Model(&discardDetailsRecord{}), opts). q := applyDiscardDetailsFilters(s.db.Model(&DiscardDetailsRecord{}), opts).
Order("created_at DESC"). Order("created_at DESC").
Order("id DESC"). Order("id DESC").
Limit(opts.Limit). Limit(opts.Limit).
@@ -276,13 +276,13 @@ func (s *store) ListDiscardDetails(opts listOptions) ([]discardDetailsRecord, er
return rows, q.Find(&rows).Error return rows, q.Find(&rows).Error
} }
func (s *store) CountDiscardDetails(opts listOptions) (int64, error) { func (s *Store) CountDiscardDetails(opts ListOptions) (int64, error) {
var total int64 var total int64
q := applyDiscardDetailsFilters(s.db.Model(&discardDetailsRecord{}), opts) q := applyDiscardDetailsFilters(s.db.Model(&DiscardDetailsRecord{}), opts)
return total, q.Count(&total).Error return total, q.Count(&total).Error
} }
func applyDiscardDetailsFilters(q *gorm.DB, opts listOptions) *gorm.DB { func applyDiscardDetailsFilters(q *gorm.DB, opts ListOptions) *gorm.DB {
if opts.Since != nil { if opts.Since != nil {
q = q.Where("created_at >= ?", *opts.Since) q = q.Where("created_at >= ?", *opts.Since)
} }
@@ -292,8 +292,8 @@ func applyDiscardDetailsFilters(q *gorm.DB, opts listOptions) *gorm.DB {
return q return q
} }
func (s *store) DeleteTextMessage(id uint64) error { func (s *Store) DeleteTextMessage(id uint64) error {
result := s.db.Where("id = ?", id).Delete(&textMessageRecord{}) result := s.db.Where("id = ?", id).Delete(&TextMessageRecord{})
if result.Error != nil { if result.Error != nil {
return result.Error return result.Error
} }
@@ -303,28 +303,28 @@ func (s *store) DeleteTextMessage(id uint64) error {
return nil return nil
} }
func (s *store) ListPositions(opts listOptions) ([]positionRecord, error) { func (s *Store) ListPositions(opts ListOptions) ([]PositionRecord, error) {
var rows []positionRecord var rows []PositionRecord
return rows, s.listAppendRows(opts, &rows).Error return rows, s.listAppendRows(opts, &rows).Error
} }
func (s *store) ListTelemetry(opts listOptions) ([]telemetryRecord, error) { func (s *Store) ListTelemetry(opts ListOptions) ([]TelemetryRecord, error) {
var rows []telemetryRecord var rows []TelemetryRecord
return rows, s.listAppendRows(opts, &rows).Error return rows, s.listAppendRows(opts, &rows).Error
} }
func (s *store) ListRouting(opts listOptions) ([]routingRecord, error) { func (s *Store) ListRouting(opts ListOptions) ([]RoutingRecord, error) {
var rows []routingRecord var rows []RoutingRecord
return rows, s.listAppendRows(opts, &rows).Error return rows, s.listAppendRows(opts, &rows).Error
} }
func (s *store) ListTraceroute(opts listOptions) ([]tracerouteRecord, error) { func (s *Store) ListTraceroute(opts ListOptions) ([]TracerouteRecord, error) {
var rows []tracerouteRecord var rows []TracerouteRecord
return rows, s.listAppendRows(opts, &rows).Error return rows, s.listAppendRows(opts, &rows).Error
} }
func (s *store) listAppendRows(opts listOptions, dest any) *gorm.DB { func (s *Store) listAppendRows(opts ListOptions, dest any) *gorm.DB {
opts = normalizeListOptions(opts) opts = NormalizeListOptions(opts)
q := s.db.Order("created_at DESC").Order("id DESC").Limit(opts.Limit).Offset(opts.Offset) q := s.db.Order("created_at DESC").Order("id DESC").Limit(opts.Limit).Offset(opts.Offset)
if opts.NodeID != "" { if opts.NodeID != "" {
q = q.Where("from_id = ?", opts.NodeID) q = q.Where("from_id = ?", opts.NodeID)
+45
View File
@@ -0,0 +1,45 @@
package store
import (
"strings"
"golang.org/x/crypto/bcrypt"
)
// 测试 helper —— 为从 main 包搬过来的测试提供它们原本依赖的小写函数。
// 这些 helper 不暴露给生产代码使用;它们的行为应当与 main 包对应实现保持一致。
// verifyPassword 复刻 auth.go 中的 bcrypt 校验,用于 user_store 的测试。
func verifyPassword(hash, password string) bool {
return bcrypt.CompareHashAndPassword([]byte(hash), []byte(password)) == nil
}
// publicMapTileSourceDTO 复刻 admin_map_source_routes.go 中的同名函数,
// 仅供 map_source_store_test.go 验证 ProxyEnabled 时 URL 是否被改写。
// 这里返回 map[string]any 而非 gin.H 以避免引入 gin 依赖。
func publicMapTileSourceDTO(row MapTileSourceRecord) map[string]any {
urlTemplate := row.URLTemplate
if row.ProxyEnabled {
hash := row.URLTemplateHash
if hash == "" {
hash = MapTileSourceHash(row.URLTemplate)
}
urlTemplate = "/api/map/" + hash + "?x={x}&y={y}&z={z}"
}
return map[string]any{
"id": row.ID,
"name": row.Name,
"url_template": urlTemplate,
"attribution": row.Attribution,
"max_zoom": row.MaxZoom,
"enabled": row.Enabled,
"is_default": row.IsDefault,
"proxy_enabled": row.ProxyEnabled,
}
}
// newDBWriteQueue 是 db_write_queue_test.go 期望的旧名字。重新导出供测试使用。
var newDBWriteQueue = NewWriteQueue
// 让 strings 不会被 import-but-not-used(如果上面用不到,就算了——保留以应对将来扩展)
var _ = strings.TrimSpace
+16 -16
View File
@@ -1,4 +1,4 @@
package main package store
import ( import (
"errors" "errors"
@@ -9,30 +9,30 @@ import (
"gorm.io/gorm" "gorm.io/gorm"
) )
var errUserAlreadyExists = errors.New("user already exists") var ErrUserAlreadyExists = errors.New("user already exists")
func (s *store) GetUserByUsername(username string) (*userRecord, error) { func (s *Store) GetUserByUsername(username string) (*UserRecord, error) {
var user userRecord var user UserRecord
if err := s.db.Where("username = ?", username).Take(&user).Error; err != nil { if err := s.db.Where("username = ?", username).Take(&user).Error; err != nil {
return nil, err return nil, err
} }
return &user, nil return &user, nil
} }
func (s *store) GetUserByID(id uint64) (*userRecord, error) { func (s *Store) GetUserByID(id uint64) (*UserRecord, error) {
var user userRecord var user UserRecord
if err := s.db.Where("id = ?", id).Take(&user).Error; err != nil { if err := s.db.Where("id = ?", id).Take(&user).Error; err != nil {
return nil, err return nil, err
} }
return &user, nil return &user, nil
} }
func (s *store) ListUsers() ([]userRecord, error) { func (s *Store) ListUsers() ([]UserRecord, error) {
var users []userRecord var users []UserRecord
return users, s.db.Order("id ASC").Find(&users).Error return users, s.db.Order("id ASC").Find(&users).Error
} }
func (s *store) CreateAdminUser(username, password string) (*userRecord, error) { func (s *Store) CreateAdminUser(username, password string) (*UserRecord, error) {
username = strings.TrimSpace(username) username = strings.TrimSpace(username)
if username == "" { if username == "" {
return nil, fmt.Errorf("username is required") return nil, fmt.Errorf("username is required")
@@ -41,7 +41,7 @@ func (s *store) CreateAdminUser(username, password string) (*userRecord, error)
return nil, fmt.Errorf("password is required") return nil, fmt.Errorf("password is required")
} }
if _, err := s.GetUserByUsername(username); err == nil { if _, err := s.GetUserByUsername(username); err == nil {
return nil, errUserAlreadyExists return nil, ErrUserAlreadyExists
} else if !errors.Is(err, gorm.ErrRecordNotFound) { } else if !errors.Is(err, gorm.ErrRecordNotFound) {
return nil, err return nil, err
} }
@@ -49,14 +49,14 @@ func (s *store) CreateAdminUser(username, password string) (*userRecord, error)
if err != nil { if err != nil {
return nil, fmt.Errorf("hash user password: %w", err) return nil, fmt.Errorf("hash user password: %w", err)
} }
user := userRecord{Username: username, PasswordHash: hash, Role: adminRole} user := UserRecord{Username: username, PasswordHash: hash, Role: AdminRole}
if err := s.db.Create(&user).Error; err != nil { if err := s.db.Create(&user).Error; err != nil {
return nil, err return nil, err
} }
return &user, nil return &user, nil
} }
func (s *store) UpdateUserPassword(id uint64, password string) (*userRecord, error) { func (s *Store) UpdateUserPassword(id uint64, password string) (*UserRecord, error) {
if id == 0 { if id == 0 {
return nil, fmt.Errorf("user id is required") return nil, fmt.Errorf("user id is required")
} }
@@ -71,15 +71,15 @@ func (s *store) UpdateUserPassword(id uint64, password string) (*userRecord, err
if err != nil { if err != nil {
return nil, fmt.Errorf("hash user password: %w", err) return nil, fmt.Errorf("hash user password: %w", err)
} }
if err := s.db.Model(&userRecord{}).Where("id = ?", id).Updates(map[string]any{"password_hash": hash, "updated_at": time.Now()}).Error; err != nil { if err := s.db.Model(&UserRecord{}).Where("id = ?", id).Updates(map[string]any{"password_hash": hash, "updated_at": time.Now()}).Error; err != nil {
return nil, err return nil, err
} }
user.PasswordHash = hash user.PasswordHash = hash
return s.GetUserByID(id) return s.GetUserByID(id)
} }
func (s *store) EnsureDefaultAdmin(username, password string) error { func (s *Store) EnsureDefaultAdmin(username, password string) error {
var existing userRecord var existing UserRecord
err := s.db.Where("username = ?", username).Take(&existing).Error err := s.db.Where("username = ?", username).Take(&existing).Error
if err == nil { if err == nil {
return nil return nil
@@ -91,7 +91,7 @@ func (s *store) EnsureDefaultAdmin(username, password string) error {
if err != nil { if err != nil {
return fmt.Errorf("hash admin password: %w", err) return fmt.Errorf("hash admin password: %w", err)
} }
user := userRecord{Username: username, PasswordHash: hash, Role: adminRole} user := UserRecord{Username: username, PasswordHash: hash, Role: AdminRole}
if err := s.db.Create(&user).Error; err != nil { if err := s.db.Create(&user).Error; err != nil {
return fmt.Errorf("create default admin user: %w", err) return fmt.Errorf("create default admin user: %w", err)
} }
+4 -4
View File
@@ -207,7 +207,7 @@ func parseArgs() (*config, error) {
if err != nil { if err != nil {
return nil, err return nil, err
} }
cfg.key = key cfg.Key = key
return cfg, nil return cfg, nil
} }
@@ -238,7 +238,7 @@ func run(cfg *config) error {
if err != nil { if err != nil {
return err return err
} }
botSender := newBotService(store, server, cfg.key) botSender := newBotService(store, server, cfg.Key)
mqttHook.autoAcker = botSender.MaybeAutoAck mqttHook.autoAcker = botSender.MaybeAutoAck
botCtx, stopBotBroadcaster := context.WithCancel(context.Background()) botCtx, stopBotBroadcaster := context.WithCancel(context.Background())
defer stopBotBroadcaster() defer stopBotBroadcaster()
@@ -301,7 +301,7 @@ func run(cfg *config) error {
DataDir: cfg.DataDir, DataDir: cfg.DataDir,
Enabled: cfg.AI.Enabled, Enabled: cfg.AI.Enabled,
ToolConfigStore: store, ToolConfigStore: store,
}, store.db, botSenderAdapter) }, store.DB(), botSenderAdapter)
if err != nil { if err != nil {
fmt.Fprintf(os.Stderr, "Warning: failed to initialize AI service: %v\n", err) fmt.Fprintf(os.Stderr, "Warning: failed to initialize AI service: %v\n", err)
} else { } else {
@@ -384,7 +384,7 @@ func startMQTTServer(cfg *config, store *store, dbQueue *dbWriteQueue, stats *me
return nil, nil, "", err return nil, nil, "", err
} }
hook := &meshtasticFilterHook{ hook := &meshtasticFilterHook{
key: cfg.key, key: cfg.Key,
dbQueue: dbQueue, dbQueue: dbQueue,
stats: stats, stats: stats,
blocking: blocking, blocking: blocking,
+145
View File
@@ -0,0 +1,145 @@
package main
// 桥接到 internal/store —— 让根目录其余文件无须修改即可继续使用旧的小写类型/函数名。
// 当各领域包逐步迁出根目录后,可以删除这些别名。
import (
storepkg "meshtastic_mqtt_server/internal/store"
)
// ---- 类型别名 ----
type (
store = storepkg.Store
mqttClientInfo = storepkg.MQTTClientInfo
dbWriteQueue = storepkg.WriteQueue
listOptions = storepkg.ListOptions
mapReportViewportOptions = storepkg.MapReportViewportOptions
mapReportViewportResult = storepkg.MapReportViewportResult
mapReportClusterRecord = storepkg.MapReportClusterRecord
userRecord = storepkg.UserRecord
loginLogRecord = storepkg.LoginLogRecord
helpContentRecord = storepkg.HelpContentRecord
runtimeSettingRecord = storepkg.RuntimeSettingRecord
mapTileSourceRecord = storepkg.MapTileSourceRecord
mapTileSourceInput = storepkg.MapTileSourceInput
discardDetailsRecord = storepkg.DiscardDetailsRecord
signRecord = storepkg.SignRecord
nodeBlockingRecord = storepkg.NodeBlockingRecord
ipBlockingRecord = storepkg.IPBlockingRecord
forbiddenWordBlockingRecord = storepkg.ForbiddenWordBlockingRecord
mqttForwarderRecord = storepkg.MQTTForwarderRecord
mqttForwarderInput = storepkg.MQTTForwarderInput
mqttForwardTopicRecord = storepkg.MQTTForwardTopicRecord
mqttForwardTopicInput = storepkg.MQTTForwardTopicInput
botNodeRecord = storepkg.BotNodeRecord
botNodeInput = storepkg.BotNodeInput
botMessageRecord = storepkg.BotMessageRecord
botDirectMessageRecord = storepkg.BotDirectMessageRecord
llmMessageQueueRecord = storepkg.LLMMessageQueueRecord
nodeInfoRecord = storepkg.NodeInfoRecord
mapReportRecord = storepkg.MapReportRecord
textMessageRecord = storepkg.TextMessageRecord
llmProviderRecord = storepkg.LLMProviderRecord
llmToolRouterRecord = storepkg.LLMToolRouterRecord
llmPrimaryConfigRecord = storepkg.LLMPrimaryConfigRecord
positionRecord = storepkg.PositionRecord
telemetryRecord = storepkg.TelemetryRecord
routingRecord = storepkg.RoutingRecord
tracerouteRecord = storepkg.TracerouteRecord
)
// AppendPacketFields / MQTTClientRecordFields 已经是导出名,不需要别名。
// 直接用包限定名供 root 文件调用:
type (
AppendPacketFields = storepkg.AppendPacketFields
MQTTClientRecordFields = storepkg.MQTTClientRecordFields
)
// 其它额外导出类型的别名(旧的小写形式仍被根目录文件直接使用)。
type (
runtimeSettingsSnapshot = storepkg.RuntimeSettingsSnapshot
mqttForwarderConfig = storepkg.MQTTForwarderConfig
botDirectMessageListOptions = storepkg.BotDirectMessageListOptions
botMessageListOptions = storepkg.BotMessageListOptions
botDirectConversation = storepkg.BotDirectConversation
signDayCount = storepkg.SignDayCount
)
var errBlockingAlreadyExists = storepkg.ErrBlockingAlreadyExists
var errBotNodeAlreadyExists = storepkg.ErrBotNodeAlreadyExists
var (
errMapTileSourceAlreadyExists = storepkg.ErrMapTileSourceAlreadyExists
errMapTileSourceCannotDeleteDefault = storepkg.ErrMapTileSourceCannotDeleteDefault
errMapTileSourceCannotDisableDefault = storepkg.ErrMapTileSourceCannotDisableDefault
errMapTileSourceDefaultMustBeEnabled = storepkg.ErrMapTileSourceDefaultMustBeEnabled
)
func mapTileSourceHash(urlTemplate string) string {
return storepkg.MapTileSourceHash(urlTemplate)
}
const (
forbiddenWordMatchContains = storepkg.ForbiddenWordMatchContains
runtimeSettingAllowEncryptedForwarding = storepkg.RuntimeSettingAllowEncryptedForwarding
runtimeSettingLLMQueueEnabled = storepkg.RuntimeSettingLLMQueueEnabled
runtimeSettingLLMQueueIncludeChannel = storepkg.RuntimeSettingLLMQueueIncludeChannel
botDefaultPSK = storepkg.BotDefaultPSK
botMessageTypeChannel = storepkg.BotMessageTypeChannel
botMessageTypeDirect = storepkg.BotMessageTypeDirect
botMessageStatusPending = storepkg.BotMessageStatusPending
botMessageStatusPublished = storepkg.BotMessageStatusPublished
botMessageStatusFailed = storepkg.BotMessageStatusFailed
botDirectMessageDirectionInbound = storepkg.BotDirectMessageDirectionInbound
botDirectMessageDirectionOutbound = storepkg.BotDirectMessageDirectionOutbound
botDefaultTopicPrefix = storepkg.BotDefaultTopicPrefix
mqttForwardDirectionSourceToTarget = storepkg.MQTTForwardDirectionSourceToTarget
mqttForwardDirectionBidirectional = storepkg.MQTTForwardDirectionBidirectional
)
const botDefaultNodeInfoBroadcastSeconds = storepkg.BotDefaultNodeInfoBroadcastSeconds
func validateBotNodeNum(n int64) error { return storepkg.ValidateBotNodeNum(n) }
var errUserAlreadyExists = storepkg.ErrUserAlreadyExists
var (
errMQTTForwarderAlreadyExists = storepkg.ErrMQTTForwarderAlreadyExists
errMQTTForwardTopicAlreadyExists = storepkg.ErrMQTTForwardTopicAlreadyExists
)
func nullableString(v any) *string { return storepkg.NullableString(v) }
func nullableStringValue(v any) *string { return storepkg.NullableStringValue(v) }
func decodeBotPublicKey(row botNodeRecord) ([]byte, error) {
return storepkg.DecodeBotPublicKey(row)
}
// LLM 消息队列状态字符串常量。
const (
llmMessageStatusPending = storepkg.LLMMessageStatusPending
llmMessageStatusProcessing = storepkg.LLMMessageStatusProcessing
llmMessageStatusProcessed = storepkg.LLMMessageStatusProcessed
llmMessageStatusError = storepkg.LLMMessageStatusError
)
const defaultHelpMarkdown = storepkg.DefaultHelpMarkdown
// 旧名 llmMessageDTO 现在通过 store 提供。
func llmMessageDTO(row llmMessageQueueRecord) map[string]any {
return storepkg.LLMMessageDTO(row)
}
// ---- 工厂函数包装 ----
func openStore(cfg databaseConfig) (*store, error) { return storepkg.OpenStore(cfg) }
func normalizeListOptions(o listOptions) listOptions {
return storepkg.NormalizeListOptions(o)
}
func normalizeMapReportViewportOptions(o mapReportViewportOptions) mapReportViewportOptions {
return storepkg.NormalizeMapReportViewportOptions(o)
}
// newDBWriteQueue 在新包里更名,提供旧名字避免改 main.go。
func newDBWriteQueue(s *store) *dbWriteQueue { return storepkg.NewWriteQueue(s) }
+21
View File
@@ -0,0 +1,21 @@
package main
import (
"path/filepath"
"testing"
storepkg "meshtastic_mqtt_server/internal/store"
)
// openTestStore 在根目录的测试中沿用旧函数名,但底层调用 internal/store 的实现。
func openTestStore(t *testing.T) *store {
t.Helper()
st, err := storepkg.OpenStore(databaseConfig{
Driver: databaseDriverSQLite,
SQLite: sqliteConfig{Path: filepath.Join(t.TempDir(), "mesh_mqtt_go.db")},
})
if err != nil {
t.Fatalf("OpenStore() error = %v", err)
}
return st
}