25 Commits
Author SHA1 Message Date
kevin 92e766dfa8 修复聊天消息不显示最新数据:所有列表查询排序从 created_at DESC 改为 id DESC,记录创建函数显式设置 CreatedAt,后端 v1.2.1 / 前端 v0.3.1 2026-08-07 22:21:26 +08:00
kevin b6fabf59b7 丢弃数据页批量管理:多选删除+一键清空(ConfirmDeleteModal 二次确认),后端 v1.2.0 / 前端 v0.3.0 2026-08-07 21:36:16 +08:00
kevin 62b6ff44d1 频道自动补全:channels 表(DBv2 回填去重) + GET /api/channels + 聊天卡片输入框下拉补全(键盘导航+点击选中) 2026-08-07 19:55:32 +08:00
kevin 0e622329a7 修复频道筛选框被遮挡:header+filter 包裹在 sticky 容器中一起固定 2026-08-07 19:38:00 +08:00
kevin c107f6d6af 修复聊天卡片频道筛选框被 sticky header 遮挡的问题 2026-08-07 19:34:52 +08:00
kevin d2b5832df3 频道筛选框从 topbar 移至聊天卡片内,ChatPanel v-model:channelFilter 双向绑定 2026-08-07 19:31:33 +08:00
kevin 4ce341208b 版本化数据库迁移系统:schema_migrations 表 + DBMigration 注册表,DBv1 修复 channel_id TEXT->VARCHAR(255) 解决 MySQL Error 1170 2026-08-07 19:26:18 +08:00
kevin d919ea2c53 频道筛选:text_message 表增加 (channel_id, created_at) 复合索引,主页增加频道筛选输入框,服务端筛选+防抖重载,版本号 backend 1.1.0 / frontend 0.2.0 2026-08-07 19:10:49 +08:00
kevin 1c9f94dcf8 footer增加GitHub图标链接 2026-08-07 18:17:51 +08:00
kevin 362c8c7a57 节点列表去除横向滚动,公钥截断显示+复制按钮 2026-08-07 18:09:28 +08:00
kevin ea2d33d1bf 加入版本显示 2026-08-07 17:37:58 +08:00
kevin 01aee47024 删除所有测试文件 2026-08-07 16:26:30 +08:00
kevin 7bc2e53ce6 屏蔽机器人消息 2026-07-02 11:02:48 +08:00
kevin 6052bf90ec 调用签到🔧注入时间 2026-07-01 12:43:05 +08:00
kevin 16d0d0ec0b up 2026-07-01 12:36:55 +08:00
kevin f11c2ed138 回复ai回复emoji的能力 2026-07-01 12:24:38 +08:00
kevin 04e105c6ba 更新迁移脚本 2026-07-01 12:12:23 +08:00
kevin 99fb474bcf ai服务状态更新 2026-07-01 11:55:37 +08:00
kevin f6fa167d76 更新ai服务提醒 2026-07-01 11:26:08 +08:00
kevin 716f711373 up 安装脚本 2026-07-01 11:05:40 +08:00
kevin a75ab812d2 添加 MySQL 8 到 MySQL 5.7 数据库迁移脚本 2026-06-30 12:36:43 +08:00
kevin 01c7275763 admin 页面 MQTT 服务状态增加去重队列长度显示 2026-06-30 12:18:40 +08:00
kevin 4782e84c15 添加 MQTT 消息去重队列:Hook 层基于 payload+topic hash 去重,TTL 15 秒,定时清理 2026-06-30 12:10:29 +08:00
kevin ca59c5f316 修改机器人hops为7 2026-06-29 10:43:18 +08:00
kevinandClaude Fable 5 6620192322 签到检查功能增强:返回签到时间和内容
问题:
- 用户询问'我什么时候签到的'时,AI 说'系统只记录是否签到,没有时间戳'
- 实际上数据库 signs 表有完整的 sign_time 和 sign_text 字段

解决方案:
- 修改 executeCheck 方法,改用 ListSigns 查询今天的签到记录
- 返回完整的签到详情:签到时间(HH:MM:SS格式)+ 签到内容
- 未签到时仍返回简单提示

改动内容:
- executeCheck 使用 ListSigns 替代 HasSignedOnDay
- 构建 ListOptions 查询今天的签到记录(过滤 NodeID + 时间范围)
- 返回格式:'XXX 今天已经签到过了。\n签到时间:10:30:45\n签到内容:上海-Kevin-GAT562签到'
- 更新测试验证返回内容包含时间和签到文本
- 更新 mockSignStore.ListSigns 支持 NodeID 和 Limit 过滤

使用效果:
- 用户:'我今天签到了吗?' → 返回是否签到 + 时间 + 内容
- 用户:'我什么时候签到的?' → 返回签到时间:10:30:45

测试:
-  未签到场景测试通过
-  已签到场景测试通过(验证时间和内容)
-  所有签到工具测试通过
-  项目编译成功

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-06-23 21:47:16 +08:00
55 changed files with 1611 additions and 4126 deletions
+10 -4
View File
@@ -32,9 +32,9 @@
### 3. 检查操作 (action=check) **新增** ### 3. 检查操作 (action=check) **新增**
检查当前节点今天是否已签到: 检查当前节点今天是否已签到:
- 用于回答"我今天签到了吗"之类的问题 - 用于回答"我今天签到了吗"、"我什么时候签到的"之类的问题
- 直接查询数据库,不依赖对话历史 - 直接查询数据库,不依赖对话历史
- 返回明确的已签到/未签到状态 - 返回明确的签到状态**包括签到时间和签到内容**
## 修改内容 ## 修改内容
@@ -102,15 +102,21 @@ AI 调用:
``` ```
返回(未签到): 返回(未签到):
``` ```text
Test Node 今天还没有签到。 Test Node 今天还没有签到。
``` ```
返回(已签到): 返回(已签到):
``` ```text
Test Node 今天已经签到过了。 Test Node 今天已经签到过了。
签到时间:10:30:45
签到内容:上海闵行-Kevin-GAT562签到
``` ```
**用户**:"我什么时候签到的?"
AI 同样调用 check 操作,返回包含签到时间和内容的完整信息。
#### 查询今天的签到情况 #### 查询今天的签到情况
```json ```json
{ {
+14 -6
View File
@@ -3,6 +3,7 @@ set -euo pipefail
SERVICE_NAME="mesh_mqtt_go" SERVICE_NAME="mesh_mqtt_go"
SERVICE_USER="mesh_mqtt_go" SERVICE_USER="mesh_mqtt_go"
SERVICE_GROUP="${SERVICE_USER}"
CONFIG_DIR="/etc/${SERVICE_NAME}" CONFIG_DIR="/etc/${SERVICE_NAME}"
DATA_DIR="/srv/${SERVICE_NAME}" DATA_DIR="/srv/${SERVICE_NAME}"
INSTALL_DIR="/opt/${SERVICE_NAME}" INSTALL_DIR="/opt/${SERVICE_NAME}"
@@ -17,11 +18,18 @@ if [[ "${EUID}" -ne 0 ]]; then
exit 1 exit 1
fi fi
if id -u "www" >/dev/null 2>&1; then
SERVICE_USER="www"
SERVICE_GROUP=$(id -gn "www")
echo "检测到 www 用户,将以 www 用户运行服务"
fi
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
cd "${SCRIPT_DIR}" cd "${SCRIPT_DIR}"
echo "拉取最新代码..." echo "拉取最新代码..."
git pull git pull
COMMIT_HASH=$(git rev-parse --short HEAD)
echo "编译前端..." echo "编译前端..."
cd "${SCRIPT_DIR}/${FRONTEND_DIR}" cd "${SCRIPT_DIR}/${FRONTEND_DIR}"
@@ -34,7 +42,7 @@ npm run build
echo "编译 Go 程序..." echo "编译 Go 程序..."
cd "${SCRIPT_DIR}" cd "${SCRIPT_DIR}"
go build -o "${BINARY_NAME}" . go build -ldflags "-X meshtastic_mqtt_server/internal/web.CommitVersion=${COMMIT_HASH}" -o "${BINARY_NAME}" .
echo "检查系统用户..." echo "检查系统用户..."
if ! id -u "${SERVICE_USER}" >/dev/null 2>&1; then if ! id -u "${SERVICE_USER}" >/dev/null 2>&1; then
@@ -42,8 +50,8 @@ if ! id -u "${SERVICE_USER}" >/dev/null 2>&1; then
fi fi
echo "创建目录..." echo "创建目录..."
install -d -m 0750 -o "${SERVICE_USER}" -g "${SERVICE_USER}" "${CONFIG_DIR}" "${DATA_DIR}" install -d -m 0750 -o "${SERVICE_USER}" -g "${SERVICE_GROUP}" "${CONFIG_DIR}" "${DATA_DIR}"
install -d -m 0755 -o "${SERVICE_USER}" -g "${SERVICE_USER}" "${INSTALL_DIR}" install -d -m 0755 -o "${SERVICE_USER}" -g "${SERVICE_GROUP}" "${INSTALL_DIR}"
echo "安装程序和前端文件..." echo "安装程序和前端文件..."
install -m 0755 -o root -g root "${SCRIPT_DIR}/${BINARY_NAME}" "${INSTALL_DIR}/${BINARY_NAME}" install -m 0755 -o root -g root "${SCRIPT_DIR}/${BINARY_NAME}" "${INSTALL_DIR}/${BINARY_NAME}"
@@ -51,7 +59,7 @@ rm -rf "${INSTALL_DIR}/dist"
cp -a "${SCRIPT_DIR}/${FRONTEND_DIST_DIR}" "${INSTALL_DIR}/dist" cp -a "${SCRIPT_DIR}/${FRONTEND_DIST_DIR}" "${INSTALL_DIR}/dist"
chown root:root "${INSTALL_DIR}/${BINARY_NAME}" chown root:root "${INSTALL_DIR}/${BINARY_NAME}"
chown -R root:root "${INSTALL_DIR}/dist" chown -R root:root "${INSTALL_DIR}/dist"
chown "${SERVICE_USER}:${SERVICE_USER}" "${INSTALL_DIR}" chown "${SERVICE_USER}:${SERVICE_GROUP}" "${INSTALL_DIR}"
chmod 0755 "${INSTALL_DIR}" chmod 0755 "${INSTALL_DIR}"
find "${INSTALL_DIR}/dist" -type d -exec chmod 0755 {} \; find "${INSTALL_DIR}/dist" -type d -exec chmod 0755 {} \;
find "${INSTALL_DIR}/dist" -type f -exec chmod 0644 {} \; find "${INSTALL_DIR}/dist" -type f -exec chmod 0644 {} \;
@@ -91,7 +99,7 @@ console_log:
sql: true sql: true
meshtastic: true meshtastic: true
EOF EOF
chown "${SERVICE_USER}:${SERVICE_USER}" "${CONFIG_DIR}/config.yaml" chown "${SERVICE_USER}:${SERVICE_GROUP}" "${CONFIG_DIR}/config.yaml"
chmod 0640 "${CONFIG_DIR}/config.yaml" chmod 0640 "${CONFIG_DIR}/config.yaml"
fi fi
@@ -105,7 +113,7 @@ Wants=network-online.target
[Service] [Service]
Type=simple Type=simple
User=${SERVICE_USER} User=${SERVICE_USER}
Group=${SERVICE_USER} Group=${SERVICE_GROUP}
WorkingDirectory=${INSTALL_DIR} WorkingDirectory=${INSTALL_DIR}
ExecStart=${INSTALL_DIR}/${BINARY_NAME} -web-socket-path ${SOCKET_PATH} -web-static-dir ${INSTALL_DIR}/dist ExecStart=${INSTALL_DIR}/${BINARY_NAME} -web-socket-path ${SOCKET_PATH} -web-static-dir ${INSTALL_DIR}/dist
Restart=on-failure Restart=on-failure
-187
View File
@@ -1,187 +0,0 @@
package active
import (
"context"
"encoding/json"
"testing"
"time"
"meshtastic_mqtt_server/internal/agenttool"
)
// mockActiveStore 是用于测试的 mock store
type mockActiveStore struct {
activeNodeCount int64
activeUserCount int64
}
func (m *mockActiveStore) CountActiveNodes(since time.Time) (int64, error) {
return m.activeNodeCount, nil
}
func (m *mockActiveStore) CountActiveUsers(since time.Time) (int64, error) {
return m.activeUserCount, nil
}
func TestActiveTool_Query(t *testing.T) {
now := time.Date(2024, 6, 23, 12, 0, 0, 0, time.UTC)
store := &mockActiveStore{
activeNodeCount: 25,
activeUserCount: 15,
}
tool := &Tool{
enabled: true,
store: store,
}
tests := []struct {
name string
hours float64
queryType string
expectNodes bool
expectUsers bool
expectError bool
}{
{
name: "默认查询1小时(both",
hours: 0, // 0 表示使用默认值
queryType: "",
expectNodes: true,
expectUsers: true,
expectError: false,
},
{
name: "查询6小时",
hours: 6,
queryType: "both",
expectNodes: true,
expectUsers: true,
expectError: false,
},
{
name: "仅查询节点",
hours: 1,
queryType: "nodes",
expectNodes: true,
expectUsers: false,
expectError: false,
},
{
name: "仅查询人数",
hours: 1,
queryType: "users",
expectNodes: false,
expectUsers: true,
expectError: false,
},
{
name: "查询24小时(最大值)",
hours: 24,
queryType: "both",
expectNodes: true,
expectUsers: true,
expectError: false,
},
{
name: "超过24小时应限制到24小时",
hours: 48,
queryType: "both",
expectNodes: true,
expectUsers: true,
expectError: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
params := activeParams{
Hours: tt.hours,
QueryType: tt.queryType,
}
argsJSON, _ := json.Marshal(params)
runtime := agenttool.Runtime{Now: now}
result, err := tool.Execute(context.Background(), string(argsJSON), runtime)
if tt.expectError && err == nil {
t.Errorf("Expected error but got none")
}
if !tt.expectError && err != nil {
t.Errorf("Unexpected error: %v", err)
}
if !tt.expectError {
t.Logf("Query result:\n%s", result)
// 验证结果包含预期的内容
if tt.expectNodes && result != "" {
// 应该包含节点统计
if !contains(result, "活跃节点") {
t.Errorf("Expected result to contain node count")
}
}
if tt.expectUsers && result != "" {
// 应该包含人数统计
if !contains(result, "活跃人数") {
t.Errorf("Expected result to contain user count")
}
}
}
})
}
}
func TestActiveTool_Enabled(t *testing.T) {
// 测试工具启用状态
tests := []struct {
name string
enabled bool
store ActiveStore
expect bool
}{
{
name: "启用且有store",
enabled: true,
store: &mockActiveStore{},
expect: true,
},
{
name: "启用但无store",
enabled: true,
store: nil,
expect: false,
},
{
name: "禁用且有store",
enabled: false,
store: &mockActiveStore{},
expect: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
tool := &Tool{
enabled: tt.enabled,
store: tt.store,
}
if tool.Enabled() != tt.expect {
t.Errorf("Expected Enabled() = %v, got %v", tt.expect, tool.Enabled())
}
})
}
}
func contains(s, substr string) bool {
return len(s) > 0 && len(substr) > 0 && (s == substr || len(s) >= len(substr) && (s[:len(substr)] == substr || s[len(s)-len(substr):] == substr || containsMiddle(s, substr)))
}
func containsMiddle(s, substr string) bool {
for i := 0; i <= len(s)-len(substr); i++ {
if s[i:i+len(substr)] == substr {
return true
}
}
return false
}
+27 -7
View File
@@ -179,7 +179,7 @@ func (t *Tool) executeSign(ctx context.Context, params signParams, runtime agent
return fmt.Sprintf("签到成功!%s\n签到内容:%s", displayName(node), record.SignText), nil return fmt.Sprintf("签到成功!%s\n签到内容:%s", displayName(node), record.SignText), nil
} }
// executeCheck 检查当前节点今天是否已签到 // executeCheck 检查当前节点今天是否已签到,并返回签到详情
func (t *Tool) executeCheck(ctx context.Context, params signParams, runtime agenttool.Runtime) (string, error) { func (t *Tool) executeCheck(ctx context.Context, params signParams, runtime agenttool.Runtime) (string, error) {
// 节点身份来自消息上下文 // 节点身份来自消息上下文
node, ok := agenttool.NodeContextFromContext(ctx) node, ok := agenttool.NodeContextFromContext(ctx)
@@ -192,18 +192,38 @@ func (t *Tool) executeCheck(ctx context.Context, params signParams, runtime agen
now = time.Now() now = time.Now()
} }
// 查询数据库检查今天是否已签到 // 构建查询选项:查询今天的签到记录
signed, err := t.store.HasSignedOnDay(node.NodeID, now) loc := now.Location()
if err != nil { if loc == nil {
return "", fmt.Errorf("检查签到状态失败:%w", err) loc = time.Local
}
start := time.Date(now.Year(), now.Month(), now.Day(), 0, 0, 0, 0, loc)
end := start.AddDate(0, 0, 1)
opts := storepkg.ListOptions{
NodeID: node.NodeID,
Since: &start,
Until: &end,
Limit: 1,
} }
if signed { // 查询数据库获取今天的签到记录
return fmt.Sprintf("%s 今天已经签到过了。", displayName(node)), nil signs, err := t.store.ListSigns(opts)
if err != nil {
return "", fmt.Errorf("查询签到记录失败:%w", err)
} }
if len(signs) == 0 {
return fmt.Sprintf("%s 今天还没有签到。", displayName(node)), nil return fmt.Sprintf("%s 今天还没有签到。", displayName(node)), nil
} }
// 返回签到详情
sign := signs[0]
signTimeStr := sign.SignTime.Format("15:04:05")
return fmt.Sprintf("%s 今天已经签到过了。\n签到时间:%s\n签到内容:%s",
displayName(node), signTimeStr, sign.SignText), nil
}
// executeQuery 执行查询操作 // executeQuery 执行查询操作
func (t *Tool) executeQuery(ctx context.Context, params signParams, runtime agenttool.Runtime) (string, error) { func (t *Tool) executeQuery(ctx context.Context, params signParams, runtime agenttool.Runtime) (string, error) {
now := runtime.Now now := runtime.Now
-317
View File
@@ -1,317 +0,0 @@
package sign
import (
"context"
"encoding/json"
"testing"
"time"
"meshtastic_mqtt_server/internal/agenttool"
storepkg "meshtastic_mqtt_server/internal/store"
)
// mockSignStore 是用于测试的 mock store
type mockSignStore struct {
signs []storepkg.SignRecord
nodeInfoMap map[string]*storepkg.NodeInfoRecord
}
func (m *mockSignStore) CreateSign(nodeID string, longName, shortName *string, signText string, signTime time.Time) (*storepkg.SignRecord, error) {
record := storepkg.SignRecord{
NodeID: nodeID,
LongName: longName,
ShortName: shortName,
SignText: signText,
SignTime: signTime,
}
m.signs = append(m.signs, record)
return &record, nil
}
func (m *mockSignStore) HasSignedOnDay(nodeID string, day time.Time) (bool, error) {
loc := day.Location()
if loc == nil {
loc = time.Local
}
start := time.Date(day.Year(), day.Month(), day.Day(), 0, 0, 0, 0, loc)
end := start.AddDate(0, 0, 1)
for _, sign := range m.signs {
if sign.NodeID == nodeID && sign.SignTime.After(start) && sign.SignTime.Before(end) {
return true, nil
}
}
return false, nil
}
func (m *mockSignStore) GetNodeInfo(nodeID string) (*storepkg.NodeInfoRecord, error) {
return m.nodeInfoMap[nodeID], nil
}
func (m *mockSignStore) CountSigns(opts storepkg.ListOptions) (int64, error) {
count := int64(0)
for _, sign := range m.signs {
if opts.Since != nil && sign.SignTime.Before(*opts.Since) {
continue
}
if opts.Until != nil && sign.SignTime.After(*opts.Until) {
continue
}
count++
}
return count, nil
}
func (m *mockSignStore) CountSignsByDay(opts storepkg.ListOptions) ([]storepkg.SignDayCount, error) {
dayCounts := make(map[string]int64)
for _, sign := range m.signs {
if opts.Since != nil && sign.SignTime.Before(*opts.Since) {
continue
}
if opts.Until != nil && sign.SignTime.After(*opts.Until) {
continue
}
dateStr := sign.SignTime.Format("2006-01-02")
dayCounts[dateStr]++
}
var result []storepkg.SignDayCount
for date, count := range dayCounts {
result = append(result, storepkg.SignDayCount{Date: date, Count: count})
}
return result, nil
}
func (m *mockSignStore) ListSigns(opts storepkg.ListOptions) ([]storepkg.SignRecord, error) {
var result []storepkg.SignRecord
for _, sign := range m.signs {
if opts.Since != nil && sign.SignTime.Before(*opts.Since) {
continue
}
if opts.Until != nil && sign.SignTime.After(*opts.Until) {
continue
}
result = append(result, sign)
}
return result, nil
}
func TestSignTool_Query(t *testing.T) {
// 创建 mock store 并添加测试数据
now := time.Date(2024, 6, 23, 12, 0, 0, 0, time.UTC)
yesterday := now.AddDate(0, 0, -1)
twoDaysAgo := now.AddDate(0, 0, -2)
store := &mockSignStore{
signs: []storepkg.SignRecord{
{NodeID: "node1", SignText: "上海-Alice-Device1签到", SignTime: now},
{NodeID: "node2", SignText: "北京-Bob-Device2签到", SignTime: now},
{NodeID: "node3", SignText: "深圳-Charlie-Device3签到", SignTime: yesterday},
{NodeID: "node4", SignText: "广州-David-Device4签到", SignTime: twoDaysAgo},
},
nodeInfoMap: make(map[string]*storepkg.NodeInfoRecord),
}
tool := &Tool{
enabled: true,
store: store,
}
tests := []struct {
name string
action string
date string
days int
expectCount int64
expectError bool
}{
{
name: "查询今天",
action: "query",
date: "2024-06-23",
days: 0,
expectCount: 2,
expectError: false,
},
{
name: "查询最近3天",
action: "query",
date: "2024-06-23",
days: 3,
expectCount: 4,
expectError: false,
},
{
name: "查询昨天",
action: "query",
date: "2024-06-22",
days: 0,
expectCount: 1,
expectError: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
params := signParams{
Action: tt.action,
Date: tt.date,
Days: tt.days,
}
argsJSON, _ := json.Marshal(params)
runtime := agenttool.Runtime{Now: now}
result, err := tool.Execute(context.Background(), string(argsJSON), runtime)
if tt.expectError && err == nil {
t.Errorf("Expected error but got none")
}
if !tt.expectError && err != nil {
t.Errorf("Unexpected error: %v", err)
}
if !tt.expectError {
t.Logf("Query result:\n%s", result)
}
})
}
}
func TestSignTool_SignAction(t *testing.T) {
now := time.Date(2024, 6, 23, 12, 0, 0, 0, time.UTC)
store := &mockSignStore{
signs: []storepkg.SignRecord{},
nodeInfoMap: make(map[string]*storepkg.NodeInfoRecord),
}
tool := &Tool{
enabled: true,
store: store,
}
// 测试签到功能
params := signParams{
Action: "sign",
Region: "上海闵行",
Name: "TestUser",
Device: "TestDevice",
}
argsJSON, _ := json.Marshal(params)
// 创建带节点上下文的 context
nodeCtx := agenttool.NodeContext{
NodeID: "test_node_123",
LongName: "Test Node",
ShortName: "TN",
}
ctx := agenttool.WithNodeContext(context.Background(), nodeCtx)
runtime := agenttool.Runtime{Now: now}
result, err := tool.Execute(ctx, string(argsJSON), runtime)
if err != nil {
t.Errorf("Unexpected error: %v", err)
}
t.Logf("Sign result: %s", result)
// 验证签到记录已创建
if len(store.signs) != 1 {
t.Errorf("Expected 1 sign record, got %d", len(store.signs))
}
}
func TestSignTool_CheckAction(t *testing.T) {
now := time.Date(2024, 6, 23, 12, 0, 0, 0, time.UTC)
// 测试场景1:今天未签到
t.Run("今天未签到", func(t *testing.T) {
store := &mockSignStore{
signs: []storepkg.SignRecord{},
nodeInfoMap: make(map[string]*storepkg.NodeInfoRecord),
}
tool := &Tool{
enabled: true,
store: store,
}
params := signParams{
Action: "check",
}
argsJSON, _ := json.Marshal(params)
nodeCtx := agenttool.NodeContext{
NodeID: "test_node_123",
LongName: "Test Node",
ShortName: "TN",
}
ctx := agenttool.WithNodeContext(context.Background(), nodeCtx)
runtime := agenttool.Runtime{Now: now}
result, err := tool.Execute(ctx, string(argsJSON), runtime)
if err != nil {
t.Errorf("Unexpected error: %v", err)
}
t.Logf("Check result (not signed): %s", result)
if !contains(result, "还没有签到") && !contains(result, "没有签到") {
t.Errorf("Expected result to indicate not signed yet")
}
})
// 测试场景2:今天已签到
t.Run("今天已签到", func(t *testing.T) {
store := &mockSignStore{
signs: []storepkg.SignRecord{
{
NodeID: "test_node_123",
SignText: "上海-TestUser-TestDevice签到",
SignTime: now,
},
},
nodeInfoMap: make(map[string]*storepkg.NodeInfoRecord),
}
tool := &Tool{
enabled: true,
store: store,
}
params := signParams{
Action: "check",
}
argsJSON, _ := json.Marshal(params)
nodeCtx := agenttool.NodeContext{
NodeID: "test_node_123",
LongName: "Test Node",
ShortName: "TN",
}
ctx := agenttool.WithNodeContext(context.Background(), nodeCtx)
runtime := agenttool.Runtime{Now: now}
result, err := tool.Execute(ctx, string(argsJSON), runtime)
if err != nil {
t.Errorf("Unexpected error: %v", err)
}
t.Logf("Check result (already signed): %s", result)
if !contains(result, "已经签到") {
t.Errorf("Expected result to indicate already signed")
}
})
}
func contains(s, substr string) bool {
return len(s) > 0 && len(substr) > 0 && (s == substr || len(s) >= len(substr) && (s[:len(substr)] == substr || s[len(s)-len(substr):] == substr || containsMiddle(s, substr)))
}
func containsMiddle(s, substr string) bool {
for i := 0; i <= len(s)-len(substr); i++ {
if s[i:i+len(substr)] == substr {
return true
}
}
return false
}
+203
View File
@@ -0,0 +1,203 @@
package ai
import (
"context"
"fmt"
"sync"
"meshtastic_mqtt_server/internal/autoreply"
"meshtastic_mqtt_server/internal/llm"
storepkg "meshtastic_mqtt_server/internal/store"
"gorm.io/gorm"
)
// AIServiceStatus reports the current state of the AI service
type AIServiceStatus struct {
Running bool `json:"running"`
Enabled bool `json:"enabled"`
ProviderCount int `json:"provider_count"`
Message string `json:"message,omitempty"`
}
// AIManager manages the lifecycle of the AI service, supporting restart
type AIManager struct {
mu sync.Mutex
service *Service
cfg Config
db *gorm.DB
botSender autoreply.BotSender
ctx context.Context
store *storepkg.Store
}
// NewAIManager creates a new AIManager
func NewAIManager(cfg Config, db *gorm.DB, botSender autoreply.BotSender, ctx context.Context, store *storepkg.Store) *AIManager {
return &AIManager{
cfg: cfg,
db: db,
botSender: botSender,
ctx: ctx,
store: store,
}
}
// SetConfigEnabled sets the enabled flag on the config
func (m *AIManager) SetConfigEnabled(enabled bool) {
m.cfg.Enabled = enabled
}
// SetProviderConfigs sets the LLM provider configs
func (m *AIManager) SetProviderConfigs(configs []llm.ProviderConfig) {
m.cfg.LLMProviders = configs
}
// Init creates and starts the AI service
func (m *AIManager) Init() error {
m.mu.Lock()
defer m.mu.Unlock()
svc, err := NewService(m.cfg, m.db, m.botSender)
if err != nil {
return fmt.Errorf("failed to create AI service: %w", err)
}
if err := svc.Start(m.ctx); err != nil {
return fmt.Errorf("failed to start AI service: %w", err)
}
m.service = svc
return nil
}
// Stop stops the currently running AI service
func (m *AIManager) Stop() {
m.mu.Lock()
defer m.mu.Unlock()
if m.service != nil {
m.service.Stop()
m.service = nil
}
}
// Status returns the current AI service status
func (m *AIManager) Status() AIServiceStatus {
m.mu.Lock()
defer m.mu.Unlock()
providerCount := len(m.cfg.LLMProviders)
if m.service == nil {
msg := "AI 服务未运行,可在配置提供商后点击重启"
if providerCount == 0 {
msg = "尚未配置 AI 提供商,请先添加提供商配置"
}
return AIServiceStatus{
Running: false,
Enabled: m.cfg.Enabled,
ProviderCount: providerCount,
Message: msg,
}
}
enabled := m.service.Enabled()
if !enabled {
return AIServiceStatus{
Running: false,
Enabled: false,
ProviderCount: providerCount,
Message: "AI 服务未启用",
}
}
return AIServiceStatus{
Running: true,
Enabled: true,
ProviderCount: providerCount,
}
}
// Restart stops the current AI service and creates a new one from DB
func (m *AIManager) Restart() error {
m.mu.Lock()
defer m.mu.Unlock()
if m.service != nil {
m.service.Stop()
m.service = nil
}
providers, err := m.store.ListLLMProviders(true)
if err != nil {
return fmt.Errorf("加载 LLM 提供商列表失败: %w", err)
}
if len(providers) == 0 {
return fmt.Errorf("没有配置任何 LLM 提供商,请先添加配置")
}
providerConfigs := make([]llm.ProviderConfig, 0, len(providers))
for _, p := range providers {
providerConfigs = append(providerConfigs, llm.ProviderConfig{
Name: p.Name,
Active: p.Active,
APIKey: p.APIKey,
BaseURL: p.BaseURL,
Model: p.Model,
Timeout: p.Timeout,
ContextWindowTokens: p.ContextWindowTokens,
})
}
m.cfg.LLMProviders = providerConfigs
svc, err := NewService(m.cfg, m.db, m.botSender)
if err != nil {
return fmt.Errorf("创建 AI 服务失败: %w", err)
}
if err := svc.Start(m.ctx); err != nil {
return fmt.Errorf("启动 AI 服务失败: %w", err)
}
m.service = svc
return nil
}
// ReloadLLMProvider delegates to the current service
func (m *AIManager) ReloadLLMProvider(config interface{}) error {
m.mu.Lock()
svc := m.service
m.mu.Unlock()
if svc == nil {
return nil
}
return svc.ReloadLLMProvider(config)
}
// AddLLMProvider delegates to the current service
func (m *AIManager) AddLLMProvider(config interface{}) error {
m.mu.Lock()
svc := m.service
m.mu.Unlock()
if svc == nil {
return nil
}
return svc.AddLLMProvider(config)
}
// RemoveLLMProvider delegates to the current service
func (m *AIManager) RemoveLLMProvider(name string) error {
m.mu.Lock()
svc := m.service
m.mu.Unlock()
if svc == nil {
return nil
}
return svc.RemoveLLMProvider(name)
}
// AIServiceStatus returns the AI service status for the web interface
func (m *AIManager) AIServiceStatus() AIServiceStatus {
return m.Status()
}
// RestartAIService restarts the AI service for the web interface
func (m *AIManager) RestartAIService() error {
return m.Restart()
}
+9
View File
@@ -262,6 +262,9 @@ func (s *Service) Enabled() bool {
// ReloadLLMProvider reloads a specific LLM provider configuration // ReloadLLMProvider reloads a specific LLM provider configuration
func (s *Service) ReloadLLMProvider(config interface{}) error { func (s *Service) ReloadLLMProvider(config interface{}) error {
if s == nil {
return nil
}
if !s.enabled || s.LLMState == nil { if !s.enabled || s.LLMState == nil {
return nil return nil
} }
@@ -274,6 +277,9 @@ func (s *Service) ReloadLLMProvider(config interface{}) error {
// AddLLMProvider adds a new LLM provider // AddLLMProvider adds a new LLM provider
func (s *Service) AddLLMProvider(config interface{}) error { func (s *Service) AddLLMProvider(config interface{}) error {
if s == nil {
return nil
}
if !s.enabled || s.LLMState == nil { if !s.enabled || s.LLMState == nil {
return nil return nil
} }
@@ -286,6 +292,9 @@ func (s *Service) AddLLMProvider(config interface{}) error {
// RemoveLLMProvider removes an LLM provider // RemoveLLMProvider removes an LLM provider
func (s *Service) RemoveLLMProvider(name string) error { func (s *Service) RemoveLLMProvider(name string) error {
if s == nil {
return nil
}
if !s.enabled || s.LLMState == nil { if !s.enabled || s.LLMState == nil {
return nil return nil
} }
+10
View File
@@ -509,6 +509,16 @@ func cleanReplyText(text string) string {
sb.WriteRune(r) sb.WriteRune(r)
case r == 0x3002 || r == 0xFF1F || r == 0xFF01 || r == 0xFF0C || r == 0xFF1A: // Fullwidth punctuation case r == 0x3002 || r == 0xFF1F || r == 0xFF01 || r == 0xFF0C || r == 0xFF1A: // Fullwidth punctuation
sb.WriteRune(r) sb.WriteRune(r)
case r >= 0x1F600 && r <= 0x1F64F: // Emoticons
sb.WriteRune(r)
case r >= 0x1F300 && r <= 0x1F5FF: // Misc Symbols and Pictographs
sb.WriteRune(r)
case r >= 0x1F680 && r <= 0x1F6FF: // Transport and Map Symbols
sb.WriteRune(r)
case r >= 0x1F900 && r <= 0x1F9FF: // Supplemental Symbols and Pictographs
sb.WriteRune(r)
case r >= 0x2600 && r <= 0x27BF: // Misc Symbols + Dingbats
sb.WriteRune(r)
default: default:
continue // Skip all other characters continue // Skip all other characters
} }
-102
View File
@@ -1,102 +0,0 @@
package blocking
import "testing"
func TestBlockingCacheLoadsEnabledRules(t *testing.T) {
st := openTestStore(t)
defer st.Close()
nodeNum := int64(305419896)
if _, err := st.CreateNodeBlocking("!12345678", &nodeNum, "enabled", true); err != nil {
t.Fatalf("CreateNodeBlocking(enabled) error = %v", err)
}
disabledNodeNum := int64(7)
if _, err := st.CreateNodeBlocking("!00000007", &disabledNodeNum, "disabled", false); err != nil {
t.Fatalf("CreateNodeBlocking(disabled) error = %v", err)
}
if _, err := st.CreateIPBlocking("192.168.1.0/24", "lan", true); err != nil {
t.Fatalf("CreateIPBlocking(cidr) error = %v", err)
}
if _, err := st.CreateIPBlocking("10.0.0.1", "disabled", false); err != nil {
t.Fatalf("CreateIPBlocking(disabled) error = %v", err)
}
if _, err := st.CreateForbiddenWordBlocking("spam", "contains", false, "enabled", true); err != nil {
t.Fatalf("CreateForbiddenWordBlocking(enabled) error = %v", err)
}
if _, err := st.CreateForbiddenWordBlocking("blocked", "contains", false, "disabled", false); err != nil {
t.Fatalf("CreateForbiddenWordBlocking(disabled) error = %v", err)
}
cache, err := New(st)
if err != nil {
t.Fatalf("New() error = %v", err)
}
if !cache.IsNodeBlocked("!12345678", nil) {
t.Fatal("IsNodeBlocked(enabled node id) = false, want true")
}
if !cache.IsNodeBlocked("", uint32(nodeNum)) {
t.Fatal("IsNodeBlocked(enabled node num) = false, want true")
}
if cache.IsNodeBlocked("!00000007", disabledNodeNum) {
t.Fatal("IsNodeBlocked(disabled node) = true, want false")
}
if !cache.IsIPBlocked("192.168.1.42") {
t.Fatal("IsIPBlocked(CIDR member) = false, want true")
}
if cache.IsIPBlocked("10.0.0.1") {
t.Fatal("IsIPBlocked(disabled IP) = true, want false")
}
if word, ok := cache.FindForbiddenWord("This is SPAM text"); !ok || word != "spam" {
t.Fatalf("FindForbiddenWord(case-insensitive) = %q, %v, want spam, true", word, ok)
}
if _, ok := cache.FindForbiddenWord("disabled blocked text"); ok {
t.Fatal("FindForbiddenWord(disabled word) = true, want false")
}
}
func TestBlockingCacheIPExactAndCIDR(t *testing.T) {
st := openTestStore(t)
defer st.Close()
if _, err := st.CreateIPBlocking("127.0.0.1", "loopback", true); err != nil {
t.Fatalf("CreateIPBlocking(ip) error = %v", err)
}
if _, err := st.CreateIPBlocking("2001:db8::/32", "docs", true); err != nil {
t.Fatalf("CreateIPBlocking(ipv6 cidr) error = %v", err)
}
cache, err := New(st)
if err != nil {
t.Fatalf("New() error = %v", err)
}
if !cache.IsIPBlocked("127.0.0.1") {
t.Fatal("IsIPBlocked(exact IPv4) = false, want true")
}
if !cache.IsIPBlocked("2001:db8::1") {
t.Fatal("IsIPBlocked(IPv6 CIDR) = false, want true")
}
if cache.IsIPBlocked("localhost") {
t.Fatal("IsIPBlocked(hostname) = true, want false")
}
}
func TestBlockingCacheForbiddenWordCaseSensitivity(t *testing.T) {
st := openTestStore(t)
defer st.Close()
if _, err := st.CreateForbiddenWordBlocking("Spam", "contains", true, "case-sensitive", true); err != nil {
t.Fatalf("CreateForbiddenWordBlocking(case-sensitive) error = %v", err)
}
cache, err := New(st)
if err != nil {
t.Fatalf("New() error = %v", err)
}
if _, ok := cache.FindForbiddenWord("lowercase spam"); ok {
t.Fatal("FindForbiddenWord(lowercase) = true, want false")
}
if word, ok := cache.FindForbiddenWord("contains Spam"); !ok || word != "Spam" {
t.Fatalf("FindForbiddenWord(exact case) = %q, %v, want Spam, true", word, ok)
}
}
-13
View File
@@ -1,13 +0,0 @@
package blocking
import (
"testing"
"meshtastic_mqtt_server/internal/store"
"meshtastic_mqtt_server/internal/store/testutil"
)
// openTestStore 委托到 store/testutil,让本包的测试代码保持简洁。
func openTestStore(t *testing.T) *store.Store {
return testutil.OpenStore(t)
}
+2 -1
View File
@@ -416,8 +416,9 @@ func (s *Service) recordOutboundDirectMessage(bot *storepkg.BotNodeRecord, msg *
BotMessageID: botMessageID, BotMessageID: botMessageID,
CreatedBy: createdByPtr, CreatedBy: createdByPtr,
PublishedAt: msg.PublishedAt, PublishedAt: msg.PublishedAt,
// 出向消息从产生那一刻起就视为已读,未读计数只关心 inbound。 // 出向消息从产生那一刻起就视为"已读",未读计数只关心 inbound。
ReadAt: &now, ReadAt: &now,
CreatedAt: now,
} }
if err := s.store.InsertBotDirectMessage(dm); err != nil { if err := s.store.InsertBotDirectMessage(dm); err != nil {
printJSON(map[string]any{ printJSON(map[string]any{
-358
View File
@@ -1,358 +0,0 @@
package config
import (
"os"
"path/filepath"
"strings"
"testing"
)
func TestLoadConfigCreatesDefaultFile(t *testing.T) {
path := filepath.Join(t.TempDir(), "mesh_mqtt_go", FileName)
cfg, err := Load(path)
if err != nil {
t.Fatalf("Load() error = %v", err)
}
if cfg.MQTT.Host != "0.0.0.0" {
t.Fatalf("host = %q, want 0.0.0.0", cfg.MQTT.Host)
}
if cfg.MQTT.Port != 1883 {
t.Fatalf("port = %d, want 1883", cfg.MQTT.Port)
}
if cfg.MQTT.TLS.Enabled {
t.Fatalf("tls enabled = true, want false")
}
if cfg.Meshtastic.PSK != "AQ==" {
t.Fatalf("psk = %q, want AQ==", cfg.Meshtastic.PSK)
}
if cfg.Database.Driver != "sqlite" {
t.Fatalf("database driver = %q, want sqlite", cfg.Database.Driver)
}
if cfg.Database.SQLite.Path == "" {
t.Fatalf("sqlite path is empty")
}
if !cfg.Web.Enabled {
t.Fatalf("web enabled = false, want true")
}
if !cfg.Web.PortEnabled {
t.Fatalf("web port enabled = false, want true")
}
wantSocketEnabled := defaultWebSocketPath() != ""
if cfg.Web.SocketEnabled != wantSocketEnabled {
t.Fatalf("web socket enabled = %t, want %t", cfg.Web.SocketEnabled, wantSocketEnabled)
}
if cfg.Web.Port != 8080 {
t.Fatalf("web port = %d, want 8080", cfg.Web.Port)
}
if cfg.Web.SocketPath != defaultWebSocketPath() {
t.Fatalf("web socket path = %q, want %q", cfg.Web.SocketPath, defaultWebSocketPath())
}
if cfg.Web.StaticDir != "./dist" {
t.Fatalf("web static dir = %q, want ./dist", cfg.Web.StaticDir)
}
if cfg.Web.MapTileCacheDir != defaultMapTileCacheDir() {
t.Fatalf("web map tile cache dir = %q, want %q", cfg.Web.MapTileCacheDir, defaultMapTileCacheDir())
}
if _, err := os.Stat(path); err != nil {
t.Fatalf("default config was not written: %v", err)
}
}
func TestLoadConfigFillsMissingFields(t *testing.T) {
path := filepath.Join(t.TempDir(), "mesh_mqtt_go", FileName)
if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(path, []byte("mqtt:\n port: 1884\n"), 0644); err != nil {
t.Fatal(err)
}
cfg, err := Load(path)
if err != nil {
t.Fatalf("Load() error = %v", err)
}
if cfg.MQTT.Port != 1884 {
t.Fatalf("port = %d, want 1884", cfg.MQTT.Port)
}
if cfg.MQTT.Host != "0.0.0.0" {
t.Fatalf("host = %q, want 0.0.0.0", cfg.MQTT.Host)
}
if cfg.Meshtastic.PSK != "AQ==" {
t.Fatalf("psk = %q, want AQ==", cfg.Meshtastic.PSK)
}
if cfg.Database.Driver != "sqlite" {
t.Fatalf("database driver = %q, want sqlite", cfg.Database.Driver)
}
data, err := os.ReadFile(path)
if err != nil {
t.Fatal(err)
}
text := string(data)
for _, want := range []string{"host:", "tls:", "enabled:", "cert_file:", "key_file:", "meshtastic:", "psk:", "database:", "driver:", "sqlite:", "mysql:", "dsn:", "web:", "port_enabled:", "socket_enabled:", "port:", "socket_path:", "static_dir:", "map_tile_cache_dir:"} {
if !strings.Contains(text, want) {
t.Fatalf("completed config missing %q in:\n%s", want, text)
}
}
}
func TestLoadConfigPreservesExplicitFalse(t *testing.T) {
path := filepath.Join(t.TempDir(), "mesh_mqtt_go", FileName)
if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil {
t.Fatal(err)
}
content := "mqtt:\n host: 127.0.0.1\n port: 1885\n tls:\n enabled: false\n cert_file: cert.pem\n key_file: key.pem\nmeshtastic:\n psk: AQ==\ndatabase:\n driver: sqlite\n sqlite:\n path: test.db\n mysql:\n dsn: \"\"\n"
if err := os.WriteFile(path, []byte(content), 0644); err != nil {
t.Fatal(err)
}
cfg, err := Load(path)
if err != nil {
t.Fatalf("Load() error = %v", err)
}
if cfg.MQTT.TLS.Enabled {
t.Fatalf("tls enabled = true, want explicit false")
}
if cfg.MQTT.TLS.CertFile != "cert.pem" || cfg.MQTT.TLS.KeyFile != "key.pem" {
t.Fatalf("tls paths = %q/%q, want cert.pem/key.pem", cfg.MQTT.TLS.CertFile, cfg.MQTT.TLS.KeyFile)
}
}
func TestLoadConfigPreservesExplicitWebFalse(t *testing.T) {
path := filepath.Join(t.TempDir(), "mesh_mqtt_go", FileName)
if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil {
t.Fatal(err)
}
content := "web:\n enabled: false\n port_enabled: false\n socket_enabled: false\n host: 127.0.0.1\n port: 8081\n static_dir: ./public\n"
if err := os.WriteFile(path, []byte(content), 0644); err != nil {
t.Fatal(err)
}
cfg, err := Load(path)
if err != nil {
t.Fatalf("Load() error = %v", err)
}
if cfg.Web.Enabled {
t.Fatalf("web enabled = true, want explicit false")
}
if cfg.Web.PortEnabled || cfg.Web.SocketEnabled {
t.Fatalf("web listener enabled = %t/%t, want explicit false/false", cfg.Web.PortEnabled, cfg.Web.SocketEnabled)
}
if cfg.Web.Host != "127.0.0.1" || cfg.Web.Port != 8081 || cfg.Web.StaticDir != "./public" {
t.Fatalf("web config = %#v", cfg.Web)
}
}
func TestLoadConfigMalformedYAMLDoesNotOverwrite(t *testing.T) {
path := filepath.Join(t.TempDir(), "mesh_mqtt_go", FileName)
if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil {
t.Fatal(err)
}
content := "mqtt:\n port: [\n"
if err := os.WriteFile(path, []byte(content), 0644); err != nil {
t.Fatal(err)
}
_, err := Load(path)
if err == nil {
t.Fatalf("Load() error = nil, want parse error")
}
data, readErr := os.ReadFile(path)
if readErr != nil {
t.Fatal(readErr)
}
if string(data) != content {
t.Fatalf("malformed config was overwritten: %q", string(data))
}
}
func TestDefaultConfigDirForGOOS(t *testing.T) {
wantRelative := filepath.Join(".", "win", "etc", "mesh_mqtt_go")
for _, goos := range []string{"windows", "darwin"} {
path := defaultConfigDirForGOOS(goos)
if path != wantRelative {
t.Fatalf("%s config dir = %q, want %q", goos, path, wantRelative)
}
}
linuxPath := defaultConfigDirForGOOS("linux")
wantLinux := filepath.Join(string(filepath.Separator), "etc", "mesh_mqtt_go")
if linuxPath != wantLinux {
t.Fatalf("linux config dir = %q, want %q", linuxPath, wantLinux)
}
}
func TestDefaultMapTileCacheDirForGOOS(t *testing.T) {
wantRelative := filepath.Join(".", "win", "srv", "mesh_mqtt_go")
for _, goos := range []string{"windows", "darwin"} {
path := defaultMapTileCacheDirForGOOS(goos)
if path != wantRelative {
t.Fatalf("%s map tile cache dir = %q, want %q", goos, path, wantRelative)
}
}
linuxPath := defaultMapTileCacheDirForGOOS("linux")
wantLinux := filepath.Join(string(filepath.Separator), "srv", "mesh_mqtt_go")
if linuxPath != wantLinux {
t.Fatalf("linux map tile cache dir = %q, want %q", linuxPath, wantLinux)
}
}
func TestDefaultWebSocketPathForGOOS(t *testing.T) {
if windowsPath := defaultWebSocketPathForGOOS("windows"); windowsPath != "" {
t.Fatalf("windows web socket path = %q, want empty", windowsPath)
}
darwinPath := defaultWebSocketPathForGOOS("darwin")
wantDarwin := filepath.Join(".", "win", "opt", "mesh_mqtt_go", "web.sock")
if darwinPath != wantDarwin {
t.Fatalf("darwin web socket path = %q, want %q", darwinPath, wantDarwin)
}
linuxPath := defaultWebSocketPathForGOOS("linux")
want := filepath.Join(string(filepath.Separator), "opt", "mesh_mqtt_go", "web.sock")
if linuxPath != want {
t.Fatalf("linux web socket path = %q, want %q", linuxPath, want)
}
}
func TestClearWebSocketPathOnUnsupportedGOOS(t *testing.T) {
cfg := Default()
cfg.Web.SocketPath = filepath.Join(".", "win", "opt", "mesh_mqtt_go", "web.sock")
if !ClearWebSocketPathOnUnsupportedGOOS(cfg, "windows") {
t.Fatalf("ClearWebSocketPathOnUnsupportedGOOS() = false, want true")
}
if cfg.Web.SocketPath != "" {
t.Fatalf("windows web socket path = %q, want empty", cfg.Web.SocketPath)
}
if cfg.Web.SocketEnabled {
t.Fatalf("windows web socket enabled = true, want false")
}
cfg.Web.SocketPath = "/opt/mesh_mqtt_go/web.sock"
if ClearWebSocketPathOnUnsupportedGOOS(cfg, "linux") {
t.Fatalf("linux ClearWebSocketPathOnUnsupportedGOOS() = true, want false")
}
if cfg.Web.SocketPath == "" {
t.Fatalf("linux web socket path was cleared")
}
}
func TestDefaultSQLitePathForGOOS(t *testing.T) {
wantRelative := filepath.Join(".", "win", "etc", "mesh_mqtt_go", "mesh_mqtt_go.db")
for _, goos := range []string{"windows", "darwin"} {
path := defaultSQLitePathForGOOS(goos)
if path != wantRelative {
t.Fatalf("%s sqlite path = %q, want %q", goos, path, wantRelative)
}
}
linuxPath := defaultSQLitePathForGOOS("linux")
want := filepath.Join(string(filepath.Separator), "srv", "mesh_mqtt_go", "mesh_mqtt_go.db")
if linuxPath != want {
t.Fatalf("linux sqlite path = %q, want %q", linuxPath, want)
}
}
func TestValidateConfigDatabase(t *testing.T) {
cfg := Default()
cfg.Database.Driver = "postgres"
if err := Validate(cfg); err == nil || !strings.Contains(err.Error(), "database.driver") {
t.Fatalf("invalid driver error = %v, want database.driver error", err)
}
cfg = Default()
cfg.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)
}
cfg = Default()
cfg.Database.Driver = "mysql"
cfg.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)
}
}
func TestValidateConfigWeb(t *testing.T) {
cfg := Default()
cfg.Web.Port = 0
if err := Validate(cfg); err == nil || !strings.Contains(err.Error(), "web port") {
t.Fatalf("invalid web port error = %v, want web port error", err)
}
cfg = Default()
cfg.Web.PortEnabled = false
cfg.Web.Port = 0
if err := Validate(cfg); err != nil {
t.Fatalf("disabled web port with invalid port error = %v, want nil", err)
}
cfg = Default()
cfg.Web.SocketEnabled = false
cfg.Web.SocketPath = ""
if err := Validate(cfg); err != nil {
t.Fatalf("disabled web socket with empty path error = %v, want nil", err)
}
cfg = Default()
cfg.Web.PortEnabled = false
cfg.Web.SocketEnabled = true
cfg.Web.SocketPath = ""
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)
}
cfg = Default()
cfg.Web.PortEnabled = false
cfg.Web.SocketEnabled = false
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)
}
cfg = Default()
cfg.Web.StaticDir = ""
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)
}
cfg = Default()
cfg.Web.MapTileCacheDir = ""
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)
}
cfg = Default()
cfg.Web.Enabled = false
cfg.Web.PortEnabled = false
cfg.Web.SocketEnabled = false
cfg.Web.Port = 0
cfg.Web.StaticDir = ""
if err := Validate(cfg); err != nil {
t.Fatalf("disabled web validate error = %v, want nil", err)
}
}
func TestBuildTLSConfigDisabled(t *testing.T) {
cfg, err := BuildTLS(TLSConfig{})
if err != nil {
t.Fatalf("BuildTLS() error = %v", err)
}
if cfg != nil {
t.Fatalf("BuildTLS() = %#v, want nil", cfg)
}
}
func TestBuildTLSConfigRequiresCertAndKey(t *testing.T) {
_, err := BuildTLS(TLSConfig{Enabled: true})
if err == nil || !strings.Contains(err.Error(), "cert_file") {
t.Fatalf("missing cert error = %v, want cert_file error", err)
}
_, err = BuildTLS(TLSConfig{Enabled: true, CertFile: "cert.pem"})
if err == nil || !strings.Contains(err.Error(), "key_file") {
t.Fatalf("missing key error = %v, want key_file error", err)
}
}
-160
View File
@@ -1,160 +0,0 @@
package llm
import (
"testing"
)
func TestUpdateProvider(t *testing.T) {
// Create initial state with one provider
configs := []ProviderConfig{
{
Name: "test-provider",
Active: true,
APIKey: "test-key",
BaseURL: "https://test.example.com",
Model: "test-model",
Timeout: 120,
ContextWindowTokens: 4096,
},
}
state, err := NewState(configs)
if err != nil {
t.Fatalf("failed to create state: %v", err)
}
// Get the initial profile
profile := state.ActiveProfile()
if profile.Config.APIKey != "test-key" {
t.Errorf("expected APIKey 'test-key', got '%s'", profile.Config.APIKey)
}
// Update the provider with new config
updatedConfig := ProviderConfig{
Name: "test-provider",
Active: true,
APIKey: "new-key",
BaseURL: "https://new.example.com",
Model: "new-model",
Timeout: 60,
ContextWindowTokens: 8192,
}
err = state.UpdateProvider(updatedConfig)
if err != nil {
t.Fatalf("failed to update provider: %v", err)
}
// Verify the update
profile = state.ActiveProfile()
if profile.Config.APIKey != "new-key" {
t.Errorf("expected updated APIKey 'new-key', got '%s'", profile.Config.APIKey)
}
if profile.Config.BaseURL != "https://new.example.com" {
t.Errorf("expected updated BaseURL 'https://new.example.com', got '%s'", profile.Config.BaseURL)
}
if profile.Config.Model != "new-model" {
t.Errorf("expected updated Model 'new-model', got '%s'", profile.Config.Model)
}
if profile.Config.Timeout != 60 {
t.Errorf("expected updated Timeout 60, got %d", profile.Config.Timeout)
}
}
func TestAddProvider(t *testing.T) {
// Create initial state with one provider
configs := []ProviderConfig{
{
Name: "provider1",
Active: true,
APIKey: "key1",
BaseURL: "https://example1.com",
Model: "model1",
Timeout: 120,
ContextWindowTokens: 4096,
},
}
state, err := NewState(configs)
if err != nil {
t.Fatalf("failed to create state: %v", err)
}
// Add a second provider
newConfig := ProviderConfig{
Name: "provider2",
Active: false,
APIKey: "key2",
BaseURL: "https://example2.com",
Model: "model2",
Timeout: 60,
ContextWindowTokens: 8192,
}
err = state.AddProvider(newConfig)
if err != nil {
t.Fatalf("failed to add provider: %v", err)
}
// Verify the new provider exists
profile, err := state.GetProfile("provider2")
if err != nil {
t.Fatalf("failed to get provider2: %v", err)
}
if profile.Config.Name != "provider2" {
t.Errorf("expected name 'provider2', got '%s'", profile.Config.Name)
}
// Verify active provider is still provider1
activeProfile := state.ActiveProfile()
if activeProfile.Config.Name != "provider1" {
t.Errorf("expected active provider 'provider1', got '%s'", activeProfile.Config.Name)
}
}
func TestRemoveProvider(t *testing.T) {
// Create initial state with two providers
configs := []ProviderConfig{
{
Name: "provider1",
Active: true,
APIKey: "key1",
BaseURL: "https://example1.com",
Model: "model1",
Timeout: 120,
ContextWindowTokens: 4096,
},
{
Name: "provider2",
Active: false,
APIKey: "key2",
BaseURL: "https://example2.com",
Model: "model2",
Timeout: 60,
ContextWindowTokens: 8192,
},
}
state, err := NewState(configs)
if err != nil {
t.Fatalf("failed to create state: %v", err)
}
// Remove provider2
err = state.RemoveProvider("provider2")
if err != nil {
t.Fatalf("failed to remove provider: %v", err)
}
// Verify provider2 is gone
_, err = state.GetProfile("provider2")
if err == nil {
t.Error("expected error when getting removed provider, got nil")
}
// Try to remove the last provider (should fail)
err = state.RemoveProvider("provider1")
if err == nil {
t.Error("expected error when removing last provider, got nil")
}
}
+76 -9
View File
@@ -9,6 +9,7 @@ import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"gorm.io/gorm" "gorm.io/gorm"
aipkg "meshtastic_mqtt_server/internal/ai"
storepkg "meshtastic_mqtt_server/internal/store" storepkg "meshtastic_mqtt_server/internal/store"
"meshtastic_mqtt_server/internal/webutil" "meshtastic_mqtt_server/internal/webutil"
) )
@@ -18,6 +19,8 @@ type LLMProviderReloader interface {
ReloadLLMProvider(config interface{}) error ReloadLLMProvider(config interface{}) error
AddLLMProvider(config interface{}) error AddLLMProvider(config interface{}) error
RemoveLLMProvider(name string) error RemoveLLMProvider(name string) error
AIServiceStatus() aipkg.AIServiceStatus
RestartAIService() error
} }
func RegisterRoutes(r *gin.RouterGroup, store *storepkg.Store, aiService LLMProviderReloader) { func RegisterRoutes(r *gin.RouterGroup, store *storepkg.Store, aiService LLMProviderReloader) {
@@ -49,6 +52,10 @@ func RegisterRoutes(r *gin.RouterGroup, store *storepkg.Store, aiService LLMProv
// LLM Primary Config - 主 AI 回复配置 // LLM Primary Config - 主 AI 回复配置
group.GET("/primary-config", handleGetLLMPrimaryConfig(store)) group.GET("/primary-config", handleGetLLMPrimaryConfig(store))
group.PUT("/primary-config", handleUpdateLLMPrimaryConfig(store)) group.PUT("/primary-config", handleUpdateLLMPrimaryConfig(store))
// AI Service Status
group.GET("/status", handleGetAIServiceStatus(aiService))
group.POST("/restart", handleRestartAIService(aiService))
} }
} }
@@ -298,7 +305,14 @@ func handleCreateLLMProvider(store *storepkg.Store, aiService LLMProviderReloade
} }
// Reload AI service with new provider // Reload AI service with new provider
if aiService != nil { if aiService == nil {
c.JSON(http.StatusOK, gin.H{
"status": "ok",
"item": llmProviderDTO(*record),
"warning": "AI 服务未运行,配置已保存但需重启服务后生效",
})
return
}
providerConfig := map[string]interface{}{ providerConfig := map[string]interface{}{
"Name": record.Name, "Name": record.Name,
"Active": record.Active, "Active": record.Active,
@@ -309,7 +323,6 @@ func handleCreateLLMProvider(store *storepkg.Store, aiService LLMProviderReloade
"ContextWindowTokens": record.ContextWindowTokens, "ContextWindowTokens": record.ContextWindowTokens,
} }
if err := aiService.AddLLMProvider(providerConfig); err != nil { if err := aiService.AddLLMProvider(providerConfig); err != nil {
// Log warning but don't fail the request - database is already updated
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
"status": "ok", "status": "ok",
"item": llmProviderDTO(*record), "item": llmProviderDTO(*record),
@@ -317,7 +330,6 @@ func handleCreateLLMProvider(store *storepkg.Store, aiService LLMProviderReloade
}) })
return return
} }
}
c.JSON(http.StatusOK, gin.H{"status": "ok", "item": llmProviderDTO(*record)}) c.JSON(http.StatusOK, gin.H{"status": "ok", "item": llmProviderDTO(*record)})
} }
@@ -385,7 +397,14 @@ func handleUpdateLLMProvider(store *storepkg.Store, aiService LLMProviderReloade
} }
// Reload AI service with updated provider // Reload AI service with updated provider
if aiService != nil { if aiService == nil {
c.JSON(http.StatusOK, gin.H{
"status": "ok",
"item": llmProviderDTO(*record),
"warning": "AI 服务未运行,配置已保存但需重启服务后生效",
})
return
}
providerConfig := map[string]interface{}{ providerConfig := map[string]interface{}{
"Name": record.Name, "Name": record.Name,
"Active": record.Active, "Active": record.Active,
@@ -396,7 +415,6 @@ func handleUpdateLLMProvider(store *storepkg.Store, aiService LLMProviderReloade
"ContextWindowTokens": record.ContextWindowTokens, "ContextWindowTokens": record.ContextWindowTokens,
} }
if err := aiService.ReloadLLMProvider(providerConfig); err != nil { if err := aiService.ReloadLLMProvider(providerConfig); err != nil {
// Log warning but don't fail the request - database is already updated
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
"status": "ok", "status": "ok",
"item": llmProviderDTO(*record), "item": llmProviderDTO(*record),
@@ -404,7 +422,6 @@ func handleUpdateLLMProvider(store *storepkg.Store, aiService LLMProviderReloade
}) })
return return
} }
}
c.JSON(http.StatusOK, gin.H{"status": "ok", "item": llmProviderDTO(*record)}) c.JSON(http.StatusOK, gin.H{"status": "ok", "item": llmProviderDTO(*record)})
} }
@@ -424,16 +441,20 @@ func handleDeleteLLMProvider(store *storepkg.Store, aiService LLMProviderReloade
} }
// Remove provider from AI service // Remove provider from AI service
if aiService != nil { if aiService == nil {
c.JSON(http.StatusOK, gin.H{
"status": "ok",
"warning": "AI 服务未运行,配置已删除但需重启服务后生效",
})
return
}
if err := aiService.RemoveLLMProvider(name); err != nil { if err := aiService.RemoveLLMProvider(name); err != nil {
// Log warning but don't fail the request - database is already updated
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
"status": "ok", "status": "ok",
"warning": "provider deleted but failed to reload AI service: " + err.Error(), "warning": "provider deleted but failed to reload AI service: " + err.Error(),
}) })
return return
} }
}
c.JSON(http.StatusOK, gin.H{"status": "ok"}) c.JSON(http.StatusOK, gin.H{"status": "ok"})
} }
@@ -814,3 +835,49 @@ func llmPrimaryConfigDTO(row storepkg.LLMPrimaryConfigRecord) map[string]any {
"updated_at": row.UpdatedAt, "updated_at": row.UpdatedAt,
} }
} }
// ============================================
// AI Service Status & Restart Handlers
// ============================================
func handleGetAIServiceStatus(aiService LLMProviderReloader) gin.HandlerFunc {
return func(c *gin.Context) {
if aiService == nil {
c.JSON(http.StatusOK, gin.H{
"running": false,
"enabled": false,
"provider_count": 0,
"message": "AI 服务未初始化",
})
return
}
status := aiService.AIServiceStatus()
c.JSON(http.StatusOK, gin.H{
"running": status.Running,
"enabled": status.Enabled,
"provider_count": status.ProviderCount,
"message": status.Message,
})
}
}
func handleRestartAIService(aiService LLMProviderReloader) gin.HandlerFunc {
return func(c *gin.Context) {
if aiService == nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "AI 服务未初始化,无法重启"})
return
}
if err := aiService.RestartAIService(); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "重启 AI 服务失败: " + err.Error()})
return
}
status := aiService.AIServiceStatus()
c.JSON(http.StatusOK, gin.H{
"status": "ok",
"running": status.Running,
"enabled": status.Enabled,
"provider_count": status.ProviderCount,
"message": "AI 服务已重启",
})
}
}
+9
View File
@@ -20,6 +20,7 @@ type PacketBuildOptions struct {
PSK []byte PSK []byte
Encrypt bool Encrypt bool
ViaMQTT bool ViaMQTT bool
HopLimit uint32
} }
type TextMessageBuildOptions struct { type TextMessageBuildOptions struct {
@@ -252,6 +253,14 @@ func buildMeshPacket(opts PacketBuildOptions, data []byte) ([]byte, error) {
out = protowire.AppendTag(out, 14, protowire.VarintType) out = protowire.AppendTag(out, 14, protowire.VarintType)
out = protowire.AppendVarint(out, 1) out = protowire.AppendVarint(out, 1)
} }
hop := opts.HopLimit
if hop == 0 {
hop = 7
}
out = protowire.AppendTag(out, 9, protowire.VarintType)
out = protowire.AppendVarint(out, uint64(hop))
out = protowire.AppendTag(out, 15, protowire.VarintType)
out = protowire.AppendVarint(out, uint64(hop))
return out, nil return out, nil
} }
-215
View File
@@ -1,215 +0,0 @@
package mqtpp
import "testing"
func TestBuildTextMessageServiceEnvelopeRoundTrip(t *testing.T) {
key, err := ExpandPSK("AQ==")
if err != nil {
t.Fatalf("ExpandPSK() error = %v", err)
}
raw, err := BuildTextMessageServiceEnvelope(TextMessageBuildOptions{
PacketBuildOptions: PacketBuildOptions{
FromNodeNum: 0x12345678,
ToNodeNum: NodeNumBroadcast,
PacketID: 0x87654321,
ChannelID: "LongFast",
GatewayID: "!12345678",
PSK: key,
Encrypt: true,
ViaMQTT: true,
},
Text: "hello from bot",
})
if err != nil {
t.Fatalf("BuildTextMessageServiceEnvelope() error = %v", err)
}
valid, _, record := MQTTPP("msh/2/e/LongFast/!12345678", raw, key, Options{})
if !valid {
t.Fatalf("MQTTPP() valid = false, record = %#v", record)
}
if record["type"] != "text_message" {
t.Fatalf("record type = %v", record["type"])
}
if record["text"] != "hello from bot" {
t.Fatalf("text = %v", record["text"])
}
if record["from_num"] != uint32(0x12345678) {
t.Fatalf("from_num = %v", record["from_num"])
}
if record["packet_to_num"] != uint32(NodeNumBroadcast) {
t.Fatalf("packet_to_num = %v", record["packet_to_num"])
}
if record["decrypt_success"] != true {
t.Fatalf("decrypt_success = %v", record["decrypt_success"])
}
}
func TestBuildTextMessageServiceEnvelopeDirectRoundTrip(t *testing.T) {
key, err := ExpandPSK("AQ==")
if err != nil {
t.Fatalf("ExpandPSK() error = %v", err)
}
raw, err := BuildTextMessageServiceEnvelope(TextMessageBuildOptions{
PacketBuildOptions: PacketBuildOptions{
FromNodeNum: 0x12345678,
ToNodeNum: 0x10203040,
PacketID: 0x11111111,
ChannelID: "LongFast",
GatewayID: "!12345678",
PSK: key,
Encrypt: true,
ViaMQTT: true,
},
Text: "direct hello",
})
if err != nil {
t.Fatalf("BuildTextMessageServiceEnvelope() error = %v", err)
}
valid, _, record := MQTTPP("msh/2/e/LongFast/!12345678", raw, key, Options{})
if !valid {
t.Fatalf("MQTTPP() valid = false, record = %#v", record)
}
if record["text"] != "direct hello" {
t.Fatalf("text = %v", record["text"])
}
if record["packet_to"] != "!10203040" {
t.Fatalf("packet_to = %v", record["packet_to"])
}
if record["packet_to_num"] != uint32(0x10203040) {
t.Fatalf("packet_to_num = %v", record["packet_to_num"])
}
}
func TestBuildNodeInfoServiceEnvelopeRoundTrip(t *testing.T) {
key, err := ExpandPSK("AQ==")
if err != nil {
t.Fatalf("ExpandPSK() error = %v", err)
}
raw, err := BuildNodeInfoServiceEnvelope(NodeInfoBuildOptions{
PacketBuildOptions: PacketBuildOptions{
FromNodeNum: 0x12345678,
ToNodeNum: NodeNumBroadcast,
PacketID: 0x22222222,
ChannelID: "LongFast",
GatewayID: "!12345678",
PSK: key,
Encrypt: true,
ViaMQTT: true,
},
NodeID: "!12345678",
LongName: "MQTT Bot",
ShortName: "BT",
HWModel: 255,
Role: 0,
IsLicensed: false,
PublicKey: []byte{1, 2, 3},
})
if err != nil {
t.Fatalf("BuildNodeInfoServiceEnvelope() error = %v", err)
}
valid, _, record := MQTTPP("msh/2/e/LongFast/!12345678", raw, key, Options{})
if !valid {
t.Fatalf("MQTTPP() valid = false, record = %#v", record)
}
if record["type"] != "nodeinfo" {
t.Fatalf("record type = %v", record["type"])
}
if record["long_name"] != "MQTT Bot" {
t.Fatalf("long_name = %v", record["long_name"])
}
if record["short_name"] != "BT" {
t.Fatalf("short_name = %v", record["short_name"])
}
if record["hw_model"] != "PRIVATE_HW" {
t.Fatalf("hw_model = %v", record["hw_model"])
}
if record["role"] != "CLIENT" {
t.Fatalf("role = %v", record["role"])
}
if record["is_licensed"] != false {
t.Fatalf("is_licensed = %v", record["is_licensed"])
}
if record["public_key"] != "010203" {
t.Fatalf("public_key = %v", record["public_key"])
}
}
func TestBuildNodeInfoTruncatesNanopbStrings(t *testing.T) {
key, err := ExpandPSK("AQ==")
if err != nil {
t.Fatalf("ExpandPSK() error = %v", err)
}
raw, err := BuildNodeInfoServiceEnvelope(NodeInfoBuildOptions{
PacketBuildOptions: PacketBuildOptions{FromNodeNum: 0x12345678, ToNodeNum: NodeNumBroadcast, PacketID: 0x33333333, ChannelID: "LongFast", GatewayID: "!12345678", PSK: key, Encrypt: true, ViaMQTT: true},
NodeID: "!12345678",
LongName: "这是一个非常非常非常非常长的机器人节点名称",
ShortName: "机器人",
})
if err != nil {
t.Fatalf("BuildNodeInfoServiceEnvelope() error = %v", err)
}
valid, _, record := MQTTPP("msh/2/e/LongFast/!12345678", raw, key, Options{})
if !valid {
t.Fatalf("MQTTPP() valid = false, record = %#v", record)
}
if len([]byte(record["long_name"].(string))) > 40 {
t.Fatalf("long_name byte length = %d", len([]byte(record["long_name"].(string))))
}
if len([]byte(record["short_name"].(string))) > 5 {
t.Fatalf("short_name byte length = %d", len([]byte(record["short_name"].(string))))
}
}
func TestBuildAckServiceEnvelopeRoundTrip(t *testing.T) {
key, err := ExpandPSK("AQ==")
if err != nil {
t.Fatalf("ExpandPSK: %v", err)
}
const requestID uint32 = 0xabcd1234
raw, err := BuildAckServiceEnvelope(AckBuildOptions{
PacketBuildOptions: PacketBuildOptions{
FromNodeNum: 0x10101010,
ToNodeNum: 0x20202020,
PacketID: 0x30303030,
ChannelID: "LongFast",
GatewayID: "!10101010",
PSK: key,
Encrypt: true,
ViaMQTT: true,
},
RequestID: requestID,
})
if err != nil {
t.Fatalf("BuildAckServiceEnvelope: %v", err)
}
valid, _, record := MQTTPP("msh/2/e/LongFast/!10101010", raw, key, Options{})
if !valid {
t.Fatalf("MQTTPP not valid: %#v", record)
}
if record["portnum"] != "ROUTING_APP" {
t.Fatalf("portnum = %v", record["portnum"])
}
if record["type"] != "routing" {
t.Fatalf("type = %v", record["type"])
}
}
func TestParseNodeID(t *testing.T) {
num, err := ParseNodeID("!1234abcd")
if err != nil {
t.Fatalf("ParseNodeID() error = %v", err)
}
if num != 0x1234abcd {
t.Fatalf("num = %#x", num)
}
if NodeNumToID(num) != "!1234abcd" {
t.Fatalf("NodeNumToID() = %s", NodeNumToID(num))
}
}
-48
View File
@@ -1,48 +0,0 @@
package mqtpp
import (
"testing"
"google.golang.org/protobuf/encoding/protowire"
)
func TestMQTTPPEncryptedPacketDefaultRejected(t *testing.T) {
raw := encryptedServiceEnvelopeTestPayload()
valid, payload, record := MQTTPP("msh/test", raw, nil, Options{})
if valid {
t.Fatalf("valid = true, want false")
}
if payload != nil {
t.Fatalf("payload = %v, want nil", payload)
}
if record["type"] != "encrypted_packet" {
t.Fatalf("type = %v, want encrypted_packet", record["type"])
}
if record["error"] != "cannot be decrypted" {
t.Fatalf("error = %v, want cannot be decrypted", record["error"])
}
}
func TestMQTTPPEncryptedPacketAllowed(t *testing.T) {
raw := encryptedServiceEnvelopeTestPayload()
valid, payload, record := MQTTPP("msh/test", raw, nil, Options{AllowEncryptedForwarding: true})
if !valid {
t.Fatalf("valid = false, want true: %+v", record)
}
if string(payload) != string(raw) {
t.Fatalf("payload = %v, want raw payload", payload)
}
if record["type"] != "encrypted_packet" {
t.Fatalf("type = %v, want encrypted_packet", record["type"])
}
if record["error"] != nil {
t.Fatalf("error = %v, want nil", record["error"])
}
}
func encryptedServiceEnvelopeTestPayload() []byte {
packet := protowire.AppendTag(nil, 5, protowire.BytesType)
packet = protowire.AppendBytes(packet, []byte{1, 2, 3, 4})
envelope := protowire.AppendTag(nil, 1, protowire.BytesType)
return protowire.AppendBytes(envelope, packet)
}
-273
View File
@@ -1,273 +0,0 @@
package mqtpp
import (
"bytes"
"crypto/ecdh"
"crypto/rand"
"encoding/binary"
"testing"
"google.golang.org/protobuf/encoding/protowire"
)
func TestBuildPKITextMessageRoundTrip(t *testing.T) {
curve := ecdh.X25519()
senderPriv, err := curve.GenerateKey(rand.Reader)
if err != nil {
t.Fatalf("generate sender key: %v", err)
}
recipientPriv, err := curve.GenerateKey(rand.Reader)
if err != nil {
t.Fatalf("generate recipient key: %v", err)
}
const text = "hello over PKI 你好"
const fromNum uint32 = 0x12345678
const toNum uint32 = 0xa1b2c3d4
const packetID uint32 = 0xdeadbeef
raw, err := BuildPKITextMessageServiceEnvelope(PKITextMessageBuildOptions{
FromNodeNum: fromNum,
ToNodeNum: toNum,
PacketID: packetID,
GatewayID: NodeNumToID(fromNum),
ViaMQTT: true,
SenderPrivate: senderPriv.Bytes(),
RecipientPub: recipientPriv.PublicKey().Bytes(),
SenderPublic: senderPriv.PublicKey().Bytes(),
Text: text,
})
if err != nil {
t.Fatalf("BuildPKITextMessageServiceEnvelope: %v", err)
}
env, err := parseServiceEnvelope(raw)
if err != nil {
t.Fatalf("parseServiceEnvelope: %v", err)
}
if env.ChannelID != PKIChannelID {
t.Fatalf("channel_id = %q want %q", env.ChannelID, PKIChannelID)
}
if env.GatewayID != NodeNumToID(fromNum) {
t.Fatalf("gateway_id = %q", env.GatewayID)
}
pkt := env.Packet
if pkt.From != fromNum || pkt.To != toNum || pkt.ID != packetID {
t.Fatalf("packet header mismatch: %+v", pkt)
}
if !pkt.PKIEncrypted {
t.Fatalf("pki_encrypted = false")
}
if !pkt.ViaMQTT {
t.Fatalf("via_mqtt = false")
}
if pkt.Channel != 0 {
t.Fatalf("channel = %d want 0", pkt.Channel)
}
if pkt.PayloadVariant != "encrypted" || len(pkt.Encrypted) <= pkcOverhead {
t.Fatalf("encrypted payload missing: %+v", pkt)
}
// 收件人用对端私钥 + 发件人公钥推导共享密钥并解密
sharedKey, err := pkiSharedKey(recipientPriv.Bytes(), senderPriv.PublicKey().Bytes())
if err != nil {
t.Fatalf("pkiSharedKey: %v", err)
}
encryptedLen := len(pkt.Encrypted) - pkcOverhead
ciphertext := pkt.Encrypted[:encryptedLen]
auth := pkt.Encrypted[encryptedLen : encryptedLen+8]
extraNonce := binary.LittleEndian.Uint32(pkt.Encrypted[encryptedLen+8:])
plaintext, err := aesCCMDecrypt(sharedKey, pkiNonce(packetID, fromNum, extraNonce), ciphertext, auth)
if err != nil {
t.Fatalf("aesCCMDecrypt: %v", err)
}
data, err := parseDataPacket(plaintext)
if err != nil {
t.Fatalf("parseDataPacket: %v", err)
}
if data.Portnum != textMessageApp {
t.Fatalf("portnum = %d", data.Portnum)
}
if string(data.Payload) != text {
t.Fatalf("text = %q want %q", string(data.Payload), text)
}
// 同样用 MQTTPP 解析路径:PKI 包对外应被识别为 encrypted_packet(无法解密),
// 但用错的 PSK 不应误报“channel hash mismatch” 之外的奇怪错误。
dummyPSK, _ := ExpandPSK("AQ==")
_, _, record := MQTTPP("msh/2/e/PKI/!12345678", raw, dummyPSK, Options{AllowEncryptedForwarding: true})
if record["channel_id"] != PKIChannelID {
t.Fatalf("MQTTPP record channel_id = %v", record["channel_id"])
}
if record["pki_encrypted"] != true {
t.Fatalf("pki_encrypted record = %v", record["pki_encrypted"])
}
}
func TestPKINonceLayoutMatchesFirmware(t *testing.T) {
// 复刻 firmware initNonce(fromNode, packetId, extraNonce) 期望的字节布局:
// nonce[0..8) = packetId(uint64 LE)
// nonce[4..8) 被 extraNonce(uint32 LE) 覆盖(当 extraNonce != 0
// nonce[8..12) = fromNode(uint32 LE)
// nonce[12] = 0
got := pkiNonce(0xaabbccdd, 0x11223344, 0x55667788)
want := []byte{
0xdd, 0xcc, 0xbb, 0xaa, // packetId low 4 bytes,未被 extraNonce 覆盖前
0x88, 0x77, 0x66, 0x55, // extraNonce 覆盖 nonce[4..8)
0x44, 0x33, 0x22, 0x11, // fromNode
0x00,
}
if !bytes.Equal(got, want) {
t.Fatalf("pkiNonce = % x\nwant % x", got, want)
}
}
func TestBuildPKITextMessageRejectsBroadcast(t *testing.T) {
curve := ecdh.X25519()
priv, _ := curve.GenerateKey(rand.Reader)
pub, _ := curve.GenerateKey(rand.Reader)
if _, err := BuildPKITextMessageServiceEnvelope(PKITextMessageBuildOptions{
FromNodeNum: 0x1,
ToNodeNum: NodeNumBroadcast,
PacketID: 0x2,
SenderPrivate: priv.Bytes(),
RecipientPub: pub.PublicKey().Bytes(),
Text: "hi",
}); err == nil {
t.Fatalf("expected error for broadcast destination")
}
}
// 确认 MeshPacket 中确实带上 pki_encrypted (tag 17) 与 public_key (tag 16)
func TestBuildPKIMeshPacketTags(t *testing.T) {
encrypted := []byte{0x01, 0x02, 0x03}
pub := make([]byte, 32)
for i := range pub {
pub[i] = byte(i)
}
raw := buildPKIMeshPacket(0x11, 0x22, 0x33, true, encrypted, pub)
tags := map[protowire.Number]bool{}
if err := walkFields(raw, func(num protowire.Number, _ protowire.Type, _ any) error {
tags[num] = true
return nil
}); err != nil {
t.Fatalf("walkFields: %v", err)
}
for _, want := range []protowire.Number{1, 2, 5, 6, 14, 16, 17} {
if !tags[want] {
t.Fatalf("missing tag %d", want)
}
}
}
// 端到端:发送方构造 PKI 包,接收方通过 PKIKeyResolver 解密并还原文本消息记录。
func TestMQTTPPDecryptsPKIWithResolver(t *testing.T) {
curve := ecdh.X25519()
senderPriv, _ := curve.GenerateKey(rand.Reader)
recipientPriv, _ := curve.GenerateKey(rand.Reader)
const text = "hello PKI inbound"
const fromNum uint32 = 0xaaaa1111
const toNum uint32 = 0xbbbb2222
const packetID uint32 = 0x77777777
raw, err := BuildPKITextMessageServiceEnvelope(PKITextMessageBuildOptions{
FromNodeNum: fromNum,
ToNodeNum: toNum,
PacketID: packetID,
GatewayID: NodeNumToID(fromNum),
ViaMQTT: true,
SenderPrivate: senderPriv.Bytes(),
RecipientPub: recipientPriv.PublicKey().Bytes(),
SenderPublic: senderPriv.PublicKey().Bytes(),
Text: text,
})
if err != nil {
t.Fatalf("build: %v", err)
}
resolver := func(to, from uint32) ([]byte, []byte, bool) {
if to != toNum || from != fromNum {
return nil, nil, false
}
return recipientPriv.Bytes(), senderPriv.PublicKey().Bytes(), true
}
dummyPSK, _ := ExpandPSK("AQ==")
valid, _, record := MQTTPP("msh/2/e/PKI/!aaaa1111", raw, dummyPSK, Options{PKIKeyResolver: resolver})
if !valid {
t.Fatalf("MQTTPP not valid: %#v", record)
}
if record["type"] != "text_message" {
t.Fatalf("type = %v, want text_message", record["type"])
}
if record["text"] != text {
t.Fatalf("text = %v", record["text"])
}
if record["pki_encrypted"] != true {
t.Fatalf("pki_encrypted = %v", record["pki_encrypted"])
}
}
func TestBuildPKIAckRoundTrip(t *testing.T) {
curve := ecdh.X25519()
botPriv, _ := curve.GenerateKey(rand.Reader)
devicePriv, _ := curve.GenerateKey(rand.Reader)
const fromNum uint32 = 0x0000beef // bot
const toNum uint32 = 0xfeed0000 // 原 device
const ackPacketID uint32 = 0xaaaa5555
const requestID uint32 = 0xdeadbeef
raw, err := BuildPKIAckServiceEnvelope(PKIAckBuildOptions{
FromNodeNum: fromNum,
ToNodeNum: toNum,
PacketID: ackPacketID,
RequestID: requestID,
GatewayID: NodeNumToID(fromNum),
ViaMQTT: true,
SenderPrivate: botPriv.Bytes(),
RecipientPub: devicePriv.PublicKey().Bytes(),
SenderPublic: botPriv.PublicKey().Bytes(),
})
if err != nil {
t.Fatalf("BuildPKIAckServiceEnvelope: %v", err)
}
// 设备侧解密
env, err := parseServiceEnvelope(raw)
if err != nil {
t.Fatalf("parse: %v", err)
}
if env.ChannelID != PKIChannelID {
t.Fatalf("channel_id = %q", env.ChannelID)
}
pkt := env.Packet
if !pkt.PKIEncrypted || pkt.From != fromNum || pkt.To != toNum || pkt.ID != ackPacketID {
t.Fatalf("ack header mismatch: %+v", pkt)
}
encryptedLen := len(pkt.Encrypted) - pkcOverhead
cipher := pkt.Encrypted[:encryptedLen]
auth := pkt.Encrypted[encryptedLen : encryptedLen+8]
extraNonce := binary.LittleEndian.Uint32(pkt.Encrypted[encryptedLen+8:])
sharedKey, err := pkiSharedKey(devicePriv.Bytes(), botPriv.PublicKey().Bytes())
if err != nil {
t.Fatalf("shared: %v", err)
}
plain, err := aesCCMDecrypt(sharedKey, pkiNonce(ackPacketID, fromNum, extraNonce), cipher, auth)
if err != nil {
t.Fatalf("decrypt: %v", err)
}
data, err := parseDataPacket(plain)
if err != nil {
t.Fatalf("data: %v", err)
}
if data.Portnum != routingApp {
t.Fatalf("portnum = %d, want ROUTING_APP(%d)", data.Portnum, routingApp)
}
// Routing payload 解析: 期望 oneof error_reason=NONE(0),即 wire 字节 0x18 0x00
wantRouting := []byte{0x18, 0x00}
if !bytes.Equal(data.Payload, wantRouting) {
t.Fatalf("routing payload = % x, want % x", data.Payload, wantRouting)
}
}
+78
View File
@@ -0,0 +1,78 @@
package mqttforward
import (
"crypto/sha256"
"encoding/hex"
"sync"
"time"
)
const dedupTTL = 15 * time.Second
type DedupQueue struct {
mu sync.Mutex
entries map[string]time.Time
stopCh chan struct{}
}
func NewDedupQueue() *DedupQueue {
return &DedupQueue{
entries: make(map[string]time.Time),
stopCh: make(chan struct{}),
}
}
func (dq *DedupQueue) TryForward(topic string, payload []byte) bool {
hash := dedupHash(topic, payload)
now := time.Now()
dq.mu.Lock()
defer dq.mu.Unlock()
if expiry, ok := dq.entries[hash]; ok && now.Before(expiry) {
return false
}
dq.entries[hash] = now.Add(dedupTTL)
return true
}
func (dq *DedupQueue) Start() {
go func() {
ticker := time.NewTicker(dedupTTL)
defer ticker.Stop()
for {
select {
case <-ticker.C:
dq.cleanup()
case <-dq.stopCh:
return
}
}
}()
}
func (dq *DedupQueue) Stop() {
close(dq.stopCh)
}
func (dq *DedupQueue) Len() int {
dq.mu.Lock()
defer dq.mu.Unlock()
return len(dq.entries)
}
func (dq *DedupQueue) cleanup() {
now := time.Now()
dq.mu.Lock()
defer dq.mu.Unlock()
for hash, expiry := range dq.entries {
if now.After(expiry) {
delete(dq.entries, hash)
}
}
}
func dedupHash(topic string, payload []byte) string {
h := sha256.New()
h.Write([]byte(topic))
h.Write(payload)
return hex.EncodeToString(h.Sum(nil))
}
@@ -1,45 +0,0 @@
package runtimesettings
import (
"testing"
storepkg "meshtastic_mqtt_server/internal/store"
"meshtastic_mqtt_server/internal/store/testutil"
)
func openTestStore(t *testing.T) *storepkg.Store {
return testutil.OpenStore(t)
}
func TestRuntimeSettingsCacheReload(t *testing.T) {
st := openTestStore(t)
defer st.Close()
cache, err := New(st)
if err != nil {
t.Fatalf("New() error = %v", err)
}
if cache.AllowEncryptedForwarding() {
t.Fatalf("AllowEncryptedForwarding() = true, want false")
}
if _, err := st.SetBoolRuntimeSetting(storepkg.RuntimeSettingAllowEncryptedForwarding, true, "test setting"); err != nil {
t.Fatalf("SetBoolRuntimeSetting(true) error = %v", err)
}
if err := cache.Reload(st); err != nil {
t.Fatalf("Reload() after true error = %v", err)
}
if !cache.AllowEncryptedForwarding() {
t.Fatalf("AllowEncryptedForwarding() = false, want true")
}
if _, err := st.SetBoolRuntimeSetting(storepkg.RuntimeSettingAllowEncryptedForwarding, false, "test setting"); err != nil {
t.Fatalf("SetBoolRuntimeSetting(false) error = %v", err)
}
if err := cache.Reload(st); err != nil {
t.Fatalf("Reload() after false error = %v", err)
}
if cache.AllowEncryptedForwarding() {
t.Fatalf("AllowEncryptedForwarding() = true, want false")
}
}
-207
View File
@@ -1,207 +0,0 @@
package store
import (
"errors"
"testing"
"gorm.io/gorm"
)
func TestNodeBlockingCRUD(t *testing.T) {
st := openTestStore(t)
defer st.Close()
nodeNum := int64(305419896)
rule, err := st.CreateNodeBlocking(" !12345678 ", &nodeNum, " noisy node ", true)
if err != nil {
t.Fatalf("CreateNodeBlocking() error = %v", err)
}
if rule.NodeID != "!12345678" || rule.NodeNum == nil || *rule.NodeNum != nodeNum || rule.Reason != "noisy node" || !rule.Enabled {
t.Fatalf("created node rule = %+v, want normalized fields", rule)
}
if _, err := st.CreateNodeBlocking("!12345678", nil, "duplicate", true); !errors.Is(err, ErrBlockingAlreadyExists) {
t.Fatalf("duplicate CreateNodeBlocking() error = %v, want ErrBlockingAlreadyExists", err)
}
updatedNum := int64(7)
updated, err := st.UpdateNodeBlocking(rule.ID, "!00000007", &updatedNum, "updated", false)
if err != nil {
t.Fatalf("UpdateNodeBlocking() error = %v", err)
}
if updated.NodeID != "!00000007" || updated.NodeNum == nil || *updated.NodeNum != updatedNum || updated.Reason != "updated" || updated.Enabled {
t.Fatalf("updated node rule = %+v, want updated fields", updated)
}
rows, err := st.ListNodeBlocking(ListOptions{})
if err != nil {
t.Fatalf("ListNodeBlocking() error = %v", err)
}
if len(rows) != 1 || rows[0].ID != rule.ID {
t.Fatalf("ListNodeBlocking() = %+v, want one updated rule", rows)
}
total, err := st.CountNodeBlocking(ListOptions{})
if err != nil || total != 1 {
t.Fatalf("CountNodeBlocking() = %d, %v, want 1, nil", total, err)
}
if err := st.DeleteNodeBlocking(rule.ID); err != nil {
t.Fatalf("DeleteNodeBlocking() error = %v", err)
}
if err := st.DeleteNodeBlocking(rule.ID); !errors.Is(err, gorm.ErrRecordNotFound) {
t.Fatalf("DeleteNodeBlocking(missing) error = %v, want record not found", err)
}
}
func TestNodeBlockingValidation(t *testing.T) {
st := openTestStore(t)
defer st.Close()
if _, err := st.CreateNodeBlocking(" ", nil, "", true); err == nil {
t.Fatal("CreateNodeBlocking(empty) error = nil, want error")
}
if _, err := st.UpdateNodeBlocking(1, "!missing", nil, "", true); !errors.Is(err, gorm.ErrRecordNotFound) {
t.Fatalf("UpdateNodeBlocking(missing) error = %v, want record not found", err)
}
}
func TestIPBlockingCRUDAndValidation(t *testing.T) {
st := openTestStore(t)
defer st.Close()
rule, err := st.CreateIPBlocking(" 127.0.0.1 ", "local", true)
if err != nil {
t.Fatalf("CreateIPBlocking(ip) error = %v", err)
}
if rule.IPValue != "127.0.0.1" || rule.Reason != "local" || !rule.Enabled {
t.Fatalf("created ip rule = %+v, want normalized IP", rule)
}
cidr, err := st.CreateIPBlocking("192.168.1.99/24", "cidr", true)
if err != nil {
t.Fatalf("CreateIPBlocking(cidr) error = %v", err)
}
if cidr.IPValue != "192.168.1.0/24" {
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) {
t.Fatalf("duplicate CreateIPBlocking() error = %v, want ErrBlockingAlreadyExists", err)
}
if _, err := st.CreateIPBlocking("not-an-ip", "invalid", true); err == nil {
t.Fatal("CreateIPBlocking(invalid) error = nil, want error")
}
updated, err := st.UpdateIPBlocking(rule.ID, "10.0.0.0/8", "updated", false)
if err != nil {
t.Fatalf("UpdateIPBlocking() error = %v", err)
}
if updated.IPValue != "10.0.0.0/8" || updated.Reason != "updated" || updated.Enabled {
t.Fatalf("updated ip rule = %+v, want updated fields", updated)
}
rows, err := st.ListIPBlocking(ListOptions{})
if err != nil {
t.Fatalf("ListIPBlocking() error = %v", err)
}
if len(rows) != 2 {
t.Fatalf("ListIPBlocking() length = %d, want 2", len(rows))
}
total, err := st.CountIPBlocking(ListOptions{})
if err != nil || total != 2 {
t.Fatalf("CountIPBlocking() = %d, %v, want 2, nil", total, err)
}
if err := st.DeleteIPBlocking(rule.ID); err != nil {
t.Fatalf("DeleteIPBlocking() error = %v", err)
}
if err := st.DeleteIPBlocking(rule.ID); !errors.Is(err, gorm.ErrRecordNotFound) {
t.Fatalf("DeleteIPBlocking(missing) error = %v, want record not found", err)
}
}
func TestListEnabledBlockingRules(t *testing.T) {
st := openTestStore(t)
defer st.Close()
nodeNum := int64(1)
if _, err := st.CreateNodeBlocking("!00000001", &nodeNum, "enabled", true); err != nil {
t.Fatalf("CreateNodeBlocking(enabled) error = %v", err)
}
if _, err := st.CreateNodeBlocking("!00000002", nil, "disabled", false); err != nil {
t.Fatalf("CreateNodeBlocking(disabled) error = %v", err)
}
if rows, err := st.ListEnabledNodeBlocking(); err != nil || len(rows) != 1 || rows[0].NodeID != "!00000001" {
t.Fatalf("ListEnabledNodeBlocking() = %+v, %v, want only enabled node", rows, err)
}
if _, err := st.CreateIPBlocking("127.0.0.1", "enabled", true); err != nil {
t.Fatalf("CreateIPBlocking(enabled) error = %v", err)
}
if _, err := st.CreateIPBlocking("192.168.1.1", "disabled", false); err != nil {
t.Fatalf("CreateIPBlocking(disabled) error = %v", err)
}
if rows, err := st.ListEnabledIPBlocking(); err != nil || len(rows) != 1 || rows[0].IPValue != "127.0.0.1" {
t.Fatalf("ListEnabledIPBlocking() = %+v, %v, want only enabled IP", rows, err)
}
if _, err := st.CreateForbiddenWordBlocking("spam", "contains", false, "enabled", true); err != nil {
t.Fatalf("CreateForbiddenWordBlocking(enabled) error = %v", err)
}
if _, err := st.CreateForbiddenWordBlocking("eggs", "contains", false, "disabled", false); err != nil {
t.Fatalf("CreateForbiddenWordBlocking(disabled) error = %v", err)
}
if rows, err := st.ListEnabledForbiddenWordBlocking(); err != nil || len(rows) != 1 || rows[0].Word != "spam" {
t.Fatalf("ListEnabledForbiddenWordBlocking() = %+v, %v, want only enabled word", rows, err)
}
}
func TestForbiddenWordBlockingCRUDAndValidation(t *testing.T) {
st := openTestStore(t)
defer st.Close()
rule, err := st.CreateForbiddenWordBlocking(" spam ", "", false, "junk", true)
if err != nil {
t.Fatalf("CreateForbiddenWordBlocking() error = %v", err)
}
if rule.Word != "spam" || rule.MatchType != ForbiddenWordMatchContains || rule.CaseSensitive || rule.Reason != "junk" || !rule.Enabled {
t.Fatalf("created word rule = %+v, want normalized fields", rule)
}
if _, err := st.CreateForbiddenWordBlocking("spam", "contains", false, "duplicate", true); !errors.Is(err, ErrBlockingAlreadyExists) {
t.Fatalf("duplicate CreateForbiddenWordBlocking() error = %v, want ErrBlockingAlreadyExists", err)
}
if _, err := st.CreateForbiddenWordBlocking(" ", "contains", false, "empty", true); err == nil {
t.Fatal("CreateForbiddenWordBlocking(empty) error = nil, want error")
}
if _, err := st.CreateForbiddenWordBlocking("regex", "regex", false, "unsupported", true); err == nil {
t.Fatal("CreateForbiddenWordBlocking(unsupported match type) error = nil, want error")
}
updated, err := st.UpdateForbiddenWordBlocking(rule.ID, "Spam", "contains", true, "updated", false)
if err != nil {
t.Fatalf("UpdateForbiddenWordBlocking() error = %v", err)
}
if updated.Word != "Spam" || updated.MatchType != "contains" || !updated.CaseSensitive || updated.Reason != "updated" || updated.Enabled {
t.Fatalf("updated word rule = %+v, want updated fields", updated)
}
rows, err := st.ListForbiddenWordBlocking(ListOptions{})
if err != nil {
t.Fatalf("ListForbiddenWordBlocking() error = %v", err)
}
if len(rows) != 1 || rows[0].ID != rule.ID {
t.Fatalf("ListForbiddenWordBlocking() = %+v, want one updated rule", rows)
}
total, err := st.CountForbiddenWordBlocking(ListOptions{})
if err != nil || total != 1 {
t.Fatalf("CountForbiddenWordBlocking() = %d, %v, want 1, nil", total, err)
}
if err := st.DeleteForbiddenWordBlocking(rule.ID); err != nil {
t.Fatalf("DeleteForbiddenWordBlocking() error = %v", err)
}
if err := st.DeleteForbiddenWordBlocking(rule.ID); !errors.Is(err, gorm.ErrRecordNotFound) {
t.Fatalf("DeleteForbiddenWordBlocking(missing) error = %v, want record not found", err)
}
}
+1 -5
View File
@@ -69,7 +69,6 @@ func (s *Store) ListBotDirectMessagesByConversation(opts BotDirectMessageListOpt
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("id DESC"). Order("id DESC").
Limit(opts.Limit). Limit(opts.Limit).
Offset(opts.Offset) Offset(opts.Offset)
@@ -159,7 +158,6 @@ func (s *Store) ListBotDirectConversations(botID uint64, opts ListOptions) ([]Bo
q := s.db.Table("(?) AS agg", subLast). q := s.db.Table("(?) AS agg", subLast).
Select("agg.bot_id AS bot_id, agg.peer_node_id AS peer_node_id, agg.peer_node_num AS peer_node_num, m.created_at AS last_message_at, m.text AS last_text, m.direction AS last_direction, agg.unread_count AS unread_count, agg.total_count AS total_count"). Select("agg.bot_id AS bot_id, agg.peer_node_id AS peer_node_id, agg.peer_node_num AS peer_node_num, m.created_at AS last_message_at, m.text AS last_text, m.direction AS last_direction, agg.unread_count AS unread_count, agg.total_count AS total_count").
Joins("JOIN bot_direct_messages m ON m.id = agg.last_id"). Joins("JOIN bot_direct_messages m ON m.id = agg.last_id").
Order("m.created_at DESC").
Order("m.id DESC"). Order("m.id DESC").
Limit(opts.Limit). Limit(opts.Limit).
Offset(opts.Offset) Offset(opts.Offset)
@@ -284,14 +282,13 @@ func insertInboundBotDirectMessage(s *Store, record map[string]any, clientInfo M
Status: BotMessageStatusPublished, Status: BotMessageStatusPublished,
ReceivedAt: &now, ReceivedAt: &now,
ContentJSON: contentPtr, ContentJSON: contentPtr,
CreatedAt: now,
} }
if err := s.InsertBotDirectMessage(dm); err != nil { if err := s.InsertBotDirectMessage(dm); err != nil {
return fmt.Errorf("insert bot direct message from %s: %w", peerNodeID, err) return fmt.Errorf("insert bot direct message from %s: %w", peerNodeID, err)
} }
_ = clientInfo // mqtt 元数据已经记录在 content_json 里,这里保留参数以保持队列签名一致 _ = clientInfo // mqtt 元数据已经记录在 content_json 里,这里保留参数以保持队列签名一致
// 同时将消息添加到 LLM 队列(忽略机器人自己发送的消息)
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"])
@@ -310,7 +307,6 @@ func insertInboundBotDirectMessage(s *Store, record map[string]any, clientInfo M
MessageType: "direct", MessageType: "direct",
ContentJSON: contentPtr, ContentJSON: contentPtr,
}) })
}
if err != nil { if err != nil {
printJSON(map[string]any{ printJSON(map[string]any{
"event": "llm_queue_enqueue_failed", "event": "llm_queue_enqueue_failed",
+6 -1
View File
@@ -66,6 +66,12 @@ func (s *Store) CountBotNodes(opts ListOptions) (int64, error) {
return total, s.db.Model(&BotNodeRecord{}).Count(&total).Error return total, s.db.Model(&BotNodeRecord{}).Count(&total).Error
} }
func (s *Store) IsBotNodeID(nodeID string) bool {
var count int64
s.db.Model(&BotNodeRecord{}).Where("node_id = ?", nodeID).Count(&count)
return count > 0
}
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 {
@@ -205,7 +211,6 @@ func (s *Store) ListBotMessages(opts BotMessageListOptions) ([]BotMessageRecord,
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("id DESC"). Order("id DESC").
Limit(opts.Limit). Limit(opts.Limit).
Offset(opts.Offset) Offset(opts.Offset)
+89 -3
View File
@@ -439,7 +439,7 @@ type TextMessageRecord struct {
Text *string `gorm:"column:text"` Text *string `gorm:"column:text"`
PayloadHex *string `gorm:"column:payload_hex"` PayloadHex *string `gorm:"column:payload_hex"`
Topic string `gorm:"column:topic;not null"` Topic string `gorm:"column:topic;not null"`
ChannelID *string `gorm:"column:channel_id"` ChannelID *string `gorm:"column:channel_id;type:varchar(255);index:idx_text_message_channel_id_created_at,priority:1"`
GatewayID *string `gorm:"column:gateway_id"` GatewayID *string `gorm:"column:gateway_id"`
PacketID *int64 `gorm:"column:packet_id;index:idx_text_message_packet_id"` PacketID *int64 `gorm:"column:packet_id;index:idx_text_message_packet_id"`
PacketTo *string `gorm:"column:packet_to"` PacketTo *string `gorm:"column:packet_to"`
@@ -458,13 +458,32 @@ type TextMessageRecord struct {
MQTTRemoteHost *string `gorm:"column:mqtt_remote_host"` MQTTRemoteHost *string `gorm:"column:mqtt_remote_host"`
MQTTRemotePort *string `gorm:"column:mqtt_remote_port"` MQTTRemotePort *string `gorm:"column:mqtt_remote_port"`
ContentJSON string `gorm:"column:content_json;not null"` ContentJSON string `gorm:"column:content_json;not null"`
CreatedAt time.Time `gorm:"column:created_at;autoCreateTime;index:idx_text_message_from_num_created_at,priority:2;index:idx_text_message_created_at"` CreatedAt time.Time `gorm:"column:created_at;autoCreateTime;index:idx_text_message_from_num_created_at,priority:2;index:idx_text_message_created_at;index:idx_text_message_channel_id_created_at,priority:2"`
} }
func (TextMessageRecord) TableName() string { func (TextMessageRecord) TableName() string {
return "text_message" return "text_message"
} }
type SchemaMigration struct {
Version int `gorm:"column:version;primaryKey"`
AppliedAt time.Time `gorm:"column:applied_at;autoCreateTime"`
}
func (SchemaMigration) TableName() string {
return "schema_migrations"
}
type ChannelRecord struct {
ID uint `gorm:"column:id;primaryKey;autoIncrement"`
ChannelID string `gorm:"column:channel_id;type:varchar(255);not null;uniqueIndex"`
CreatedAt time.Time `gorm:"column:created_at;autoCreateTime"`
}
func (ChannelRecord) TableName() string {
return "channels"
}
// LLMProviderRecord 保存 LLM API 配置,支持多个 AI 提供商 // LLMProviderRecord 保存 LLM API 配置,支持多个 AI 提供商
type LLMProviderRecord struct { type LLMProviderRecord struct {
Name string `gorm:"column:name;primaryKey;size:64;not null"` // 配置名称,如 "default"、"openai"、"ark" 等 Name string `gorm:"column:name;primaryKey;size:64;not null"` // 配置名称,如 "default"、"openai"、"ark" 等
@@ -691,6 +710,7 @@ func (s *Store) migrate() error {
{label: "nodeinfo", model: &NodeInfoRecord{}}, {label: "nodeinfo", model: &NodeInfoRecord{}},
{label: "map_report", model: &MapReportRecord{}}, {label: "map_report", model: &MapReportRecord{}},
{label: "text_message", model: &TextMessageRecord{}}, {label: "text_message", model: &TextMessageRecord{}},
{label: "channels", model: &ChannelRecord{}},
{label: "position", model: &PositionRecord{}}, {label: "position", model: &PositionRecord{}},
{label: "telemetry", model: &TelemetryRecord{}}, {label: "telemetry", model: &TelemetryRecord{}},
{label: "routing", model: &RoutingRecord{}}, {label: "routing", model: &RoutingRecord{}},
@@ -702,12 +722,15 @@ func (s *Store) migrate() error {
} }
} }
} }
if err := runDBMigrations(tx, s.driver); err != nil {
return err
}
for _, item := range []struct { for _, item := range []struct {
label string label string
model any model any
indexes []string indexes []string
}{ }{
{label: "text_message", model: &TextMessageRecord{}, indexes: []string{"idx_text_message_from_num_created_at", "idx_text_message_created_at", "idx_text_message_packet_id"}}, {label: "text_message", model: &TextMessageRecord{}, indexes: []string{"idx_text_message_from_num_created_at", "idx_text_message_created_at", "idx_text_message_packet_id", "idx_text_message_channel_id_created_at"}},
{label: "bot_direct_messages", model: &BotDirectMessageRecord{}, indexes: []string{"idx_bot_dm_bot_peer", "idx_bot_dm_bot_created_at"}}, {label: "bot_direct_messages", model: &BotDirectMessageRecord{}, indexes: []string{"idx_bot_dm_bot_peer", "idx_bot_dm_bot_created_at"}},
{label: "llm_message_queue", model: &LLMMessageQueueRecord{}, indexes: []string{"idx_llm_queue_bot_created"}}, {label: "llm_message_queue", model: &LLMMessageQueueRecord{}, indexes: []string{"idx_llm_queue_bot_created"}},
} { } {
@@ -878,6 +901,64 @@ func createMissingIndexes(migrator gorm.Migrator, model any, label string, index
return nil return nil
} }
type DBMigration struct {
Version int
Up func(tx *gorm.DB, driver string) error
}
var dbMigrations = []DBMigration{
{Version: 1, Up: migrateDB1},
{Version: 2, Up: migrateDB2},
}
func migrateDB1(tx *gorm.DB, driver string) error {
if driver == config.DriverMySQL {
if err := tx.Exec("ALTER TABLE text_message MODIFY COLUMN channel_id VARCHAR(255)").Error; err != nil {
return fmt.Errorf("alter text_message.channel_id to varchar: %w", err)
}
}
return nil
}
func migrateDB2(tx *gorm.DB, driver string) error {
var sql string
if driver == config.DriverSQLite {
sql = "INSERT OR IGNORE INTO channels (channel_id, created_at) SELECT DISTINCT channel_id, datetime('now') FROM text_message WHERE channel_id IS NOT NULL AND channel_id != ''"
} else {
sql = "INSERT IGNORE INTO channels (channel_id, created_at) SELECT DISTINCT channel_id, NOW() FROM text_message WHERE channel_id IS NOT NULL AND channel_id != ''"
}
if err := tx.Exec(sql).Error; err != nil {
return fmt.Errorf("backfill channels: %w", err)
}
return nil
}
func runDBMigrations(tx *gorm.DB, driver string) error {
if !tx.Migrator().HasTable(&SchemaMigration{}) {
if err := tx.Migrator().CreateTable(&SchemaMigration{}); err != nil {
return fmt.Errorf("create schema_migrations table: %w", err)
}
}
applied := make(map[int]bool)
var rows []SchemaMigration
tx.Find(&rows)
for _, r := range rows {
applied[r.Version] = true
}
for _, m := range dbMigrations {
if applied[m.Version] {
continue
}
if err := m.Up(tx, driver); err != nil {
return fmt.Errorf("db migration %d: %w", m.Version, err)
}
if err := tx.Create(&SchemaMigration{Version: m.Version}).Error; err != nil {
return fmt.Errorf("record db migration %d: %w", m.Version, err)
}
}
return nil
}
func (s *Store) UpsertNodeInfo(record map[string]any) error { func (s *Store) UpsertNodeInfo(record map[string]any) error {
node, err := nodeInfoFromRecord(record) node, err := nodeInfoFromRecord(record)
if err != nil { if err != nil {
@@ -1024,6 +1105,9 @@ func (s *Store) InsertTextMessage(record map[string]any, clientInfo MQTTClientIn
if err := s.db.Create(message).Error; err != nil { if err := s.db.Create(message).Error; err != nil {
return fmt.Errorf("insert text_message from %s: %w", message.FromID, err) return fmt.Errorf("insert text_message from %s: %w", message.FromID, err)
} }
if message.ChannelID != nil && *message.ChannelID != "" {
s.db.FirstOrCreate(&ChannelRecord{}, ChannelRecord{ChannelID: *message.ChannelID})
}
return nil return nil
} }
@@ -1207,6 +1291,7 @@ func textMessageFromRecord(record map[string]any, clientInfo MQTTClientInfo) (*T
MQTTRemoteHost: clientFields.MQTTRemoteHost, MQTTRemoteHost: clientFields.MQTTRemoteHost,
MQTTRemotePort: clientFields.MQTTRemotePort, MQTTRemotePort: clientFields.MQTTRemotePort,
ContentJSON: common.ContentJSON, ContentJSON: common.ContentJSON,
CreatedAt: time.Now(),
}, nil }, nil
} }
@@ -1317,6 +1402,7 @@ func AppendPacketFieldsFromRecord(record map[string]any, wantType string, client
DecryptSuccess: nullableBool(record["decrypt_success"]), DecryptSuccess: nullableBool(record["decrypt_success"]),
DecryptStatus: NullableString(record["decrypt_status"]), DecryptStatus: NullableString(record["decrypt_status"]),
ContentJSON: string(contentJSON), ContentJSON: string(contentJSON),
CreatedAt: time.Now(),
}, MQTTClientRecordFields{ }, MQTTClientRecordFields{
MQTTClientID: NullableString(clientInfo.ClientID), MQTTClientID: NullableString(clientInfo.ClientID),
MQTTUsername: NullableString(clientInfo.Username), MQTTUsername: NullableString(clientInfo.Username),
File diff suppressed because it is too large Load Diff
-104
View File
@@ -1,104 +0,0 @@
package store
import (
"database/sql"
"testing"
)
func TestDBWriteQueueWritesRecordsAsync(t *testing.T) {
st := openTestStore(t)
defer st.Close()
queue := newDBWriteQueue(st)
record := textMessageTestRecord("queued")
queue.EnqueueRecord(record, MQTTClientInfo{ClientID: "client-1"})
record["text"] = "mutated after enqueue"
queue.Close()
var text, clientID string
if err := rawTestDB(t, st).QueryRow("SELECT text, mqtt_client_id FROM text_message WHERE from_id = ?", "!12345678").Scan(&text, &clientID); err != nil {
t.Fatal(err)
}
if text != "queued" || clientID != "client-1" {
t.Fatalf("queued row = text %q client %q, want queued/client-1", text, clientID)
}
}
func TestDBWriteQueueWritesDiscardAsync(t *testing.T) {
st := openTestStore(t)
defer st.Close()
queue := newDBWriteQueue(st)
record := map[string]any{"topic": "msh/test", "error": "bad packet"}
queue.EnqueueDiscard(record, []byte{1, 2, 3}, MQTTClientInfo{RemoteAddr: "127.0.0.1:1883"})
record["error"] = "mutated after enqueue"
queue.Close()
var topic, reason, rawBase64, remoteAddr string
if err := rawTestDB(t, st).QueryRow("SELECT topic, error, raw_base64, mqtt_remote_addr FROM discard_details").Scan(&topic, &reason, &rawBase64, &remoteAddr); err != nil {
t.Fatal(err)
}
if topic != "msh/test" || reason != "bad packet" || rawBase64 != "AQID" || remoteAddr != "127.0.0.1:1883" {
t.Fatalf("discard row = %q/%q/%q/%q, want queued values", topic, reason, rawBase64, remoteAddr)
}
}
func TestDBWriteQueueLen(t *testing.T) {
queue := &WriteQueue{jobs: make(chan writeJob, 1)}
queue.enqueue(writeJob{run: func() error { return nil }})
if queue.Len() != 1 {
t.Fatalf("queue.Len() = %d, want 1", queue.Len())
}
}
func TestDBWriteQueueIgnoresUnsupportedRecordType(t *testing.T) {
st := openTestStore(t)
defer st.Close()
queue := newDBWriteQueue(st)
queue.EnqueueRecord(map[string]any{"type": "empty_packet", "from": "!12345678"}, MQTTClientInfo{})
queue.Close()
var count int
if err := rawTestDB(t, st).QueryRow("SELECT COUNT(*) FROM text_message").Scan(&count); err != nil {
t.Fatal(err)
}
if count != 0 {
t.Fatalf("text_message count = %d, want 0", count)
}
}
func TestDBWriteQueueNilStore(t *testing.T) {
if queue := newDBWriteQueue(nil); queue != nil {
t.Fatalf("newDBWriteQueue(nil) = %#v, want nil", queue)
}
var queue *WriteQueue
queue.EnqueueRecord(textMessageTestRecord("ignored"), MQTTClientInfo{})
queue.EnqueueDiscard(map[string]any{"topic": "ignored"}, []byte{1}, MQTTClientInfo{})
queue.Close()
}
func TestDBWriteQueueRecordValidationErrorDoesNotStopWorker(t *testing.T) {
st := openTestStore(t)
defer st.Close()
queue := newDBWriteQueue(st)
badRecord := textMessageTestRecord("bad")
delete(badRecord, "from")
queue.EnqueueRecord(badRecord, MQTTClientInfo{})
queue.EnqueueRecord(textMessageTestRecord("good"), MQTTClientInfo{})
queue.Close()
var text string
if err := rawTestDB(t, st).QueryRow("SELECT text FROM text_message").Scan(&text); err != nil {
t.Fatal(err)
}
if text != "good" {
t.Fatalf("text = %q, want good", text)
}
var missing sql.NullString
if err := rawTestDB(t, st).QueryRow("SELECT text FROM text_message WHERE text = ?", "bad").Scan(&missing); err != sql.ErrNoRows {
t.Fatalf("bad row error = %v, want sql.ErrNoRows", err)
}
}
+2
View File
@@ -4,6 +4,7 @@ import (
"encoding/base64" "encoding/base64"
"encoding/json" "encoding/json"
"fmt" "fmt"
"time"
) )
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 {
@@ -34,6 +35,7 @@ func discardDetailsFromRecord(record map[string]any, raw []byte, clientInfo MQTT
MQTTRemoteAddr: NullableStringValue(clientInfo.RemoteAddr), MQTTRemoteAddr: NullableStringValue(clientInfo.RemoteAddr),
MQTTRemoteHost: NullableStringValue(clientInfo.RemoteHost), MQTTRemoteHost: NullableStringValue(clientInfo.RemoteHost),
MQTTRemotePort: NullableStringValue(clientInfo.RemotePort), MQTTRemotePort: NullableStringValue(clientInfo.RemotePort),
CreatedAt: time.Now(),
}, nil }, nil
} }
+7 -7
View File
@@ -306,8 +306,7 @@ func (s *Store) EnqueueLLMMessage(input LLMMessageQueueInput) (*LLMMessageQueueR
return nil, nil // 机器人的 LLM 队列未启用,静默返回 return nil, nil // 机器人的 LLM 队列未启用,静默返回
} }
// 忽略机器人自己发送的消息,避免自循环 if s.IsBotNodeID(input.FromNodeID) {
if input.FromNodeID == bot.NodeID {
return nil, nil return nil, nil
} }
@@ -370,6 +369,7 @@ func (s *Store) EnqueueLLMMessage(input LLMMessageQueueInput) (*LLMMessageQueueR
Status: LLMMessageStatusPending, Status: LLMMessageStatusPending,
ReceivedAt: now, ReceivedAt: now,
ContentJSON: input.ContentJSON, ContentJSON: input.ContentJSON,
CreatedAt: now,
} }
if err := s.db.Create(record).Error; err != nil { if err := s.db.Create(record).Error; err != nil {
@@ -397,7 +397,7 @@ func (s *Store) ListLLMMessages(opts ListOptions, botID uint64, includeDeleted b
} }
// 排序和分页 // 排序和分页
query = query.Order("created_at DESC") query = query.Order("id DESC")
if opts.Limit > 0 { if opts.Limit > 0 {
query = query.Limit(opts.Limit) query = query.Limit(opts.Limit)
} }
@@ -526,6 +526,10 @@ func enqueueChannelMessageToLLM(s *Store, record map[string]any) error {
shortName = &sn shortName = &sn
} }
if s.IsBotNodeID(fromNodeID) {
return nil
}
var channelID *string var channelID *string
if cid, ok := record["channel_id"].(string); ok && cid != "" { if cid, ok := record["channel_id"].(string); ok && cid != "" {
channelID = &cid channelID = &cid
@@ -546,11 +550,7 @@ func enqueueChannelMessageToLLM(s *Store, record map[string]any) error {
return fmt.Errorf("query bots for channel message enqueue: %w", err) return fmt.Errorf("query bots for channel message enqueue: %w", err)
} }
// 为每个符合条件的机器人创建一条队列记录(忽略机器人自己发送的消息)
for _, bot := range bots { for _, bot := range bots {
if fromNodeID == bot.NodeID {
continue
}
_, err = s.EnqueueLLMMessage(LLMMessageQueueInput{ _, err = s.EnqueueLLMMessage(LLMMessageQueueInput{
BotID: bot.ID, BotID: bot.ID,
BotNodeID: bot.NodeID, BotNodeID: bot.NodeID,
+6 -1
View File
@@ -1,13 +1,18 @@
package store package store
import "time"
func (s *Store) InsertLoginLog(log LoginLogRecord) error { func (s *Store) InsertLoginLog(log LoginLogRecord) error {
if log.CreatedAt.IsZero() {
log.CreatedAt = time.Now()
}
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("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)
} }
-258
View File
@@ -1,258 +0,0 @@
package store
import (
"errors"
"strings"
"testing"
"gorm.io/gorm"
)
func TestMapTileSourceDefaultSeeded(t *testing.T) {
st := openTestStore(t)
defer st.Close()
row, err := st.GetDefaultMapTileSource()
if err != nil {
t.Fatalf("GetDefaultMapTileSource() error = %v", err)
}
if row.Name != defaultMapTileSourceName || row.URLTemplate != defaultMapTileSourceURLTemplate || !row.Enabled || !row.IsDefault {
t.Fatalf("default map source = %+v, want built-in default", row)
}
}
func TestCreateMapTileSourceValidation(t *testing.T) {
st := openTestStore(t)
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 {
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 {
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 {
t.Fatal("CreateMapTileSource() credentials error = nil, want error")
}
}
func TestListEnabledMapTileSources(t *testing.T) {
st := openTestStore(t)
defer st.Close()
disabled, err := st.CreateMapTileSource(MapTileSourceInput{Name: "Disabled", URLTemplate: "https://disabled.example.com/{z}/{x}/{y}.png", MaxZoom: 18, Enabled: false})
if err != nil {
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})
if err != nil {
t.Fatalf("CreateMapTileSource(custom) error = %v", err)
}
if _, err := st.SetDefaultMapTileSource(custom.ID); err != nil {
t.Fatalf("SetDefaultMapTileSource() error = %v", err)
}
rows, err := st.ListEnabledMapTileSources()
if err != nil {
t.Fatalf("ListEnabledMapTileSources() error = %v", err)
}
if len(rows) < 2 {
t.Fatalf("ListEnabledMapTileSources() length = %d, want at least 2", len(rows))
}
if rows[0].ID != custom.ID {
t.Fatalf("first enabled source id = %d, want default %d", rows[0].ID, custom.ID)
}
for _, row := range rows {
if row.ID == disabled.ID {
t.Fatalf("disabled source was returned: %+v", row)
}
if !row.Enabled {
t.Fatalf("disabled row returned: %+v", row)
}
}
}
func TestMapTileSourceDuplicateAndDefaultRules(t *testing.T) {
st := openTestStore(t)
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})
if err != nil {
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) {
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) {
t.Fatalf("duplicate url error = %v, want ErrMapTileSourceAlreadyExists", err)
}
updated, err := st.SetDefaultMapTileSource(first.ID)
if err != nil {
t.Fatalf("SetDefaultMapTileSource() error = %v", err)
}
if !updated.IsDefault {
t.Fatalf("updated default = %+v, want is_default", updated)
}
oldDefault, err := st.GetDefaultMapTileSource()
if err != nil {
t.Fatalf("GetDefaultMapTileSource() error = %v", err)
}
if 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) {
t.Fatalf("disable default error = %v, want ErrMapTileSourceCannotDisableDefault", err)
}
if err := st.DeleteMapTileSource(first.ID); !errors.Is(err, ErrMapTileSourceCannotDeleteDefault) {
t.Fatalf("delete default error = %v, want ErrMapTileSourceCannotDeleteDefault", err)
}
}
func TestMapTileSourceHashIsSetOnCreate(t *testing.T) {
st := openTestStore(t)
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})
if err != nil {
t.Fatalf("CreateMapTileSource() error = %v", err)
}
want := MapTileSourceHash("https://test.example.com/{z}/{x}/{y}.png")
if row.URLTemplateHash != want {
t.Fatalf("URLTemplateHash = %q, want %q", row.URLTemplateHash, want)
}
if !row.ProxyEnabled {
t.Fatal("ProxyEnabled = false, want true")
}
}
func TestMapTileSourceDefaultHasHash(t *testing.T) {
st := openTestStore(t)
defer st.Close()
row, err := st.GetDefaultMapTileSource()
if err != nil {
t.Fatalf("GetDefaultMapTileSource() error = %v", err)
}
want := MapTileSourceHash(defaultMapTileSourceURLTemplate)
if row.URLTemplateHash != want {
t.Fatalf("default URLTemplateHash = %q, want %q", row.URLTemplateHash, want)
}
}
func TestGetEnabledMapTileSourceByHash(t *testing.T) {
st := openTestStore(t)
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})
if err != nil {
t.Fatalf("CreateMapTileSource() error = %v", err)
}
found, err := st.GetEnabledMapTileSourceByHash(row.URLTemplateHash)
if err != nil {
t.Fatalf("GetEnabledMapTileSourceByHash() error = %v", err)
}
if found.ID != row.ID {
t.Fatalf("found ID = %d, want %d", found.ID, row.ID)
}
}
func TestGetEnabledMapTileSourceByHashDisabled(t *testing.T) {
st := openTestStore(t)
defer st.Close()
row, err := st.CreateMapTileSource(MapTileSourceInput{Name: "DisabledHash", URLTemplate: "https://disabled-hash.example.com/{z}/{x}/{y}.png", MaxZoom: 18, Enabled: false})
if err != nil {
t.Fatalf("CreateMapTileSource() error = %v", err)
}
_, err = st.GetEnabledMapTileSourceByHash(row.URLTemplateHash)
if !errors.Is(err, gorm.ErrRecordNotFound) {
t.Fatalf("GetEnabledMapTileSourceByHash(disabled) = %v, want gorm.ErrRecordNotFound", err)
}
}
func TestGetEnabledMapTileSourceByHashProxyDisabled(t *testing.T) {
st := openTestStore(t)
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})
if err != nil {
t.Fatalf("CreateMapTileSource() error = %v", err)
}
_, err = st.GetEnabledMapTileSourceByHash(row.URLTemplateHash)
if !errors.Is(err, gorm.ErrRecordNotFound) {
t.Fatalf("GetEnabledMapTileSourceByHash(proxy disabled) = %v, want gorm.ErrRecordNotFound", err)
}
}
func TestGetEnabledMapTileSourceByHashUnknown(t *testing.T) {
st := openTestStore(t)
defer st.Close()
_, err := st.GetEnabledMapTileSourceByHash("aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa")
if !errors.Is(err, gorm.ErrRecordNotFound) {
t.Fatalf("GetEnabledMapTileSourceByHash(unknown) = %v, want gorm.ErrRecordNotFound", err)
}
}
func TestPublicMapTileSourceDTOProxyURL(t *testing.T) {
st := openTestStore(t)
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})
if err != nil {
t.Fatalf("CreateMapTileSource() error = %v", err)
}
dto := publicMapTileSourceDTO(*row)
urlTemplate, ok := dto["url_template"].(string)
if !ok {
t.Fatal("url_template is not a string")
}
wantPrefix := "/api/map/" + row.URLTemplateHash + "?x={x}&y={y}&z={z}"
if urlTemplate != wantPrefix {
t.Fatalf("url_template = %q, want %q", urlTemplate, wantPrefix)
}
if strings.Contains(urlTemplate, "proxy.example.com") {
t.Fatal("url_template should not contain upstream hostname")
}
}
func TestPublicMapTileSourceDTORawURLWhenProxyDisabled(t *testing.T) {
st := openTestStore(t)
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})
if err != nil {
t.Fatalf("CreateMapTileSource() error = %v", err)
}
dto := publicMapTileSourceDTO(*row)
urlTemplate, ok := dto["url_template"].(string)
if !ok {
t.Fatal("url_template is not a string")
}
if urlTemplate != row.URLTemplate {
t.Fatalf("url_template = %q, want raw %q", urlTemplate, row.URLTemplate)
}
}
func TestMapTileSourceHashFunction(t *testing.T) {
hash1 := 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")
if hash1 != hash2 {
t.Fatal("hash should be deterministic")
}
if len(hash1) != 64 {
t.Fatalf("hash length = %d, want 64", len(hash1))
}
if hash1 == hash3 {
t.Fatal("different URLs should produce different hashes")
}
}
@@ -1,38 +0,0 @@
package store
import "testing"
func TestRuntimeSettingsDefaultAndUpdates(t *testing.T) {
st := openTestStore(t)
defer st.Close()
settings, err := st.GetRuntimeSettings()
if err != nil {
t.Fatalf("GetRuntimeSettings() error = %v", err)
}
if settings.AllowEncryptedForwarding {
t.Fatalf("AllowEncryptedForwarding = true, want false")
}
if _, err := st.SetBoolRuntimeSetting(RuntimeSettingAllowEncryptedForwarding, true, "test setting"); err != nil {
t.Fatalf("SetBoolRuntimeSetting(true) error = %v", err)
}
settings, err = st.GetRuntimeSettings()
if err != nil {
t.Fatalf("GetRuntimeSettings() after true error = %v", err)
}
if !settings.AllowEncryptedForwarding {
t.Fatalf("AllowEncryptedForwarding = false, want true")
}
if _, err := st.SetBoolRuntimeSetting(RuntimeSettingAllowEncryptedForwarding, false, "test setting"); err != nil {
t.Fatalf("SetBoolRuntimeSetting(false) error = %v", err)
}
settings, err = st.GetRuntimeSettings()
if err != nil {
t.Fatalf("GetRuntimeSettings() after false error = %v", err)
}
if settings.AllowEncryptedForwarding {
t.Fatalf("AllowEncryptedForwarding = true, want false")
}
}
+19 -2
View File
@@ -316,7 +316,6 @@ func (s *Store) ListDiscardDetails(opts ListOptions) ([]DiscardDetailsRecord, er
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("id DESC"). Order("id DESC").
Limit(opts.Limit). Limit(opts.Limit).
Offset(opts.Offset) Offset(opts.Offset)
@@ -339,6 +338,19 @@ func applyDiscardDetailsFilters(q *gorm.DB, opts ListOptions) *gorm.DB {
return q return q
} }
func (s *Store) DeleteDiscardDetailsByIDs(ids []uint64) (int64, error) {
if len(ids) == 0 {
return 0, nil
}
result := s.db.Where("id IN ?", ids).Delete(&DiscardDetailsRecord{})
return result.RowsAffected, result.Error
}
func (s *Store) DeleteAllDiscardDetails() (int64, error) {
result := s.db.Where("1 = 1").Delete(&DiscardDetailsRecord{})
return result.RowsAffected, result.Error
}
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 {
@@ -372,7 +384,7 @@ func (s *Store) ListTraceroute(opts ListOptions) ([]TracerouteRecord, 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("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)
} }
@@ -387,3 +399,8 @@ func (s *Store) listAppendRows(opts ListOptions, dest any) *gorm.DB {
} }
return q.Find(dest) return q.Find(dest)
} }
func (s *Store) ListChannels() ([]ChannelRecord, error) {
var rows []ChannelRecord
return rows, s.db.Order("channel_id ASC").Find(&rows).Error
}
-45
View File
@@ -1,45 +0,0 @@
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
-27
View File
@@ -1,27 +0,0 @@
// Package testutil 提供给其它包测试使用的 store 临时实例工厂。
//
// 重构前 db_test.go 中的 openTestStore helper 被 8+ 个测试文件复用;
// 现在抽到这里,让 store 包外的测试也可以零样板地拿到一个临时 SQLite store。
package testutil
import (
"path/filepath"
"testing"
"meshtastic_mqtt_server/internal/config"
"meshtastic_mqtt_server/internal/store"
)
// OpenStore 返回一个写在 t.TempDir() 中的临时 SQLite store。
// 测试结束时调用方需要 defer st.Close()。
func OpenStore(t *testing.T) *store.Store {
t.Helper()
st, err := store.OpenStore(config.DatabaseConfig{
Driver: config.DriverSQLite,
SQLite: config.SQLiteConfig{Path: filepath.Join(t.TempDir(), "mesh_mqtt_go.db")},
}, false)
if err != nil {
t.Fatalf("OpenStore() error = %v", err)
}
return st
}
+1
View File
@@ -88,6 +88,7 @@ func RunAgentToolLoop(ctx context.Context, state *State, profile *llm.Profile, s
if routerPrompt == "" { if routerPrompt == "" {
routerPrompt = primaryPrompt routerPrompt = primaryPrompt
} }
routerPrompt = routerPrompt + "\n当前日期:" + time.Now().Format("2006-01-02")
if primaryPrompt != "" { if primaryPrompt != "" {
primarySystemMessage := &model.ChatCompletionMessage{ primarySystemMessage := &model.ChatCompletionMessage{
Role: "system", Role: "system",
+1 -1
View File
@@ -48,7 +48,7 @@ func NewState(cfg *Config, ai *llm.State, options ...Option) (*State, error) {
Enabled: true, Enabled: true,
Timeout: 30, Timeout: 30,
MaxTokens: 512, MaxTokens: 512,
SystemPrompt: "你可以按需直接调用可用工具来回答用户问题。\n每个工具的 description 描述了它的适用场景和调用条件。\n工具结果优先于模型内置知识;工具失败时必须如实说明,不要编造结果。\n只调用确实必要的工具。", SystemPrompt: "你可以按需直接调用可用工具来回答用户问题。\n每个工具的 description 描述了它的适用场景和调用条件。\n工具结果优先于模型内置知识;工具失败时必须如实说明,不要编造结果。\n只调用确实必要的工具。\n\n重要规则:\n- 当用户问题包含\"今天\"、\"昨天\"、\"最近\"、\"本周\"、\"本月\"等相对时间词时,你必须先调用 time 工具获取当前准确日期,然后再用该日期调用其他工具。绝不可以使用模型内置知识猜测日期。\n- 你不知道\"今天\"是哪一天,必须通过 time 工具查询。",
} }
} }
if ai == nil { if ai == nil {
-167
View File
@@ -1,167 +0,0 @@
package web
import (
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"testing"
configpkg "meshtastic_mqtt_server/internal/config"
storepkg "meshtastic_mqtt_server/internal/store"
"meshtastic_mqtt_server/internal/store/testutil"
)
func openTestStore(t *testing.T) *storepkg.Store {
return testutil.OpenStore(t)
}
func TestMapTileProxyFetchesAndCaches(t *testing.T) {
st := openTestStore(t)
defer st.Close()
requests := 0
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
requests++
if r.URL.Path != "/3/1/2.png" {
t.Fatalf("upstream path = %q, want /3/1/2.png", r.URL.Path)
}
w.Header().Set("Content-Type", "image/png")
_, _ = w.Write([]byte("tile-data"))
}))
defer upstream.Close()
row, err := st.CreateMapTileSource(storepkg.MapTileSourceInput{Name: "Tiles", URLTemplate: upstream.URL + "/{z}/{x}/{y}.png", MaxZoom: 18, Enabled: true, ProxyEnabled: true})
if err != nil {
t.Fatalf("CreateMapTileSource() error = %v", err)
}
cacheDir := t.TempDir()
router := NewRouter(configpkg.WebConfig{StaticDir: t.TempDir(), MapTileCacheDir: cacheDir}, false, st, nil, nil, nil, nil, nil, nil, nil)
url := "/api/map/" + row.URLTemplateHash + "?x=1&y=2&z=3"
for i := 0; i < 2; i++ {
recorder := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, url, nil)
router.ServeHTTP(recorder, req)
if recorder.Code != http.StatusOK {
t.Fatalf("request %d status = %d, body = %s", i+1, recorder.Code, recorder.Body.String())
}
if recorder.Body.String() != "tile-data" {
t.Fatalf("request %d body = %q, want tile-data", i+1, recorder.Body.String())
}
}
if requests != 1 {
t.Fatalf("upstream requests = %d, want 1", requests)
}
cachePath := filepath.Join(cacheDir, row.URLTemplateHash, "3", "1", "2.tile")
data, err := os.ReadFile(cachePath)
if err != nil {
t.Fatalf("read cache file %s: %v", cachePath, err)
}
if string(data) != "tile-data" {
t.Fatalf("cache file = %q, want tile-data", string(data))
}
}
func TestMapTileProxyRejectsInvalidCoordinates(t *testing.T) {
st := openTestStore(t)
defer st.Close()
row, err := st.CreateMapTileSource(storepkg.MapTileSourceInput{Name: "Tiles", URLTemplate: "https://tiles.example.com/{z}/{x}/{y}.png", MaxZoom: 3, Enabled: true, ProxyEnabled: true})
if err != nil {
t.Fatalf("CreateMapTileSource() error = %v", err)
}
router := NewRouter(configpkg.WebConfig{StaticDir: t.TempDir(), MapTileCacheDir: t.TempDir()}, false, st, nil, nil, nil, nil, nil, nil, nil)
cases := []string{
"/api/map/" + row.URLTemplateHash + "?y=0&z=0",
"/api/map/" + row.URLTemplateHash + "?x=-1&y=0&z=0",
"/api/map/" + row.URLTemplateHash + "?x=0&y=0&z=4",
"/api/map/" + row.URLTemplateHash + "?x=2&y=0&z=1",
}
for _, url := range cases {
recorder := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, url, nil)
router.ServeHTTP(recorder, req)
if recorder.Code != http.StatusBadRequest {
t.Fatalf("%s status = %d, want 400; body = %s", url, recorder.Code, recorder.Body.String())
}
}
}
func TestMapTileProxyUnknownAndDisabledSource(t *testing.T) {
st := openTestStore(t)
defer st.Close()
disabled, err := st.CreateMapTileSource(storepkg.MapTileSourceInput{Name: "Disabled", URLTemplate: "https://disabled.example.com/{z}/{x}/{y}.png", MaxZoom: 3, Enabled: false})
if err != nil {
t.Fatalf("CreateMapTileSource(disabled) error = %v", err)
}
proxyDisabled, err := st.CreateMapTileSource(storepkg.MapTileSourceInput{Name: "ProxyDisabled", URLTemplate: "https://proxy-disabled.example.com/{z}/{x}/{y}.png", MaxZoom: 3, Enabled: true, ProxyEnabled: false})
if err != nil {
t.Fatalf("CreateMapTileSource(proxy disabled) error = %v", err)
}
router := NewRouter(configpkg.WebConfig{StaticDir: t.TempDir(), MapTileCacheDir: t.TempDir()}, false, st, nil, nil, nil, nil, nil, nil, nil)
cases := []string{
"/api/map/not-a-hash?x=0&y=0&z=0",
"/api/map/aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa?x=0&y=0&z=0",
"/api/map/" + disabled.URLTemplateHash + "?x=0&y=0&z=0",
"/api/map/" + proxyDisabled.URLTemplateHash + "?x=0&y=0&z=0",
}
wantStatus := []int{http.StatusBadRequest, http.StatusNotFound, http.StatusNotFound, http.StatusNotFound}
for i, url := range cases {
recorder := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, url, nil)
router.ServeHTTP(recorder, req)
if recorder.Code != wantStatus[i] {
t.Fatalf("%s status = %d, want %d; body = %s", url, recorder.Code, wantStatus[i], recorder.Body.String())
}
}
}
func TestMapTileProxyUpstreamStatus(t *testing.T) {
st := openTestStore(t)
defer st.Close()
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if strings.Contains(r.URL.Path, "/404/") {
http.NotFound(w, r)
return
}
http.Error(w, "upstream error", http.StatusInternalServerError)
}))
defer upstream.Close()
row404, err := st.CreateMapTileSource(storepkg.MapTileSourceInput{Name: "NotFoundTiles", URLTemplate: upstream.URL + "/404/{z}/{x}/{y}.png", MaxZoom: 18, Enabled: true, ProxyEnabled: true})
if err != nil {
t.Fatalf("CreateMapTileSource(404) error = %v", err)
}
row500, err := st.CreateMapTileSource(storepkg.MapTileSourceInput{Name: "StatusTiles", URLTemplate: upstream.URL + "/{z}/{x}/{y}.png", MaxZoom: 18, Enabled: true, ProxyEnabled: true})
if err != nil {
t.Fatalf("CreateMapTileSource(500) error = %v", err)
}
router := NewRouter(configpkg.WebConfig{StaticDir: t.TempDir(), MapTileCacheDir: t.TempDir()}, false, st, nil, nil, nil, nil, nil, nil, nil)
cases := []struct {
url string
want int
}{
{url: "/api/map/" + row404.URLTemplateHash + "?x=0&y=0&z=0", want: http.StatusNotFound},
{url: "/api/map/" + row500.URLTemplateHash + "?x=0&y=0&z=0", want: http.StatusBadGateway},
}
for _, tc := range cases {
recorder := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, tc.url, nil)
router.ServeHTTP(recorder, req)
if recorder.Code != tc.want {
t.Fatalf("%s status = %d, want %d; body = %s", tc.url, recorder.Code, tc.want, recorder.Body.String())
}
}
}
+11 -1
View File
@@ -30,6 +30,7 @@ type MQTTRuntimeStatus struct {
Stats *mqttforwardpkg.Stats Stats *mqttforwardpkg.Stats
ClientStats *mqttforwardpkg.ClientStats ClientStats *mqttforwardpkg.ClientStats
DBQueue *storepkg.WriteQueue DBQueue *storepkg.WriteQueue
DedupQueue *mqttforwardpkg.DedupQueue
} }
// AdminMQTTStatus 是 admin 路由 GET /admin/mqtt-status 返回的 JSON 视图。 // AdminMQTTStatus 是 admin 路由 GET /admin/mqtt-status 返回的 JSON 视图。
@@ -50,6 +51,7 @@ type AdminMQTTStatus struct {
MessagesSent int64 `json:"messages_sent"` MessagesSent int64 `json:"messages_sent"`
MessagesDropped int64 `json:"messages_dropped"` MessagesDropped int64 `json:"messages_dropped"`
DBWriteQueueLength int `json:"db_write_queue_length"` DBWriteQueueLength int `json:"db_write_queue_length"`
DedupQueueLength int `json:"dedup_queue_len"`
Retained int64 `json:"retained"` Retained int64 `json:"retained"`
Inflight int64 `json:"inflight"` Inflight int64 `json:"inflight"`
InflightDropped int64 `json:"inflight_dropped"` InflightDropped int64 `json:"inflight_dropped"`
@@ -71,7 +73,7 @@ type AdminMQTTClient struct {
// Status 实现 MQTTStatusProvider。 // Status 实现 MQTTStatusProvider。
func (m MQTTRuntimeStatus) Status() AdminMQTTStatus { func (m MQTTRuntimeStatus) Status() AdminMQTTStatus {
if m.Server == nil || m.Server.Info == nil { if m.Server == nil || m.Server.Info == nil {
return AdminMQTTStatus{Running: false, Address: m.Address, TLS: m.TLS, DBWriteQueueLength: m.DBQueue.Len()} return AdminMQTTStatus{Running: false, Address: m.Address, TLS: m.TLS, DBWriteQueueLength: m.DBQueue.Len(), DedupQueueLength: m.dedupQueueLen()}
} }
info := m.Server.Info.Clone() info := m.Server.Info.Clone()
status := AdminMQTTStatus{ status := AdminMQTTStatus{
@@ -91,6 +93,7 @@ func (m MQTTRuntimeStatus) Status() AdminMQTTStatus {
MessagesSent: m.Stats.Forwarded(), MessagesSent: m.Stats.Forwarded(),
MessagesDropped: m.Stats.Dropped(), MessagesDropped: m.Stats.Dropped(),
DBWriteQueueLength: m.DBQueue.Len(), DBWriteQueueLength: m.DBQueue.Len(),
DedupQueueLength: m.dedupQueueLen(),
Retained: info.Retained, Retained: info.Retained,
Inflight: info.Inflight, Inflight: info.Inflight,
InflightDropped: info.InflightDropped, InflightDropped: info.InflightDropped,
@@ -136,6 +139,13 @@ func mqttClientInfo(c *mqtt.Client) mqttClientInfoView {
} }
} }
func (m MQTTRuntimeStatus) dedupQueueLen() int {
if m.DedupQueue == nil {
return 0
}
return m.DedupQueue.Len()
}
// DisconnectClient 实现 MQTTStatusProvider:发送 Disconnect 报文并关闭连接。 // DisconnectClient 实现 MQTTStatusProvider:发送 Disconnect 报文并关闭连接。
// 使用 ErrAdministrativeAction 作为断开理由,便于日志区分。 // 使用 ErrAdministrativeAction 作为断开理由,便于日志区分。
func (m MQTTRuntimeStatus) DisconnectClient(clientID string) bool { func (m MQTTRuntimeStatus) DisconnectClient(clientID string) bool {
+50
View File
@@ -13,6 +13,7 @@ import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"gorm.io/gorm" "gorm.io/gorm"
aipkg "meshtastic_mqtt_server/internal/ai"
"meshtastic_mqtt_server/internal/auth" "meshtastic_mqtt_server/internal/auth"
blockingpkg "meshtastic_mqtt_server/internal/blocking" blockingpkg "meshtastic_mqtt_server/internal/blocking"
botpkg "meshtastic_mqtt_server/internal/bot" botpkg "meshtastic_mqtt_server/internal/bot"
@@ -32,6 +33,8 @@ type LLMProviderReloader interface {
ReloadLLMProvider(config interface{}) error ReloadLLMProvider(config interface{}) error
AddLLMProvider(config interface{}) error AddLLMProvider(config interface{}) error
RemoveLLMProvider(name string) error RemoveLLMProvider(name string) error
AIServiceStatus() aipkg.AIServiceStatus
RestartAIService() error
} }
func NewHTTPServer(cfg configpkg.WebConfig, consoleLog bool, store *storepkg.Store, sessions *auth.Manager, mqttStatus MQTTStatusProvider, blocking *blockingpkg.Cache, forwarder mqttforwardpkg.Reloader, settings *rspkg.Cache, botSender botpkg.TextSender, aiService LLMProviderReloader) *http.Server { func NewHTTPServer(cfg configpkg.WebConfig, consoleLog bool, store *storepkg.Store, sessions *auth.Manager, mqttStatus MQTTStatusProvider, blocking *blockingpkg.Cache, forwarder mqttforwardpkg.Reloader, settings *rspkg.Cache, botSender botpkg.TextSender, aiService LLMProviderReloader) *http.Server {
@@ -81,6 +84,10 @@ func NewRouter(cfg configpkg.WebConfig, consoleLog bool, store *storepkg.Store,
return r return r
} }
const BackendVersion = "1.2.1"
var CommitVersion = "dev"
func registerAPIRoutes(r gin.IRouter, store *storepkg.Store, mapTileCacheDir string) { func registerAPIRoutes(r gin.IRouter, store *storepkg.Store, mapTileCacheDir string) {
r.GET("/health", func(c *gin.Context) { r.GET("/health", func(c *gin.Context) {
status := gin.H{"status": "ok", "database": "ok"} status := gin.H{"status": "ok", "database": "ok"}
@@ -93,6 +100,10 @@ func registerAPIRoutes(r gin.IRouter, store *storepkg.Store, mapTileCacheDir str
c.JSON(http.StatusOK, status) c.JSON(http.StatusOK, status)
}) })
r.GET("/version", func(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"version": BackendVersion, "commit": CommitVersion})
})
registerNodeInfoRoutes(r, store, "/nodeinfo") registerNodeInfoRoutes(r, store, "/nodeinfo")
registerNodeInfoRoutes(r, store, "/nodes") registerNodeInfoRoutes(r, store, "/nodes")
registerMapReportRoutes(r, store) registerMapReportRoutes(r, store)
@@ -128,6 +139,18 @@ func registerAPIRoutes(r gin.IRouter, store *storepkg.Store, mapTileCacheDir str
rows, err := store.ListTextMessages(opts) rows, err := store.ListTextMessages(opts)
writeListResponse(c, rows, opts, err, textMessageDTO) writeListResponse(c, rows, opts, err, textMessageDTO)
}) })
r.GET("/channels", func(c *gin.Context) {
rows, err := store.ListChannels()
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
items := make([]gin.H, 0, len(rows))
for _, row := range rows {
items = append(items, gin.H{"channel_id": row.ChannelID})
}
c.JSON(http.StatusOK, gin.H{"items": items})
})
r.GET("/discard-details", func(c *gin.Context) { r.GET("/discard-details", func(c *gin.Context) {
opts, ok := parseListOptions(c) opts, ok := parseListOptions(c)
if !ok { if !ok {
@@ -376,6 +399,33 @@ func registerAdminRoutes(r gin.IRouter, store *storepkg.Store, sessions *auth.Ma
rows, err := store.ListLoginLogs(opts) rows, err := store.ListLoginLogs(opts)
writeListResponse(c, rows, opts, err, loginLogDTO) writeListResponse(c, rows, opts, err, loginLogDTO)
}) })
protected.POST("/discard-details/batch-delete", func(c *gin.Context) {
var req struct {
IDs []uint64 `json:"ids"`
}
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
if len(req.IDs) == 0 {
c.JSON(http.StatusBadRequest, gin.H{"error": "ids is empty"})
return
}
count, err := store.DeleteDiscardDetailsByIDs(req.IDs)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"status": "ok", "deleted_count": count})
})
protected.DELETE("/discard-details", func(c *gin.Context) {
count, err := store.DeleteAllDiscardDetails()
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"status": "ok", "deleted_count": count})
})
protected.DELETE("/text-messages/:id", func(c *gin.Context) { protected.DELETE("/text-messages/:id", func(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64) id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil || id == 0 { if err != nil || id == 0 {
+36 -30
View File
@@ -59,6 +59,7 @@ type meshtasticFilterHook struct {
autoAcker func(record map[string]any) autoAcker func(record map[string]any)
consoleLog bool // 控制台是否打印 MQTT 连接/订阅事件 consoleLog bool // 控制台是否打印 MQTT 连接/订阅事件
packetConsoleLog bool // 控制台是否打印 Meshtastic 数据包 packetConsoleLog bool // 控制台是否打印 Meshtastic 数据包
dedupQueue *mqttforwardpkg.DedupQueue
} }
// ID 返回用于识别 Meshtastic payload 过滤器的 hook 名称。 // ID 返回用于识别 Meshtastic payload 过滤器的 hook 名称。
@@ -225,6 +226,10 @@ func (h *meshtasticFilterHook) OnPublish(cl *mqtt.Client, pk packets.Packet) (pa
h.rejectPublish(cl, pk, record) h.rejectPublish(cl, pk, record)
return pk, packets.ErrRejectPacket return pk, packets.ErrRejectPacket
} }
if h.dedupQueue != nil && !h.dedupQueue.TryForward(pk.TopicName, pk.Payload) {
h.stats.IncDropped()
return pk, packets.ErrRejectPacket
}
h.stats.IncForwarded() h.stats.IncForwarded()
h.dbQueue.EnqueueRecord(record, mqttClientInfoFromClient(cl)) h.dbQueue.EnqueueRecord(record, mqttClientInfoFromClient(cl))
@@ -393,27 +398,6 @@ func run(cfg *configpkg.Config) error {
defer forwardManager.StopAll() defer forwardManager.StopAll()
// Initialize AI Service // Initialize AI Service
var aiService *ai.Service
if cfg.AI.Enabled {
// Get LLM providers from database
llmProviders, err := store.ListLLMProviders(true)
if err != nil {
fmt.Fprintf(os.Stderr, "Warning: failed to load LLM providers: %v\n", err)
} else if len(llmProviders) > 0 {
// Convert database records to provider configs
providerConfigs := make([]llm.ProviderConfig, 0, len(llmProviders))
for _, p := range llmProviders {
providerConfigs = append(providerConfigs, llm.ProviderConfig{
Name: p.Name,
Active: p.Active,
APIKey: p.APIKey,
BaseURL: p.BaseURL,
Model: p.Model,
Timeout: p.Timeout,
ContextWindowTokens: p.ContextWindowTokens,
})
}
// Create bot sender adapter - 支持频道消息和私聊消息两种发送方式 // Create bot sender adapter - 支持频道消息和私聊消息两种发送方式
botSenderAdapter := autoreply.NewBotServiceAdapter( botSenderAdapter := autoreply.NewBotServiceAdapter(
// SendDirectText: 发送私聊消息 // SendDirectText: 发送私聊消息
@@ -438,8 +422,7 @@ func run(cfg *configpkg.Config) error {
}, },
) )
aiService, err = ai.NewService(ai.Config{ aiManager := ai.NewAIManager(ai.Config{
LLMProviders: providerConfigs,
DataDir: cfg.AI.DataDir, DataDir: cfg.AI.DataDir,
Enabled: cfg.AI.Enabled, Enabled: cfg.AI.Enabled,
ConsoleLog: cfg.ConsoleLog.LLM, ConsoleLog: cfg.ConsoleLog.LLM,
@@ -447,14 +430,31 @@ func run(cfg *configpkg.Config) error {
ToolRouterStore: store, ToolRouterStore: store,
TopicRouterStore: store, TopicRouterStore: store,
Store: store, Store: store,
}, store.DB(), botSenderAdapter) }, store.DB(), botSenderAdapter, botCtx, store)
if cfg.AI.Enabled {
aiManager.SetConfigEnabled(true)
providers, err := store.ListLLMProviders(true)
if err != nil { if err != nil {
fmt.Fprintf(os.Stderr, "Warning: failed to load LLM providers: %v\n", err)
} else if len(providers) > 0 {
providerConfigs := make([]llm.ProviderConfig, 0, len(providers))
for _, p := range providers {
providerConfigs = append(providerConfigs, llm.ProviderConfig{
Name: p.Name,
Active: p.Active,
APIKey: p.APIKey,
BaseURL: p.BaseURL,
Model: p.Model,
Timeout: p.Timeout,
ContextWindowTokens: p.ContextWindowTokens,
})
}
aiManager.SetProviderConfigs(providerConfigs)
if err := aiManager.Init(); 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 {
if err := aiService.Start(botCtx); err != nil { defer aiManager.Stop()
fmt.Fprintf(os.Stderr, "Warning: failed to start AI service: %v\n", err)
}
defer aiService.Stop()
printJSON(map[string]any{"event": "ai_service_started", "providers": len(providerConfigs)}) printJSON(map[string]any{"event": "ai_service_started", "providers": len(providerConfigs)})
} }
} else { } else {
@@ -469,8 +469,8 @@ func run(cfg *configpkg.Config) error {
if err != nil { if err != nil {
return err return err
} }
mqttStatus := webpkg.MQTTRuntimeStatus{Server: server, Address: mqttAddr, TLS: cfg.MQTT.TLS.Enabled, Stats: messageStats, ClientStats: clientStats, DBQueue: dbQueue} mqttStatus := webpkg.MQTTRuntimeStatus{Server: server, Address: mqttAddr, TLS: cfg.MQTT.TLS.Enabled, Stats: messageStats, ClientStats: clientStats, DBQueue: dbQueue, DedupQueue: mqttHook.dedupQueue}
handler := webpkg.NewRouter(cfg.Web, cfg.ConsoleLog.Web, store, sessions, mqttStatus, blocking, forwardManager, settings, botSender, aiService) handler := webpkg.NewRouter(cfg.Web, cfg.ConsoleLog.Web, store, sessions, mqttStatus, blocking, forwardManager, settings, botSender, aiManager)
webAddresses := []string{} webAddresses := []string{}
if cfg.Web.PortEnabled { if cfg.Web.PortEnabled {
httpServer := &http.Server{ httpServer := &http.Server{
@@ -521,6 +521,9 @@ func run(cfg *configpkg.Config) error {
if err := server.Close(); err != nil && runErr == nil { if err := server.Close(); err != nil && runErr == nil {
runErr = err runErr = err
} }
if mqttHook.dedupQueue != nil {
mqttHook.dedupQueue.Stop()
}
return runErr return runErr
} }
@@ -529,6 +532,8 @@ func startMQTTServer(cfg *configpkg.Config, store *storepkg.Store, dbQueue *stor
if err := server.AddHook(new(mqttauth.AllowHook), nil); err != nil { if err := server.AddHook(new(mqttauth.AllowHook), nil); err != nil {
return nil, nil, "", err return nil, nil, "", err
} }
dedupQueue := mqttforwardpkg.NewDedupQueue()
dedupQueue.Start()
hook := &meshtasticFilterHook{ hook := &meshtasticFilterHook{
server: server, server: server,
key: cfg.Key, key: cfg.Key,
@@ -540,6 +545,7 @@ func startMQTTServer(cfg *configpkg.Config, store *storepkg.Store, dbQueue *stor
pkiResolver: botpkg.NewPKIKeyResolver(store), pkiResolver: botpkg.NewPKIKeyResolver(store),
consoleLog: cfg.ConsoleLog.MQTT, consoleLog: cfg.ConsoleLog.MQTT,
packetConsoleLog: cfg.ConsoleLog.Meshtastic, packetConsoleLog: cfg.ConsoleLog.Meshtastic,
dedupQueue: dedupQueue,
} }
if err := server.AddHook(hook, nil); err != nil { if err := server.AddHook(hook, nil); err != nil {
return nil, nil, "", err return nil, nil, "", err
-113
View File
@@ -1,113 +0,0 @@
package main
import (
"testing"
mqtt "github.com/mochi-mqtt/server/v2"
blockingpkg "meshtastic_mqtt_server/internal/blocking"
storepkg "meshtastic_mqtt_server/internal/store"
"meshtastic_mqtt_server/internal/store/testutil"
)
func TestMQTTClientInfoFromClientNil(t *testing.T) {
info := mqttClientInfoFromClient(nil)
if info != (storepkg.MQTTClientInfo{}) {
t.Fatalf("info = %#v, want zero value", info)
}
}
func TestMQTTClientInfoFromClientIPv4(t *testing.T) {
info := mqttClientInfoFromClient(&mqtt.Client{
ID: "client-1",
Properties: mqtt.ClientProperties{Username: []byte("user-1")},
Net: mqtt.ClientConnection{Listener: "tcp", Remote: "127.0.0.1:1234"},
})
if info.ClientID != "client-1" || info.Username != "user-1" || info.Listener != "tcp" {
t.Fatalf("client fields = %#v", info)
}
if info.RemoteAddr != "127.0.0.1:1234" || info.RemoteHost != "127.0.0.1" || info.RemotePort != "1234" {
t.Fatalf("remote fields = %#v", info)
}
}
func TestMQTTClientInfoFromClientIPv6(t *testing.T) {
info := mqttClientInfoFromClient(&mqtt.Client{Net: mqtt.ClientConnection{Remote: "[::1]:1234"}})
if info.RemoteHost != "::1" || info.RemotePort != "1234" {
t.Fatalf("remote fields = %#v, want host ::1 and port 1234", info)
}
}
func TestMQTTClientInfoFromClientUnsplitRemote(t *testing.T) {
info := mqttClientInfoFromClient(&mqtt.Client{Net: mqtt.ClientConnection{Remote: "localhost"}})
if info.RemoteHost != "localhost" || info.RemotePort != "" {
t.Fatalf("remote fields = %#v, want host localhost and empty port", info)
}
}
// blockingViolationForRecord 的测试用真实 *Store + blocking.Cache 走完整路径,
// 不依赖 cache 的未导出字段。
func TestBlockingViolationForRecordNode(t *testing.T) {
st := testutil.OpenStore(t)
defer st.Close()
nodeNum := int64(305419896)
if _, err := st.CreateNodeBlocking("!12345678", &nodeNum, "blocked", true); err != nil {
t.Fatalf("CreateNodeBlocking() error = %v", err)
}
cache, err := blockingpkg.New(st)
if err != nil {
t.Fatalf("blocking.New() error = %v", err)
}
record := map[string]any{"type": "position", "from": "!12345678", "from_num": uint32(305419896)}
violation := blockingViolationForRecord(cache, record)
if violation == nil || violation["blocking_type"] != "node" {
t.Fatalf("blockingViolationForRecord() = %#v, want node violation", violation)
}
}
func TestBlockingViolationForRecordForbiddenWordFields(t *testing.T) {
st := testutil.OpenStore(t)
defer st.Close()
if _, err := st.CreateForbiddenWordBlocking("spam", "contains", false, "blocked", true); err != nil {
t.Fatalf("CreateForbiddenWordBlocking() error = %v", err)
}
cache, err := blockingpkg.New(st)
if err != nil {
t.Fatalf("blocking.New() error = %v", err)
}
for _, tc := range []struct {
name string
record map[string]any
field string
}{
{name: "text", record: map[string]any{"type": "text_message", "from": "!1", "text": "has SPAM"}, field: "text"},
{name: "nodeinfo", record: map[string]any{"type": "nodeinfo", "from": "!1", "long_name": "has SPAM"}, field: "long_name"},
{name: "map_report", record: map[string]any{"type": "map_report", "from": "!1", "long_name": "has SPAM"}, field: "long_name"},
} {
t.Run(tc.name, func(t *testing.T) {
violation := blockingViolationForRecord(cache, tc.record)
if violation == nil || violation["blocking_type"] != "forbidden_word" || violation["blocking_field"] != tc.field || violation["matched_word"] != "spam" {
t.Fatalf("blockingViolationForRecord() = %#v, want forbidden word on %s", violation, tc.field)
}
})
}
}
func TestBlockingViolationForRecordAllowed(t *testing.T) {
st := testutil.OpenStore(t)
defer st.Close()
if _, err := st.CreateForbiddenWordBlocking("spam", "contains", false, "blocked", true); err != nil {
t.Fatalf("CreateForbiddenWordBlocking() error = %v", err)
}
cache, err := blockingpkg.New(st)
if err != nil {
t.Fatalf("blocking.New() error = %v", err)
}
record := map[string]any{"type": "text_message", "from": "!1", "text": "hello"}
if violation := blockingViolationForRecord(cache, record); violation != nil {
t.Fatalf("blockingViolationForRecord() = %#v, want nil", violation)
}
}
+31 -7
View File
@@ -1,6 +1,6 @@
<script setup lang="ts"> <script setup lang="ts">
import { computed, onBeforeUnmount, onMounted, ref } from 'vue' import { computed, onBeforeUnmount, onMounted, ref, watch } from 'vue'
import { adminLogout, createNodeBlockingRule, deleteNode, deleteTextMessage, getAdminMe, getHealth, getMapReportViewport, getNodeInfo, getPositions, getTextMessages, purgeNode } from './api' import { adminLogout, createNodeBlockingRule, deleteNode, deleteTextMessage, getAdminMe, getChannels, getHealth, getMapReportViewport, getNodeInfo, getPositions, getTextMessages, purgeNode } from './api'
import AdminBlockingManagement from './components/AdminBlockingManagement.vue' import AdminBlockingManagement from './components/AdminBlockingManagement.vue'
import AdminBot from './components/AdminBot.vue' import AdminBot from './components/AdminBot.vue'
import AdminBotDirect from './components/AdminBotDirect.vue' import AdminBotDirect from './components/AdminBotDirect.vue'
@@ -15,6 +15,7 @@ import AdminMapSource from './components/AdminMapSource.vue'
import AdminMqttForward from './components/AdminMqttForward.vue' import AdminMqttForward from './components/AdminMqttForward.vue'
import AdminSignManagement from './components/AdminSignManagement.vue' import AdminSignManagement from './components/AdminSignManagement.vue'
import AdminUsers from './components/AdminUsers.vue' import AdminUsers from './components/AdminUsers.vue'
import AppFooter from './components/AppFooter.vue'
import ChatPanel from './components/ChatPanel.vue' import ChatPanel from './components/ChatPanel.vue'
import ConfirmDeleteModal from './components/ConfirmDeleteModal.vue' import ConfirmDeleteModal from './components/ConfirmDeleteModal.vue'
import HelpPage from './components/HelpPage.vue' import HelpPage from './components/HelpPage.vue'
@@ -56,6 +57,7 @@ const nodePage = ref(1)
const nodePageSize = 25 const nodePageSize = 25
const nodeTotal = ref(0) const nodeTotal = ref(0)
const messages = ref<TextMessage[]>([]) const messages = ref<TextMessage[]>([])
const channels = ref<string[]>([])
const chatPageSize = 20 const chatPageSize = 20
const chatLoadingOlder = ref(false) const chatLoadingOlder = ref(false)
const chatHasMore = ref(true) const chatHasMore = ref(true)
@@ -68,6 +70,7 @@ const mapReportTotal = ref(0)
const mapSources = ref<PublicMapTileSource[]>([fallbackMapSource]) const mapSources = ref<PublicMapTileSource[]>([fallbackMapSource])
const mapSource = ref<PublicMapTileSource>(fallbackMapSource) const mapSource = ref<PublicMapTileSource>(fallbackMapSource)
const nodeFilter = ref('') const nodeFilter = ref('')
const channelFilter = ref('')
const pendingDeleteAction = ref<PendingDeleteAction | null>(null) const pendingDeleteAction = ref<PendingDeleteAction | null>(null)
type DeletableTextMessage = TextMessage & { mergedCount?: number; mergedMessages?: TextMessage[] } type DeletableTextMessage = TextMessage & { mergedCount?: number; mergedMessages?: TextMessage[] }
type NodeActionRequest = { nodeId: string; nodeNum: number | null; message?: DeletableTextMessage } type NodeActionRequest = { nodeId: string; nodeNum: number | null; message?: DeletableTextMessage }
@@ -94,6 +97,7 @@ const nodesById = computed<NodeInfoById>(() => {
}) })
const normalizedNodeFilter = computed(() => nodeFilter.value.trim().toLowerCase()) const normalizedNodeFilter = computed(() => nodeFilter.value.trim().toLowerCase())
const normalizedChannelFilter = computed(() => channelFilter.value.trim())
function nodeMatchesFilterByInfo(node: NodeInfo | null | undefined, keyword: string): boolean { function nodeMatchesFilterByInfo(node: NodeInfo | null | undefined, keyword: string): boolean {
if (!keyword) { if (!keyword) {
@@ -253,8 +257,7 @@ function toChronological(items: TextMessage[]): TextMessage[] {
} }
function compareMessages(a: TextMessage, b: TextMessage): number { function compareMessages(a: TextMessage, b: TextMessage): number {
const timeDiff = Date.parse(a.created_at) - Date.parse(b.created_at) return a.id - b.id
return timeDiff !== 0 ? timeDiff : a.id - b.id
} }
function mergeMessages(existing: TextMessage[], incoming: TextMessage[]): TextMessage[] { function mergeMessages(existing: TextMessage[], incoming: TextMessage[]): TextMessage[] {
@@ -297,7 +300,7 @@ function clearSelectedNode() {
} }
async function loadInitialChatMessages() { async function loadInitialChatMessages() {
const response = await getTextMessages(chatPageSize, 0) const response = await getTextMessages(chatPageSize, 0, normalizedChannelFilter.value ? { channelId: normalizedChannelFilter.value } : '')
messages.value = toChronological(response.items) messages.value = toChronological(response.items)
chatHasMore.value = response.items.length === chatPageSize chatHasMore.value = response.items.length === chatPageSize
chatInitialized.value = true chatInitialized.value = true
@@ -310,7 +313,7 @@ async function loadOlderMessages() {
chatLoadingOlder.value = true chatLoadingOlder.value = true
try { try {
const response = await getTextMessages(chatPageSize, messages.value.length) const response = await getTextMessages(chatPageSize, messages.value.length, normalizedChannelFilter.value ? { channelId: normalizedChannelFilter.value } : '')
messages.value = mergeMessages(messages.value, toChronological(response.items)) messages.value = mergeMessages(messages.value, toChronological(response.items))
chatHasMore.value = response.items.length === chatPageSize chatHasMore.value = response.items.length === chatPageSize
} catch (err) { } catch (err) {
@@ -321,10 +324,23 @@ async function loadOlderMessages() {
} }
async function pollLatestMessages() { async function pollLatestMessages() {
const response = await getTextMessages(chatPageSize, 0) const response = await getTextMessages(chatPageSize, 0, normalizedChannelFilter.value ? { channelId: normalizedChannelFilter.value } : '')
messages.value = mergeMessages(messages.value, toChronological(response.items)) messages.value = mergeMessages(messages.value, toChronological(response.items))
} }
let channelFilterTimer: number | undefined
watch(normalizedChannelFilter, () => {
if (channelFilterTimer !== undefined) {
window.clearTimeout(channelFilterTimer)
}
channelFilterTimer = window.setTimeout(() => {
messages.value = []
chatHasMore.value = true
chatInitialized.value = false
loadInitialChatMessages()
}, 400)
})
async function loadNodePage(page: number, showLoading = true) { async function loadNodePage(page: number, showLoading = true) {
if (showLoading) { if (showLoading) {
nodePageLoading.value = true nodePageLoading.value = true
@@ -666,6 +682,7 @@ onMounted(() => {
loadMapSource() loadMapSource()
refresh() refresh()
refreshTimer = window.setInterval(() => refresh(false), 5000) refreshTimer = window.setInterval(() => refresh(false), 5000)
getChannels().then(res => { channels.value = res.items.map(i => i.channel_id) }).catch(() => {})
}) })
onBeforeUnmount(() => { onBeforeUnmount(() => {
@@ -675,6 +692,9 @@ onBeforeUnmount(() => {
if (mapBoundsTimer !== undefined) { if (mapBoundsTimer !== undefined) {
window.clearTimeout(mapBoundsTimer) window.clearTimeout(mapBoundsTimer)
} }
if (channelFilterTimer !== undefined) {
window.clearTimeout(channelFilterTimer)
}
}) })
</script> </script>
@@ -791,6 +811,8 @@ onBeforeUnmount(() => {
<section class="workspace"> <section class="workspace">
<ChatPanel <ChatPanel
v-model:channelFilter="channelFilter"
:channels="channels"
:messages="filteredMessages" :messages="filteredMessages"
:nodes-by-id="nodesById" :nodes-by-id="nodesById"
:selected-node-id="selectedNodeId" :selected-node-id="selectedNodeId"
@@ -838,6 +860,8 @@ onBeforeUnmount(() => {
/> />
</template> </template>
<AppFooter />
<ConfirmDeleteModal <ConfirmDeleteModal
:open="!!pendingDeleteAction" :open="!!pendingDeleteAction"
:title="deleteModalTitle" :title="deleteModalTitle"
+28 -2
View File
@@ -6,6 +6,7 @@ import type {
AdminRuntimeSettingsPayload, AdminRuntimeSettingsPayload,
AdminRuntimeSettingsResponse, AdminRuntimeSettingsResponse,
AdminUsersResponse, AdminUsersResponse,
AIServiceStatus,
BotMessage, BotMessage,
BotMessageMutationResponse, BotMessageMutationResponse,
BotNode, BotNode,
@@ -130,6 +131,14 @@ export function getHealth(): Promise<HealthStatus> {
return getJSON<HealthStatus>('/api/health') return getJSON<HealthStatus>('/api/health')
} }
export function getBackendVersion(): Promise<{ version: string; commit: string }> {
return getJSON<{ version: string; commit: string }>('/api/version')
}
export function getChannels(): Promise<{ items: { channel_id: string }[] }> {
return getJSON<{ items: { channel_id: string }[] }>('/api/channels')
}
export function getHelpContent(): Promise<HelpContentResponse> { export function getHelpContent(): Promise<HelpContentResponse> {
return getJSON<HelpContentResponse>('/api/help') return getJSON<HelpContentResponse>('/api/help')
} }
@@ -226,6 +235,14 @@ export function getDiscardDetails(limit = 100, offset = 0): Promise<ListResponse
return getJSON<ListResponse<DiscardDetails>>(listPath('/api/discard-details', limit, offset)) return getJSON<ListResponse<DiscardDetails>>(listPath('/api/discard-details', limit, offset))
} }
export function deleteDiscardDetailsByIDs(ids: number[]): Promise<{ status: string; deleted_count: number }> {
return postJSON<{ status: string; deleted_count: number }>('/api/admin/discard-details/batch-delete', { ids })
}
export function clearDiscardDetails(): Promise<{ status: string; deleted_count: number }> {
return deleteJSON<{ status: string; deleted_count: number }>('/api/admin/discard-details')
}
export function getTelemetry(limit = 500, offset = 0, nodeId = ''): Promise<ListResponse<TelemetryRecord>> { export function getTelemetry(limit = 500, offset = 0, nodeId = ''): Promise<ListResponse<TelemetryRecord>> {
return getJSON<ListResponse<TelemetryRecord>>(listPath('/api/telemetry', limit, offset, nodeId)) return getJSON<ListResponse<TelemetryRecord>>(listPath('/api/telemetry', limit, offset, nodeId))
} }
@@ -515,8 +532,17 @@ export function updateLLMProvider(name: string, payload: Partial<LLMProviderPayl
return putJSON<LLMProviderResponse>(`/api/admin/llm/providers/${encodeURIComponent(name)}`, payload) return putJSON<LLMProviderResponse>(`/api/admin/llm/providers/${encodeURIComponent(name)}`, payload)
} }
export function deleteLLMProvider(name: string): Promise<{ status: string }> { export function deleteLLMProvider(name: string): Promise<{ status: string; warning?: string }> {
return deleteJSON<{ status: string }>(`/api/admin/llm/providers/${encodeURIComponent(name)}`) return deleteJSON<{ status: string; warning?: string }>(`/api/admin/llm/providers/${encodeURIComponent(name)}`)
}
// AI Service Status API
export function getAIServiceStatus(): Promise<AIServiceStatus> {
return getJSON<AIServiceStatus>('/api/admin/llm/status')
}
export function restartAIService(): Promise<AIServiceStatus & { status: string; message?: string }> {
return postJSON<AIServiceStatus & { status: string; message?: string }>('/api/admin/llm/restart')
} }
// LLM Tool Router API // LLM Tool Router API
@@ -200,6 +200,7 @@ onBeforeUnmount(() => {
<div><span>订阅数</span><strong>{{ status.subscriptions }}</strong></div> <div><span>订阅数</span><strong>{{ status.subscriptions }}</strong></div>
<div><span>转发消息</span><strong>{{ status.messages_sent }}</strong></div> <div><span>转发消息</span><strong>{{ status.messages_sent }}</strong></div>
<div><span>数据库队列</span><strong>{{ status.db_write_queue_length }}</strong></div> <div><span>数据库队列</span><strong>{{ status.db_write_queue_length }}</strong></div>
<div><span>去重队列</span><strong>{{ status.dedup_queue_len }}</strong></div>
<a class="status-card-link" href="/admin/discard_details"><span>丢弃消息</span><strong>{{ status.messages_dropped }}</strong></a> <a class="status-card-link" href="/admin/discard_details"><span>丢弃消息</span><strong>{{ status.messages_dropped }}</strong></a>
<div><span>收到包</span><strong>{{ status.packets_received }}</strong></div> <div><span>收到包</span><strong>{{ status.packets_received }}</strong></div>
<div><span>发送包</span><strong>{{ status.packets_sent }}</strong></div> <div><span>发送包</span><strong>{{ status.packets_sent }}</strong></div>
@@ -1,7 +1,8 @@
<script setup lang="ts"> <script setup lang="ts">
import { onMounted, ref } from 'vue' import { computed, onMounted, ref } from 'vue'
import { getDiscardDetails } from '../api' import { clearDiscardDetails, deleteDiscardDetailsByIDs, getDiscardDetails } from '../api'
import type { DiscardDetails } from '../types' import type { DiscardDetails } from '../types'
import ConfirmDeleteModal from './ConfirmDeleteModal.vue'
const items = ref<DiscardDetails[]>([]) const items = ref<DiscardDetails[]>([])
const loading = ref(false) const loading = ref(false)
@@ -9,9 +10,48 @@ const error = ref('')
const page = ref(1) const page = ref(1)
const pageSize = 25 const pageSize = 25
const selectedIds = ref<Set<number>>(new Set())
const pendingAction = ref<'batch' | 'clear' | null>(null)
const canPrev = () => page.value > 1 const canPrev = () => page.value > 1
const canNext = () => items.value.length === pageSize const canNext = () => items.value.length === pageSize
const allSelected = computed(
() => items.value.length > 0 && items.value.every((item) => selectedIds.value.has(item.id)),
)
const someSelected = computed(() => selectedIds.value.size > 0 && !allSelected.value)
const selectedCount = computed(() => selectedIds.value.size)
const modalTitle = computed(() =>
pendingAction.value === 'clear' ? '清空丢弃数据' : '批量删除丢弃数据',
)
const modalMessage = computed(() =>
pendingAction.value === 'clear'
? '确定要清空全部丢弃数据吗?此操作不可撤销。'
: `确定要删除选中的 ${selectedCount.value} 条丢弃数据吗?此操作不可撤销。`,
)
const modalConfirmText = computed(() =>
pendingAction.value === 'clear' ? '清空全部' : '删除',
)
function toggleAll(checked: boolean) {
if (checked) {
items.value.forEach((item) => selectedIds.value.add(item.id))
} else {
selectedIds.value.clear()
}
selectedIds.value = new Set(selectedIds.value)
}
function toggleRow(id: number, checked: boolean) {
if (checked) {
selectedIds.value.add(id)
} else {
selectedIds.value.delete(id)
}
selectedIds.value = new Set(selectedIds.value)
}
function formatTime(value: string): string { function formatTime(value: string): string {
return new Date(value).toLocaleString() return new Date(value).toLocaleString()
} }
@@ -31,9 +71,43 @@ async function refreshItems() {
function changePage(nextPage: number) { function changePage(nextPage: number) {
page.value = Math.max(1, nextPage) page.value = Math.max(1, nextPage)
selectedIds.value.clear()
selectedIds.value = new Set(selectedIds.value)
refreshItems() refreshItems()
} }
function requestBatchDelete() {
if (selectedCount.value === 0) return
pendingAction.value = 'batch'
}
function requestClearAll() {
pendingAction.value = 'clear'
}
function cancelDeleteModal() {
pendingAction.value = null
}
async function confirmDeleteModal() {
const action = pendingAction.value
pendingAction.value = null
if (!action) return
try {
if (action === 'clear') {
await clearDiscardDetails()
} else {
await deleteDiscardDetailsByIDs([...selectedIds.value])
}
selectedIds.value.clear()
selectedIds.value = new Set(selectedIds.value)
page.value = 1
await refreshItems()
} catch (err) {
error.value = err instanceof Error ? err.message : String(err)
}
}
onMounted(refreshItems) onMounted(refreshItems)
</script> </script>
@@ -45,14 +119,35 @@ onMounted(refreshItems)
<p class="eyebrow">Discard details</p> <p class="eyebrow">Discard details</p>
<h2>丢弃数据</h2> <h2>丢弃数据</h2>
</div> </div>
<div style="display: flex; gap: 8px;">
<button
class="admin-button admin-button-danger"
:disabled="loading || selectedCount === 0"
@click="requestBatchDelete"
>批量删除{{ selectedCount > 0 ? `${selectedCount}` : '' }}</button>
<button
class="admin-button admin-button-danger"
:disabled="loading"
@click="requestClearAll"
>清空全部</button>
<button class="admin-button" @click="refreshItems" :disabled="loading">{{ loading ? '刷新中...' : '刷新数据' }}</button> <button class="admin-button" @click="refreshItems" :disabled="loading">{{ loading ? '刷新中...' : '刷新数据' }}</button>
</div> </div>
</div>
<p v-if="error" class="error">{{ error }}</p> <p v-if="error" class="error">{{ error }}</p>
<div class="node-table-wrap"> <div class="node-table-wrap">
<table class="node-table"> <table class="node-table">
<thead> <thead>
<tr> <tr>
<th style="width: 40px;">
<input
type="checkbox"
:checked="allSelected"
:indeterminate.prop="someSelected"
:disabled="items.length === 0"
@change="toggleAll(($event.target as HTMLInputElement).checked)"
/>
</th>
<th>时间</th> <th>时间</th>
<th>Topic</th> <th>Topic</th>
<th>Error</th> <th>Error</th>
@@ -67,6 +162,13 @@ onMounted(refreshItems)
</thead> </thead>
<tbody> <tbody>
<tr v-for="item in items" :key="item.id"> <tr v-for="item in items" :key="item.id">
<td>
<input
type="checkbox"
:checked="selectedIds.has(item.id)"
@change="toggleRow(item.id, ($event.target as HTMLInputElement).checked)"
/>
</td>
<td>{{ formatTime(item.created_at) }}</td> <td>{{ formatTime(item.created_at) }}</td>
<td>{{ item.topic || '-' }}</td> <td>{{ item.topic || '-' }}</td>
<td>{{ item.error || '-' }}</td> <td>{{ item.error || '-' }}</td>
@@ -90,5 +192,14 @@ onMounted(refreshItems)
<button :disabled="loading || !canNext()" @click="changePage(page + 1)">下一页</button> <button :disabled="loading || !canNext()" @click="changePage(page + 1)">下一页</button>
</div> </div>
</div> </div>
<ConfirmDeleteModal
:open="pendingAction !== null"
:title="modalTitle"
:message="modalMessage"
:confirm-text="modalConfirmText"
@cancel="cancelDeleteModal"
@confirm="confirmDeleteModal"
/>
</section> </section>
</template> </template>
+115 -8
View File
@@ -11,12 +11,19 @@ import {
updateLLMToolRouter, updateLLMToolRouter,
updateLLMTopicConfig, updateLLMTopicConfig,
updateLLMPrimaryConfig, updateLLMPrimaryConfig,
getAIServiceStatus,
restartAIService,
} from '../api' } from '../api'
import type { LLMPlatformRouter, LLMTopicConfig, LLMProvider, LLMPrimaryConfig } from '../types' import type { LLMPlatformRouter, LLMTopicConfig, LLMProvider, LLMPrimaryConfig, AIServiceStatus } from '../types'
const loading = ref(false) const loading = ref(false)
const error = ref('') const error = ref('')
const success = ref('') const success = ref('')
const showWarning = ref(false)
// AI Service Status
const aiStatus = ref<AIServiceStatus>({ running: false, enabled: false, provider_count: 0 })
const restarting = ref(false)
// LLM Provider 相关 // LLM Provider 相关
const providers = ref<LLMProvider[]>([]) const providers = ref<LLMProvider[]>([])
@@ -169,12 +176,18 @@ async function saveProvider() {
} }
try { try {
let response: any
if (isCreatingProvider.value) { if (isCreatingProvider.value) {
await createLLMProvider(providerForm.value) response = await createLLMProvider(providerForm.value)
success.value = '创建成功'
} else if (editingProvider.value) { } else if (editingProvider.value) {
await updateLLMProvider(editingProvider.value.name, providerForm.value) response = await updateLLMProvider(editingProvider.value.name, providerForm.value)
success.value = '更新成功' }
if (response.warning) {
success.value = response.warning
showWarning.value = true
} else {
success.value = isCreatingProvider.value ? '创建成功' : '更新成功'
showWarning.value = false
} }
clearSuccess() clearSuccess()
closeProviderForm() closeProviderForm()
@@ -189,8 +202,14 @@ async function confirmDeleteProvider(name: string) {
return return
} }
try { try {
await deleteLLMProvider(name) const response = await deleteLLMProvider(name)
if (response.warning) {
success.value = response.warning
showWarning.value = true
} else {
success.value = '删除成功' success.value = '删除成功'
showWarning.value = false
}
clearSuccess() clearSuccess()
await loadProviders() await loadProviders()
} catch (err) { } catch (err) {
@@ -286,7 +305,31 @@ async function savePrimaryConfig() {
} }
} }
async function loadAIStatus() {
try {
aiStatus.value = await getAIServiceStatus()
} catch (err) {
console.warn('Failed to load AI service status', err)
}
}
async function handleRestartAI() {
if (!confirm('确定要重启 AI 服务吗?')) return
restarting.value = true
try {
const result = await restartAIService()
aiStatus.value = result
success.value = result.message || 'AI 服务已重启'
clearSuccess()
} catch (err) {
error.value = err instanceof Error ? err.message : String(err)
} finally {
restarting.value = false
}
}
onMounted(() => { onMounted(() => {
loadAIStatus()
loadProviders() loadProviders()
loadToolRouter() loadToolRouter()
loadTopicConfig() loadTopicConfig()
@@ -298,8 +341,23 @@ onMounted(() => {
<div class="admin-llm-api"> <div class="admin-llm-api">
<h2>LLM API 配置管理</h2> <h2>LLM API 配置管理</h2>
<div class="ai-status-bar">
<span class="status-indicator" :class="{ running: aiStatus.running, stopped: !aiStatus.running }"></span>
<span class="status-text">
<template v-if="aiStatus.running">AI 服务运行中{{ aiStatus.provider_count }} 个提供商</template>
<template v-else>{{ aiStatus.message || 'AI 服务未运行' }}</template>
</span>
<button
class="admin-button admin-button-small"
:disabled="restarting"
@click="handleRestartAI"
>
{{ restarting ? '重启中...' : '重启 AI 服务' }}
</button>
</div>
<p v-if="error" class="error">{{ error }}</p> <p v-if="error" class="error">{{ error }}</p>
<p v-if="success" class="success">{{ success }}</p> <p v-if="success" :class="showWarning ? 'warning' : 'success'">{{ success }}</p>
<!-- LLM Provider 列表 --> <!-- LLM Provider 列表 -->
<div class="admin-section"> <div class="admin-section">
@@ -690,13 +748,49 @@ onMounted(() => {
} }
.admin-llm-api h2 { .admin-llm-api h2 {
margin: 0 0 2rem; margin: 0 0 1rem;
font-size: 1.75rem; font-size: 1.75rem;
font-weight: 700; font-weight: 700;
color: #1e293b; color: #1e293b;
letter-spacing: -0.02em; letter-spacing: -0.02em;
} }
.ai-status-bar {
display: flex;
align-items: center;
gap: 0.75rem;
padding: 0.75rem 1.25rem;
background: white;
border-radius: 12px;
border: 1px solid #e2e8f0;
margin-bottom: 1.5rem;
box-shadow: 0 1px 3px rgba(0, 0, 0, 0.05);
}
.status-indicator {
width: 10px;
height: 10px;
border-radius: 50%;
flex-shrink: 0;
}
.status-indicator.running {
background: #22c55e;
box-shadow: 0 0 6px rgba(34, 197, 94, 0.5);
}
.status-indicator.stopped {
background: #ef4444;
box-shadow: 0 0 6px rgba(239, 68, 68, 0.4);
}
.status-text {
flex: 1;
font-size: 0.9rem;
color: #475569;
font-weight: 500;
}
.admin-section { .admin-section {
background: white; background: white;
padding: 1.75rem; padding: 1.75rem;
@@ -973,6 +1067,19 @@ onMounted(() => {
gap: 0.5rem; gap: 0.5rem;
} }
.warning {
color: #92400e;
padding: 1rem 1.25rem;
background: linear-gradient(135deg, #fef3c7 0%, #fde68a 100%);
border-radius: 10px;
margin-bottom: 1.25rem;
border: 1px solid #fcd34d;
font-weight: 500;
display: flex;
align-items: center;
gap: 0.5rem;
}
.admin-loading { .admin-loading {
padding: 3rem; padding: 3rem;
text-align: center; text-align: center;
@@ -0,0 +1,35 @@
<script setup lang="ts">
import { onMounted, ref } from 'vue'
import { getBackendVersion } from '../api'
import { FRONTEND_VERSION } from '../version'
const backendVersion = ref('-')
const commitVersion = ref('-')
onMounted(async () => {
try {
const res = await getBackendVersion()
backendVersion.value = res.version
commitVersion.value = res.commit
} catch {
backendVersion.value = '-'
commitVersion.value = '-'
}
})
</script>
<template>
<footer class="app-footer">
<span>前端 v{{ FRONTEND_VERSION }}</span>
<span class="app-footer-sep">·</span>
<span>后端 v{{ backendVersion }}</span>
<span class="app-footer-sep">·</span>
<span>提交 {{ commitVersion }}</span>
<span class="app-footer-sep">·</span>
<a class="app-footer-github" href="https://github.com/wuwenfengmi1998/meshtastic_mqtt_server" target="_blank" rel="noopener noreferrer" aria-label="GitHub 仓库" title="GitHub 仓库">
<svg width="14" height="14" viewBox="0 0 16 16" fill="currentColor" aria-hidden="true">
<path d="M8 0C3.58 0 0 3.58 0 8c0 3.54 2.29 6.53 5.47 7.59.4.07.55-.17.55-.38 0-.19-.01-.82-.01-1.49-2.01.37-2.53-.49-2.69-.94-.09-.23-.48-.94-.82-1.13-.28-.15-.68-.52-.01-.53.63-.01 1.08.58 1.23.82.72 1.21 1.87.87 2.33.66.07-.52.28-.87.51-1.07-1.78-.2-3.64-.89-3.64-3.95 0-.87.31-1.59.82-2.15-.08-.2-.36-1.02.08-2.12 0 0 .67-.21 2.2.82.64-.18 1.32-.27 2-.27.68 0 1.36.09 2 .27 1.53-1.04 2.2-.82 2.2-.82.44 1.1.16 1.92.08 2.12.51.56.82 1.27.82 2.15 0 3.07-1.87 3.75-3.65 3.95.29.25.54.73.54 1.48 0 1.07-.01 1.93-.01 2.2 0 .21.15.46.55.38A8.013 8.013 0 0016 8c0-4.42-3.58-8-8-8z"/>
</svg>
</a>
</footer>
</template>
@@ -9,6 +9,8 @@ const props = defineProps<{
loadingOlder: boolean loadingOlder: boolean
hasMoreMessages: boolean hasMoreMessages: boolean
isAdmin: boolean isAdmin: boolean
channelFilter: string
channels: string[]
}>() }>()
type GroupedTextMessage = TextMessage & { mergedCount: number; mergedMessages: TextMessage[] } type GroupedTextMessage = TextMessage & { mergedCount: number; mergedMessages: TextMessage[] }
@@ -18,6 +20,7 @@ const emit = defineEmits<{
'load-older': [] 'load-older': []
'delete-message': [message: GroupedTextMessage] 'delete-message': [message: GroupedTextMessage]
'delete-and-block-node': [payload: { nodeId: string; nodeNum: number | null; message: GroupedTextMessage }] 'delete-and-block-node': [payload: { nodeId: string; nodeNum: number | null; message: GroupedTextMessage }]
'update:channelFilter': [value: string]
}>() }>()
const panelRef = ref<HTMLElement | null>(null) const panelRef = ref<HTMLElement | null>(null)
@@ -28,6 +31,56 @@ const topThreshold = 8
const bottomThreshold = 40 const bottomThreshold = 40
const scrollOverflowAllowance = 1 const scrollOverflowAllowance = 1
const showSuggestions = ref(false)
const selectedSuggestion = ref(-1)
const suggestions = computed(() => {
const q = props.channelFilter.trim().toLowerCase()
if (!q) return []
return props.channels
.filter(ch => ch.toLowerCase().includes(q))
.slice(0, 10)
})
function onFilterInput(event: Event) {
const value = (event.target as HTMLInputElement).value
emit('update:channelFilter', value)
selectedSuggestion.value = -1
showSuggestions.value = value.trim().length > 0 && suggestions.value.length > 0
}
function onFilterFocus() {
if (props.channelFilter.trim() && suggestions.value.length > 0) {
showSuggestions.value = true
}
}
function onFilterBlur() {
setTimeout(() => { showSuggestions.value = false }, 150)
}
function onFilterKeydown(event: KeyboardEvent) {
if (!showSuggestions.value || suggestions.value.length === 0) return
if (event.key === 'ArrowDown') {
event.preventDefault()
selectedSuggestion.value = Math.min(selectedSuggestion.value + 1, suggestions.value.length - 1)
} else if (event.key === 'ArrowUp') {
event.preventDefault()
selectedSuggestion.value = Math.max(selectedSuggestion.value - 1, -1)
} else if (event.key === 'Enter' && selectedSuggestion.value >= 0) {
event.preventDefault()
emit('update:channelFilter', suggestions.value[selectedSuggestion.value])
showSuggestions.value = false
} else if (event.key === 'Escape') {
showSuggestions.value = false
}
}
function selectSuggestion(ch: string) {
emit('update:channelFilter', ch)
showSuggestions.value = false
}
const groupedMessages = computed<GroupedTextMessage[]>(() => { const groupedMessages = computed<GroupedTextMessage[]>(() => {
const groups = new Map<string, GroupedTextMessage>() const groups = new Map<string, GroupedTextMessage>()
for (const message of props.messages) { for (const message of props.messages) {
@@ -184,6 +237,7 @@ onUpdated(() => {
<template> <template>
<aside ref="panelRef" class="chat-panel panel" @scroll.passive="handleScroll"> <aside ref="panelRef" class="chat-panel panel" @scroll.passive="handleScroll">
<div class="chat-panel-sticky">
<div class="panel-header"> <div class="panel-header">
<div> <div>
<p class="eyebrow">Chat</p> <p class="eyebrow">Chat</p>
@@ -192,6 +246,34 @@ onUpdated(() => {
<span class="badge">{{ groupedMessages.length }}</span> <span class="badge">{{ groupedMessages.length }}</span>
</div> </div>
<div class="chat-filter">
<input
type="search"
class="chat-filter-input"
:value="channelFilter"
placeholder="筛选频道"
@input="onFilterInput"
@focus="onFilterFocus"
@blur="onFilterBlur"
@keydown="onFilterKeydown"
/>
<button
v-if="channelFilter.trim()"
type="button"
class="chat-filter-clear"
@click="emit('update:channelFilter', '')"
>清除</button>
<ul v-if="showSuggestions && suggestions.length > 0" class="chat-filter-suggestions">
<li
v-for="(ch, idx) in suggestions"
:key="ch"
:class="{ active: idx === selectedSuggestion }"
@mousedown.prevent="selectSuggestion(ch)"
>{{ ch }}</li>
</ul>
</div>
</div>
<div v-if="loadingOlder" class="chat-loading">正在加载更早消息...</div> <div v-if="loadingOlder" class="chat-loading">正在加载更早消息...</div>
<div v-else-if="!hasMoreMessages && messages.length > 0" class="chat-end">没有更多历史消息</div> <div v-else-if="!hasMoreMessages && messages.length > 0" class="chat-end">没有更多历史消息</div>
<div v-if="messages.length === 0" class="empty">暂无聊天消息</div> <div v-if="messages.length === 0" class="empty">暂无聊天消息</div>
@@ -27,11 +27,33 @@ const canNext = computed(() => props.page < totalPages.value)
const menuNode = ref<NodeInfo | null>(null) const menuNode = ref<NodeInfo | null>(null)
const menuX = ref(0) const menuX = ref(0)
const menuY = ref(0) const menuY = ref(0)
const copiedNodeId = ref<string | null>(null)
function formatTime(value: string): string { function formatTime(value: string): string {
return new Date(value).toLocaleString() return new Date(value).toLocaleString()
} }
function truncateKey(key: string): string {
if (key.length <= 12) return key
return key.slice(0, 8) + '…' + key.slice(-4)
}
async function copyPublicKey(node: NodeInfo, event: MouseEvent) {
event.stopPropagation()
if (!node.public_key) return
try {
await navigator.clipboard.writeText(node.public_key)
copiedNodeId.value = node.node_id
setTimeout(() => {
if (copiedNodeId.value === node.node_id) {
copiedNodeId.value = null
}
}, 2000)
} catch {
/* clipboard not available */
}
}
function closeNodeMenu() { function closeNodeMenu() {
menuNode.value = null menuNode.value = null
} }
@@ -95,7 +117,7 @@ onBeforeUnmount(() => {
<span class="badge"> {{ total }} </span> <span class="badge"> {{ total }} </span>
</div> </div>
<div class="node-table-wrap" @scroll="closeNodeMenu"> <div class="node-table-wrap node-list-wrap" @scroll="closeNodeMenu">
<table class="node-table"> <table class="node-table">
<thead> <thead>
<tr> <tr>
@@ -122,7 +144,18 @@ onBeforeUnmount(() => {
<td>{{ node.short_name || '-' }}</td> <td>{{ node.short_name || '-' }}</td>
<td>{{ node.hw_model || '-' }}</td> <td>{{ node.hw_model || '-' }}</td>
<td>{{ node.role || '-' }}</td> <td>{{ node.role || '-' }}</td>
<td>{{ node.public_key || '-' }}</td> <td>
<span v-if="node.public_key" class="pubkey-cell">
<code class="pubkey-text">{{ truncateKey(node.public_key) }}</code>
<button
class="copy-btn"
type="button"
:title="copiedNodeId === node.node_id ? '已复制' : '复制完整公钥'"
@click.stop="copyPublicKey(node, $event)"
>{{ copiedNodeId === node.node_id ? '已复制' : '复制' }}</button>
</span>
<span v-else>-</span>
</td>
<td>{{ formatTime(node.updated_at) }}</td> <td>{{ formatTime(node.updated_at) }}</td>
</tr> </tr>
</tbody> </tbody>
+142 -1
View File
@@ -75,6 +75,31 @@ a {
padding: 16px; padding: 16px;
} }
.app-footer {
display: flex;
justify-content: center;
align-items: center;
gap: 8px;
padding: 8px 0;
font-size: 12px;
color: var(--color-muted);
}
.app-footer-sep {
opacity: 0.5;
}
.app-footer-github {
display: inline-flex;
align-items: center;
color: var(--color-muted);
transition: color 0.15s ease;
}
.app-footer-github:hover {
color: var(--color-heading);
}
.topbar { .topbar {
display: flex; display: flex;
align-items: center; align-items: center;
@@ -274,7 +299,7 @@ h3 {
overflow-y: auto; overflow-y: auto;
} }
.chat-panel .panel-header { .chat-panel-sticky {
position: sticky; position: sticky;
z-index: 10; z-index: 10;
top: 0; top: 0;
@@ -333,6 +358,81 @@ h3 {
background: var(--color-primary-soft); background: var(--color-primary-soft);
} }
.chat-filter {
position: relative;
display: flex;
align-items: center;
gap: 6px;
padding: 8px 12px;
border-bottom: 1px solid var(--color-border);
background: var(--color-surface-soft);
}
.chat-filter-input {
flex: 1;
min-width: 0;
border: 1px solid var(--color-border-strong);
border-radius: var(--radius-sm);
padding: 6px 10px;
color: var(--color-heading);
background: var(--color-surface);
outline: none;
font-size: 13px;
transition: border-color 0.16s ease, box-shadow 0.16s ease;
}
.chat-filter-input:focus {
border-color: var(--color-primary);
box-shadow: 0 0 0 3px color-mix(in srgb, var(--color-primary) 20%, transparent);
}
.chat-filter-clear {
flex-shrink: 0;
border: 1px solid var(--color-border-strong);
border-radius: var(--radius-sm);
padding: 5px 10px;
color: var(--color-heading);
background: var(--color-surface);
font-size: 12px;
font-weight: 700;
}
.chat-filter-clear:hover {
border-color: var(--color-primary);
color: var(--color-primary-hover);
background: var(--color-primary-soft);
}
.chat-filter-suggestions {
position: absolute;
top: 100%;
left: 0;
right: 0;
z-index: 20;
margin: 0;
padding: 0;
list-style: none;
background: var(--color-surface);
border: 1px solid var(--color-border-strong);
border-radius: var(--radius-sm);
box-shadow: 0 4px 12px rgba(0, 0, 0, 0.15);
max-height: 200px;
overflow-y: auto;
}
.chat-filter-suggestions li {
padding: 6px 10px;
font-size: 13px;
cursor: pointer;
color: var(--color-heading);
}
.chat-filter-suggestions li:hover,
.chat-filter-suggestions li.active {
background: var(--color-primary-soft);
color: var(--color-primary-hover);
}
.chat-item { .chat-item {
display: grid; display: grid;
gap: 6px; gap: 6px;
@@ -1113,6 +1213,47 @@ h3 {
background: var(--color-primary-soft); background: var(--color-primary-soft);
} }
/* === 节点列表:去除横向滚动,内容自适应换行 === */
.node-list-wrap {
overflow-x: hidden;
}
.node-list-wrap .node-table th,
.node-list-wrap .node-table td {
white-space: normal;
word-break: break-all;
vertical-align: middle;
}
.pubkey-cell {
display: inline-flex;
align-items: center;
gap: 6px;
}
.pubkey-text {
font-family: var(--font-mono);
font-size: 12px;
}
.copy-btn {
flex-shrink: 0;
padding: 2px 8px;
font-size: 12px;
border: 1px solid var(--color-border);
border-radius: var(--radius-sm);
background: var(--color-surface);
color: var(--color-muted);
cursor: pointer;
transition: all 0.15s ease;
}
.copy-btn:hover {
background: var(--color-primary-soft);
color: var(--color-primary);
border-color: var(--color-primary);
}
.admin-loading { .admin-loading {
padding: 24px; padding: 24px;
color: var(--color-muted); color: var(--color-muted);
+9
View File
@@ -364,6 +364,7 @@ export interface AdminMqttStatus {
messages_sent: number messages_sent: number
messages_dropped: number messages_dropped: number
db_write_queue_length: number db_write_queue_length: number
dedup_queue_len: number
retained: number retained: number
inflight: number inflight: number
inflight_dropped: number inflight_dropped: number
@@ -643,6 +644,14 @@ export interface LLMProviderPayload {
export interface LLMProviderResponse { export interface LLMProviderResponse {
item: LLMProvider item: LLMProvider
warning?: string
}
export interface AIServiceStatus {
running: boolean
enabled: boolean
provider_count: number
message?: string
} }
// LLM Tool Router 相关类型 // LLM Tool Router 相关类型
+1
View File
@@ -0,0 +1 @@
export const FRONTEND_VERSION = '0.3.1'
+1
View File
@@ -1 +1,2 @@
__pycache__ __pycache__
db_config.py
+266
View File
@@ -0,0 +1,266 @@
#!/usr/bin/env python3
"""
MySQL 8 -> MySQL 5.7 数据库迁移脚本
依赖: pip install pymysql
使用方法:
python py/migrate_mysql8_to_mysql57.py
功能:
1. 从源 MySQL 8 读取表结构,修正 collation 后在目标 MySQL 5.7 建表
2. 逐表批量拷贝数据
3. 修正 auto-increment 值
注意:
- 目标库应为空库(无同名表)
- 运行前请确保源库和目标库都可连接
"""
from __future__ import annotations
import re
import sys
from typing import Any
import pymysql
import pymysql.cursors
from db_config import SOURCE_CONFIG, TARGET_CONFIG
BATCH_SIZE = 5000
# ============================================================
# 工具函数
# ============================================================
def _conn(config: dict[str, Any]) -> pymysql.Connection:
"""创建数据库连接,使用 DictCursor 方便按列名访问"""
return pymysql.connect(
host=config["host"],
port=config["port"],
user=config["user"],
password=config["password"],
database=config["database"],
charset=config["charset"],
cursorclass=pymysql.cursors.DictCursor,
)
def fix_collation(ddl: str) -> str:
"""将 MySQL 8 默认 collation 替换为 MySQL 5.7 兼容版本"""
ddl = ddl.replace("utf8mb4_0900_ai_ci", "utf8mb4_general_ci")
ddl = ddl.replace("utf8mb4_0900_as_ci", "utf8mb4_general_ci")
return ddl
def fix_text_blob_index(ddl: str) -> str:
"""为 TEXT/BLOB 列在索引中添加前缀长度 (255),兼容 MySQL 5.7"""
text_cols = set()
for m in re.finditer(r'`(\w+)`\s+(?:tinytext|text|mediumtext|longtext|tinyblob|blob|mediumblob|longblob)\b', ddl, re.IGNORECASE):
text_cols.add(m.group(1))
if not text_cols:
return ddl
def _add_prefix(match: re.Match) -> str:
col = match.group(1)
rest = match.group(2)
if col in text_cols and not rest.strip().startswith('('):
return f'`{col}`(255){rest}'
return match.group(0)
# 匹配 KEY 定义中的列引用: `col_name` 后面不跟 ( 的情况
ddl = re.sub(r'`(\w+)`(\s*[,\)])', _add_prefix, ddl)
return ddl
def get_all_tables(conn: pymysql.Connection) -> list[str]:
with conn.cursor() as cur:
cur.execute("SHOW TABLES")
key = list(cur.description[0])[0]
return [row[key] for row in cur.fetchall()]
def get_table_columns(conn: pymysql.Connection, table: str) -> list[str]:
"""返回表的所有列名(按顺序)"""
with conn.cursor() as cur:
cur.execute("SHOW COLUMNS FROM `%s`" % table)
return [row["Field"] for row in cur.fetchall()]
def get_auto_increment(
conn: pymysql.Connection, table: str
) -> int | None:
"""获取某张表当前的 AUTO_INCREMENT 值"""
with conn.cursor() as cur:
cur.execute(
"SELECT AUTO_INCREMENT FROM information_schema.TABLES "
"WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = %s",
(table,),
)
row = cur.fetchone()
if row and row["AUTO_INCREMENT"]:
return int(row["AUTO_INCREMENT"])
return None
def row_count(conn: pymysql.Connection, table: str) -> int:
with conn.cursor() as cur:
cur.execute("SELECT COUNT(*) AS cnt FROM `%s`" % table)
return cur.fetchone()["cnt"]
# ============================================================
# 主流程
# ============================================================
def migrate() -> int:
src = _conn(SOURCE_CONFIG)
tgt = _conn(TARGET_CONFIG)
print(f"[连接] 源 MySQL 8 @ {SOURCE_CONFIG['host']}:{SOURCE_CONFIG['port']}")
print(f"[连接] 目标 MySQL 5.7 @ {TARGET_CONFIG['host']}:{TARGET_CONFIG['port']}")
print()
# 1. 在目标库建库(如还不存在)
try:
with tgt.cursor() as cur:
cur.execute(
"CREATE DATABASE IF NOT EXISTS `%s` "
"CHARACTER SET utf8mb4 COLLATE utf8mb4_general_ci"
% TARGET_CONFIG["database"]
)
cur.execute("USE `%s`" % TARGET_CONFIG["database"])
tgt.select_db(TARGET_CONFIG["database"])
except Exception as e:
print(f"[错误] 创建目标数据库失败: {e}")
return 1
tables = get_all_tables(src)
print(f"[发现] 源库共 {len(tables)} 张表: {', '.join(tables)}")
print()
# 2. 逐表建表(修正 collation
print("=" * 60)
print("阶段 1: 在目标库创建表结构")
print("=" * 60)
for idx, table in enumerate(tables, 1):
with src.cursor() as cur:
cur.execute("SHOW CREATE TABLE `%s`" % table)
row = cur.fetchone()
ddl = row["Create Table"]
ddl = fix_collation(ddl)
ddl = fix_text_blob_index(ddl)
try:
with tgt.cursor() as cur:
cur.execute(ddl)
tgt.commit()
print(f" [{idx:2d}/{len(tables)}] OK `{table}`")
except Exception as e:
print(f" [{idx:2d}/{len(tables)}] 错误 `{table}`: {e}")
tgt.rollback()
return 1
print()
# 3. 逐表拷贝数据
print("=" * 60)
print("阶段 2: 拷贝表数据")
print("=" * 60)
total_rows_copied = 0
auto_increments: dict[str, int | None] = {}
for idx, table in enumerate(tables, 1):
columns = get_table_columns(src, table)
if not columns:
auto_increments[table] = None
print(f" [{idx:2d}/{len(tables)}] SKIP `{table}` (0 列)")
continue
col_quoted = ", ".join("`%s`" % c for c in columns)
placeholders = ", ".join(["%s"] * len(columns))
insert_sql = "INSERT INTO `%s` (%s) VALUES (%s)" % (
table,
col_quoted,
placeholders,
)
table_count = 0
with src.cursor() as read_cur:
read_cur.execute("SELECT * FROM `%s`" % table)
batch = read_cur.fetchmany(BATCH_SIZE)
while batch:
rows_values = [
[row.get(c) for c in columns] for row in batch
]
try:
with tgt.cursor() as write_cur:
write_cur.executemany(insert_sql, rows_values)
tgt.commit()
table_count += len(rows_values)
print(
f"\r [{idx:2d}/{len(tables)}] `{table}` -> {table_count}",
end="",
flush=True,
)
except Exception as e:
print()
print(f" [{idx:2d}/{len(tables)}] 错误 `{table}`: {e}")
tgt.rollback()
return 1
batch = read_cur.fetchmany(BATCH_SIZE)
# 修正 auto_increment
ai = get_auto_increment(src, table)
auto_increments[table] = ai
if ai is not None:
with tgt.cursor() as cur:
cur.execute(
"ALTER TABLE `%s` AUTO_INCREMENT = %s" % (table, ai)
)
tgt.commit()
total_rows_copied += table_count
print(
f"\r [{idx:2d}/{len(tables)}] `{table}` -> {table_count} 行 [OK]"
)
# 4. 验证
print()
print("=" * 60)
print("阶段 3: 验证行数")
print("=" * 60)
all_match = True
src_total = 0
tgt_total = 0
for table in tables:
s = row_count(src, table)
t = row_count(tgt, table)
src_total += s
tgt_total += t
status = "OK" if s == t else "不匹配!"
if s != t:
all_match = False
print(f" `{table}`: 源={s} 目标={t} [{status}]")
print()
print(f" 总计: 源={src_total} 目标={tgt_total}")
src.close()
tgt.close()
if all_match:
print()
print("迁移完成,所有表行数一致。")
return 0
else:
print()
print("警告: 部分表行数不一致,请检查。")
return 1
if __name__ == "__main__":
sys.exit(migrate())
-156
View File
@@ -1,156 +0,0 @@
package main
import (
"net"
"testing"
"time"
mqtt "github.com/mochi-mqtt/server/v2"
"github.com/mochi-mqtt/server/v2/packets"
)
// TestTCPNoDelay 测试 TCP_NODELAY 是否正确设置
func TestTCPNoDelay(t *testing.T) {
// 创建一个模拟的 TCP 连接
listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatalf("Failed to create listener: %v", err)
}
defer listener.Close()
addr := listener.Addr().String()
// 模拟客户端连接
connChan := make(chan net.Conn, 1)
go func() {
conn, err := listener.Accept()
if err != nil {
t.Errorf("Failed to accept connection: %v", err)
return
}
connChan <- conn
}()
// 客户端连接
clientConn, err := net.Dial("tcp", addr)
if err != nil {
t.Fatalf("Failed to dial: %v", err)
}
defer clientConn.Close()
// 等待服务器端接受连接
serverConn := <-connChan
defer serverConn.Close()
// 创建 MQTT Client 包装
cl := &mqtt.Client{
Net: mqtt.ClientConnection{
Conn: serverConn,
Remote: serverConn.RemoteAddr().String(),
},
}
// 创建 hook 并调用 OnConnect
hook := &meshtasticFilterHook{}
pk := packets.Packet{}
err = hook.OnConnect(cl, pk)
if err != nil {
t.Fatalf("OnConnect failed: %v", err)
}
// 验证 TCP_NODELAY 是否设置
if tcpConn, ok := serverConn.(*net.TCPConn); ok {
// 这里我们无法直接读取 TCP_NODELAY 的值,但可以验证没有错误
// 实际上,我们可以通过设置后再次设置来验证
err := tcpConn.SetNoDelay(false)
if err != nil {
t.Fatalf("Failed to set NoDelay to false: %v", err)
}
err = tcpConn.SetNoDelay(true)
if err != nil {
t.Fatalf("Failed to set NoDelay to true: %v", err)
}
t.Log("TCP_NODELAY successfully set")
} else {
t.Fatal("Connection is not a TCP connection")
}
}
// TestQoS0MessageLatency 测试 QoS0 消息的响应延迟
func TestQoS0MessageLatency(t *testing.T) {
// 创建一个简单的 TCP echo 服务器来模拟 MQTT
listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatalf("Failed to create listener: %v", err)
}
defer listener.Close()
addr := listener.Addr().String()
// 启动服务器
go func() {
conn, err := listener.Accept()
if err != nil {
return
}
defer conn.Close()
// 设置 TCP_NODELAY
if tcpConn, ok := conn.(*net.TCPConn); ok {
tcpConn.SetNoDelay(true)
}
buf := make([]byte, 1024)
for {
n, err := conn.Read(buf)
if err != nil {
return
}
// 立即写回(模拟 ACK
_, err = conn.Write(buf[:n])
if err != nil {
return
}
}
}()
// 客户端连接
conn, err := net.Dial("tcp", addr)
if err != nil {
t.Fatalf("Failed to dial: %v", err)
}
defer conn.Close()
// 测试小数据包的延迟
testData := []byte("test")
samples := 10
var totalLatency time.Duration
for i := 0; i < samples; i++ {
start := time.Now()
_, err := conn.Write(testData)
if err != nil {
t.Fatalf("Write failed: %v", err)
}
buf := make([]byte, len(testData))
_, err = conn.Read(buf)
if err != nil {
t.Fatalf("Read failed: %v", err)
}
latency := time.Since(start)
totalLatency += latency
t.Logf("Round trip %d: %v", i+1, latency)
}
avgLatency := totalLatency / time.Duration(samples)
t.Logf("Average latency: %v", avgLatency)
// 平均延迟应该小于 10ms(如果没有 Nagle 算法延迟)
if avgLatency > 10*time.Millisecond {
t.Logf("Warning: Average latency %v is higher than expected, may indicate Nagle's algorithm is active", avgLatency)
}
}