Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
dd5a331db2 | ||
|
|
7582e8b55e |
+1
-5
@@ -35,16 +35,12 @@ func runServe(ctx context.Context, cfg config.Config) error {
|
||||
return fmt.Errorf("mkdir data_dir: %w", err)
|
||||
}
|
||||
|
||||
db, err := store.OpenWriter(cfg.DataDir, cfg.SQLiteSynchronous)
|
||||
db, err := store.Open(cfg.DataDir, cfg.SQLiteSynchronous)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() { _ = db.Close() }()
|
||||
|
||||
if migErr := store.Migrate(db); migErr != nil {
|
||||
return migErr
|
||||
}
|
||||
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc("GET /healthz", func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
|
||||
@@ -73,6 +73,53 @@
|
||||
- 备选方案:协议包只做编解码,校验留给各 app 模块。
|
||||
- 影响:服务端应复用本包 `Validate`,避免重复规则。
|
||||
|
||||
### T0.3 2026-09-30
|
||||
|
||||
1. **完整表放在 0002,不改已发布的 0001**
|
||||
- 原条款:TASKS T0.3 / 4.2「`0001_init.sql` 包含第 7.7 节全部表」;T0.1 偏差曾写「T0.3 需替换 0001 正文」。
|
||||
- 实际做法:保留 `0001_init.sql` 为 `SELECT 1;`;新增 `0002_schema.sql` 写入 DEVELOPMENT 7.7 全部业务表与索引(含 `api_tokens`、`settings`、`session_hash` 等)。`schema_migrations` 仍由迁移执行器 `CREATE TABLE IF NOT EXISTS` 维护,不放入 0002。
|
||||
- 原因:T0.1 的 0001 可能已记入已有库的 `schema_migrations`;改写已发布迁移语义会导致「版本已应用但表不存在」。
|
||||
- 备选方案:对未迁移库特殊检测并改写 0001(复杂且易错)。
|
||||
- 影响:新库会有版本 1+2 两行;与 TASKS「表在 0001」字面不一致,与「不改已发布迁移」一致。
|
||||
|
||||
2. **写入队列先做一操作一事务**
|
||||
- 原条款:DEVELOPMENT 7.2 合并提交(最多 256 或凑满 2ms,SAVEPOINT);TASKS T0.3 允许简单实现,P2 换合并。
|
||||
- 实际做法:`store.Queue` 用互斥锁串行,每请求一个事务;注释与本条标明 P2 再改为写 goroutine 合并提交。
|
||||
- 原因:本任务范围;合并留给平台 P2。
|
||||
- 备选方案:T0.3 直接做合并(抢 P2 范围)。
|
||||
- 影响:高并发写入落盘次数偏多,正式压测前需完成 P2。
|
||||
|
||||
3. **空库不备份;仅已有 db 文件且有未应用版本时 VACUUM INTO**
|
||||
- 原条款:DEVELOPMENT 7.7「有未应用版本时先 VACUUM INTO」;未区分空库。
|
||||
- 实际做法:`Open` 在打开前检查 `nixmsg.db` 是否已存在;不存在则跳过备份;存在且有 pending 则写入 `<data_dir>/backup/pre-migrate-<UTC时间>.db`。迁移失败返回错误,不自动从备份恢复。
|
||||
- 原因:空库备份无意义;失败退出与文档一致,恢复交给运维。
|
||||
- 备选方案:失败时自动还原备份再退出。
|
||||
- 影响:与任务说明一致;运维需知备份路径。
|
||||
|
||||
### T0.5 2026-09-30
|
||||
|
||||
1. **admin init / 管理员密码**
|
||||
- 原条款:TASKS T0.5「自动执行 `admin init` 拿到管理员密码」;命令行启动器 JSON 含管理员密码。
|
||||
- 实际做法:定义 `AdminInitializer`(默认 `CLIAdminInit`);二进制尚无 `admin` 子命令时返回 `ErrAdminInitUnsupported`,启动器跳过并继续起 serve;`admin_password` 字段为空字符串。示例集成测试只验 `/healthz`。
|
||||
- 原因:`admin init` 属 P1,当前 main 仅有 `version`/`serve`。
|
||||
- 备选方案:harness 内嵌假密码写入库(无表结构可写);或阻塞等 P1。
|
||||
- 影响:P1 合入后无需改调用方接口,密码解析约定见 `parseAdminPassword`;Q/SDK 集成测试在拿到非空密码前勿依赖管理登录。
|
||||
|
||||
2. **MQTT 测试客户端范围**
|
||||
- 原条款:能收发 DEVELOPMENT 第 6 节应用帧。
|
||||
- 实际做法:提供 TCP / WebSocket(`/mqtt`,子协议 `mqtt`)传输层 `MQTTClient`,收发原始 MQTT 控制包字节;不实现 CONNECT/hello/主题业务。
|
||||
- 原因:内置 broker 与协议处理尚未合入(连接 N / 后续任务);T0.5 先给可连传输与占位 API。
|
||||
- 备选方案:引入完整 MQTT 客户端库并编假 broker。
|
||||
- 影响:业务级帧测试在 broker 可用后由各线基于 `Send`/`Recv` 或再包一层完成。
|
||||
|
||||
3. **进程停止方式**
|
||||
- 原条款:优雅停机(DEVELOPMENT 7.8 / P1)。
|
||||
- 实际做法:测试启动器对子进程使用 `Kill`(Windows 上 `Interrupt` 不可靠)。
|
||||
- 原因:保证并行测试与清理在 Windows 上稳定。
|
||||
- 备选方案:Unix 发 SIGTERM;Windows 用 Job Object / Ctrl+Break。
|
||||
- 影响:不覆盖优雅停机验收;该验收仍归 P1/Q。
|
||||
|
||||
|
||||
## 平台 P
|
||||
|
||||
暂无。
|
||||
|
||||
+103
-9
@@ -10,29 +10,105 @@ import (
|
||||
_ "modernc.org/sqlite"
|
||||
)
|
||||
|
||||
// OpenWriter 按 DEVELOPMENT 7.7 打开写连接(带 _txlock=immediate)。
|
||||
const dbFileName = "nixmsg.db"
|
||||
|
||||
// DB 持有读写连接与写入队列。
|
||||
type DB struct {
|
||||
Write *sql.DB
|
||||
Read *sql.DB
|
||||
Queue *Queue
|
||||
}
|
||||
|
||||
// Open 打开数据目录下的库:写连接、读连接池、迁移,并启动简单写入队列。
|
||||
// 若 nixmsg.db 尚不存在则跳过迁移前备份。
|
||||
func Open(dataDir, synchronous string) (*DB, error) {
|
||||
if err := os.MkdirAll(dataDir, 0o755); err != nil {
|
||||
return nil, fmt.Errorf("mkdir data_dir: %w", err)
|
||||
}
|
||||
dbPath := filepath.Join(dataDir, dbFileName)
|
||||
existed, err := fileExists(dbPath)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
write, err := OpenWriter(dataDir, synchronous)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if migErr := Migrate(write, dataDir, existed); migErr != nil {
|
||||
_ = write.Close()
|
||||
return nil, migErr
|
||||
}
|
||||
read, err := OpenReader(dataDir, synchronous)
|
||||
if err != nil {
|
||||
_ = write.Close()
|
||||
return nil, err
|
||||
}
|
||||
q := NewQueue(write)
|
||||
return &DB{Write: write, Read: read, Queue: q}, nil
|
||||
}
|
||||
|
||||
// Close 停止写入队列并关闭连接。
|
||||
func (d *DB) Close() error {
|
||||
var first error
|
||||
if d.Queue != nil {
|
||||
if err := d.Queue.Close(); err != nil {
|
||||
first = err
|
||||
}
|
||||
}
|
||||
if d.Read != nil {
|
||||
if err := d.Read.Close(); err != nil && first == nil {
|
||||
first = err
|
||||
}
|
||||
}
|
||||
if d.Write != nil {
|
||||
if err := d.Write.Close(); err != nil && first == nil {
|
||||
first = err
|
||||
}
|
||||
}
|
||||
return first
|
||||
}
|
||||
|
||||
// OpenWriter 按 DEVELOPMENT 7.7 打开写连接(带 _txlock=immediate,MaxOpenConns=1)。
|
||||
func OpenWriter(dataDir, synchronous string) (*sql.DB, error) {
|
||||
if err := os.MkdirAll(dataDir, 0o755); err != nil {
|
||||
return nil, fmt.Errorf("mkdir data_dir: %w", err)
|
||||
}
|
||||
dsn, err := writeDSN(dataDir, synchronous)
|
||||
dsn, err := buildDSN(dataDir, synchronous, true)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
db, err := sql.Open("sqlite", dsn)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("open sqlite: %w", err)
|
||||
return nil, fmt.Errorf("open sqlite writer: %w", err)
|
||||
}
|
||||
db.SetMaxOpenConns(1)
|
||||
db.SetMaxIdleConns(1)
|
||||
if err := db.Ping(); err != nil {
|
||||
_ = db.Close()
|
||||
return nil, fmt.Errorf("ping sqlite: %w", err)
|
||||
return nil, fmt.Errorf("ping sqlite writer: %w", err)
|
||||
}
|
||||
return db, nil
|
||||
}
|
||||
|
||||
func writeDSN(dataDir, synchronous string) (string, error) {
|
||||
// OpenReader 打开读连接池(DSN 不含 _txlock=immediate)。
|
||||
func OpenReader(dataDir, synchronous string) (*sql.DB, error) {
|
||||
dsn, err := buildDSN(dataDir, synchronous, false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
db, err := sql.Open("sqlite", dsn)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("open sqlite reader: %w", err)
|
||||
}
|
||||
if err := db.Ping(); err != nil {
|
||||
_ = db.Close()
|
||||
return nil, fmt.Errorf("ping sqlite reader: %w", err)
|
||||
}
|
||||
return db, nil
|
||||
}
|
||||
|
||||
func buildDSN(dataDir, synchronous string, writer bool) (string, error) {
|
||||
sync := strings.ToUpper(strings.TrimSpace(synchronous))
|
||||
if sync == "" {
|
||||
sync = "FULL"
|
||||
@@ -40,10 +116,28 @@ func writeDSN(dataDir, synchronous string) (string, error) {
|
||||
if sync != "FULL" && sync != "NORMAL" {
|
||||
return "", fmt.Errorf("invalid sqlite_synchronous: %s", synchronous)
|
||||
}
|
||||
dbPath := filepath.ToSlash(filepath.Join(dataDir, "nixmsg.db"))
|
||||
return fmt.Sprintf(
|
||||
"file:%s?_pragma=journal_mode(WAL)&_pragma=busy_timeout(5000)&_pragma=synchronous(%s)&_pragma=foreign_keys(ON)&_pragma=secure_delete(ON)&_txlock=immediate",
|
||||
dbPath := filepath.ToSlash(filepath.Join(dataDir, dbFileName))
|
||||
dsn := fmt.Sprintf(
|
||||
"file:%s?_pragma=journal_mode(WAL)&_pragma=busy_timeout(5000)&_pragma=synchronous(%s)&_pragma=foreign_keys(ON)&_pragma=secure_delete(ON)",
|
||||
dbPath,
|
||||
sync,
|
||||
), nil
|
||||
)
|
||||
if writer {
|
||||
dsn += "&_txlock=immediate"
|
||||
}
|
||||
return dsn, nil
|
||||
}
|
||||
|
||||
func fileExists(path string) (bool, error) {
|
||||
st, err := os.Stat(path)
|
||||
if err == nil {
|
||||
if st.IsDir() {
|
||||
return false, fmt.Errorf("db path is a directory: %s", path)
|
||||
}
|
||||
return true, nil
|
||||
}
|
||||
if os.IsNotExist(err) {
|
||||
return false, nil
|
||||
}
|
||||
return false, err
|
||||
}
|
||||
|
||||
+144
-15
@@ -1,36 +1,165 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestOpenAndMigrate(t *testing.T) {
|
||||
func TestOpenEmptyDirCreatesTables(t *testing.T) {
|
||||
t.Parallel()
|
||||
dir := t.TempDir()
|
||||
db, err := OpenWriter(dir, "FULL")
|
||||
|
||||
db, err := Open(dir, "FULL")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer func() { _ = db.Close() }()
|
||||
|
||||
if migErr := Migrate(db); migErr != nil {
|
||||
t.Fatal(migErr)
|
||||
assertMigrationCount(t, db.Write, 2)
|
||||
for _, table := range []string{
|
||||
"endpoints", "settings", "talk_grants", "groups", "group_members",
|
||||
"messages", "message_bodies", "deliveries", "receipts", "send_keys",
|
||||
"admin_sessions", "api_tokens", "schema_migrations",
|
||||
} {
|
||||
assertTableExists(t, db.Write, table)
|
||||
}
|
||||
if migErr := Migrate(db); migErr != nil {
|
||||
t.Fatal(migErr)
|
||||
// 空目录首次启动不应产生迁移备份。
|
||||
if entries, readErr := os.ReadDir(filepath.Join(dir, "backup")); readErr == nil && len(entries) > 0 {
|
||||
t.Fatalf("unexpected backup files on empty start: %d", len(entries))
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenIdempotentNoRemigrate(t *testing.T) {
|
||||
t.Parallel()
|
||||
dir := t.TempDir()
|
||||
|
||||
db1, err := Open(dir, "FULL")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
assertMigrationCount(t, db1.Write, 2)
|
||||
_ = db1.Close()
|
||||
|
||||
db2, err := Open(dir, "FULL")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer func() { _ = db2.Close() }()
|
||||
assertMigrationCount(t, db2.Write, 2)
|
||||
|
||||
if entries, readErr := os.ReadDir(filepath.Join(dir, "backup")); readErr == nil && len(entries) > 0 {
|
||||
t.Fatalf("idempotent reopen should not backup: %v", entries)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMigrateBackupWhenNewVersion(t *testing.T) {
|
||||
t.Parallel()
|
||||
dir := t.TempDir()
|
||||
|
||||
// 模拟仅应用了 0001 的旧库。
|
||||
w, err := OpenWriter(dir, "FULL")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, execErr := w.Exec(`
|
||||
CREATE TABLE schema_migrations (
|
||||
version INTEGER PRIMARY KEY,
|
||||
applied_at INTEGER NOT NULL
|
||||
)`); execErr != nil {
|
||||
t.Fatal(execErr)
|
||||
}
|
||||
if _, execErr := w.Exec(`INSERT INTO schema_migrations(version, applied_at) VALUES(1, ?)`, time.Now().UnixMilli()); execErr != nil {
|
||||
t.Fatal(execErr)
|
||||
}
|
||||
_ = w.Close()
|
||||
|
||||
db, err := Open(dir, "FULL")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer func() { _ = db.Close() }()
|
||||
|
||||
assertMigrationCount(t, db.Write, 2)
|
||||
assertTableExists(t, db.Write, "endpoints")
|
||||
assertTableExists(t, db.Write, "api_tokens")
|
||||
|
||||
backupDir := filepath.Join(dir, "backup")
|
||||
entries, err := os.ReadDir(backupDir)
|
||||
if err != nil {
|
||||
t.Fatalf("backup dir missing: %v", err)
|
||||
}
|
||||
found := false
|
||||
for _, e := range entries {
|
||||
if strings.HasPrefix(e.Name(), "pre-migrate-") && strings.HasSuffix(e.Name(), ".db") {
|
||||
found = true
|
||||
info, statErr := e.Info()
|
||||
if statErr != nil {
|
||||
t.Fatal(statErr)
|
||||
}
|
||||
if info.Size() == 0 {
|
||||
t.Fatalf("backup file empty: %s", e.Name())
|
||||
}
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Fatal("expected pre-migrate-*.db backup")
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteQueueOneOpOneTx(t *testing.T) {
|
||||
t.Parallel()
|
||||
dir := t.TempDir()
|
||||
db, err := Open(dir, "FULL")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer func() { _ = db.Close() }()
|
||||
|
||||
ctx := context.Background()
|
||||
err = db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
_, execErr := tx.Exec(
|
||||
`INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?)`,
|
||||
"registration_enabled", "0", time.Now().UnixMilli(),
|
||||
)
|
||||
return execErr
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
var n int
|
||||
if scanErr := db.QueryRow(`SELECT COUNT(*) FROM schema_migrations`).Scan(&n); scanErr != nil {
|
||||
var value string
|
||||
if scanErr := db.Read.QueryRow(`SELECT value FROM settings WHERE key = ?`, "registration_enabled").Scan(&value); scanErr != nil {
|
||||
t.Fatal(scanErr)
|
||||
}
|
||||
if n != 1 {
|
||||
t.Fatalf("want 1 migration row, got %d", n)
|
||||
if value != "0" {
|
||||
t.Fatalf("want 0, got %q", value)
|
||||
}
|
||||
}
|
||||
|
||||
func assertMigrationCount(t *testing.T, db *sql.DB, want int) {
|
||||
t.Helper()
|
||||
var n int
|
||||
if err := db.QueryRow(`SELECT COUNT(*) FROM schema_migrations`).Scan(&n); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if n != want {
|
||||
t.Fatalf("migration rows: want %d got %d", want, n)
|
||||
}
|
||||
}
|
||||
|
||||
func assertTableExists(t *testing.T, db *sql.DB, name string) {
|
||||
t.Helper()
|
||||
var got string
|
||||
err := db.QueryRow(
|
||||
`SELECT name FROM sqlite_master WHERE type='table' AND name=?`,
|
||||
name,
|
||||
).Scan(&got)
|
||||
if err != nil {
|
||||
t.Fatalf("table %s missing: %v", name, err)
|
||||
}
|
||||
nested, openErr := OpenWriter(filepath.Join(dir, "nested"), "NORMAL")
|
||||
if openErr != nil {
|
||||
t.Fatal(openErr)
|
||||
}
|
||||
_ = nested.Close()
|
||||
}
|
||||
|
||||
+84
-27
@@ -5,6 +5,8 @@ import (
|
||||
"embed"
|
||||
"fmt"
|
||||
"io/fs"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
@@ -14,8 +16,11 @@ import (
|
||||
//go:embed migrations/*.sql
|
||||
var migrationFS embed.FS
|
||||
|
||||
// Migrate 应用尚未执行的迁移。T0.1 仅为空执行器加 0001 占位;完整备份与表结构见 T0.3。
|
||||
func Migrate(db *sql.DB) error {
|
||||
// Migrate 应用尚未执行的嵌入迁移。
|
||||
// 若 dbExisted 为 true(调用 Open/Migrate 前已有 nixmsg.db)且存在未应用版本,
|
||||
// 先 VACUUM INTO <data_dir>/backup/pre-migrate-<时间>.db,再迁移。
|
||||
// 任一版本失败则返回错误,调用方不得继续带半新半旧库提供服务。
|
||||
func Migrate(db *sql.DB, dataDir string, dbExisted bool) error {
|
||||
if _, err := db.Exec(`
|
||||
CREATE TABLE IF NOT EXISTS schema_migrations (
|
||||
version INTEGER PRIMARY KEY,
|
||||
@@ -24,9 +29,38 @@ CREATE TABLE IF NOT EXISTS schema_migrations (
|
||||
return fmt.Errorf("ensure schema_migrations: %w", err)
|
||||
}
|
||||
|
||||
pending, err := pendingMigrations(db)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if len(pending) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
if dbExisted {
|
||||
if err := backupBeforeMigrate(db, dataDir); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
for _, m := range pending {
|
||||
if err := applyMigration(db, m); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
type migrationFile struct {
|
||||
version int
|
||||
name string
|
||||
body string
|
||||
}
|
||||
|
||||
func pendingMigrations(db *sql.DB) ([]migrationFile, error) {
|
||||
entries, err := fs.ReadDir(migrationFS, "migrations")
|
||||
if err != nil {
|
||||
return fmt.Errorf("read migrations: %w", err)
|
||||
return nil, fmt.Errorf("read migrations: %w", err)
|
||||
}
|
||||
var names []string
|
||||
for _, e := range entries {
|
||||
@@ -37,10 +71,11 @@ CREATE TABLE IF NOT EXISTS schema_migrations (
|
||||
}
|
||||
sort.Strings(names)
|
||||
|
||||
var pending []migrationFile
|
||||
for _, name := range names {
|
||||
version, err := parseMigrationVersion(name)
|
||||
if err != nil {
|
||||
return err
|
||||
return nil, err
|
||||
}
|
||||
var exists int
|
||||
err = db.QueryRow(`SELECT 1 FROM schema_migrations WHERE version = ?`, version).Scan(&exists)
|
||||
@@ -48,36 +83,58 @@ CREATE TABLE IF NOT EXISTS schema_migrations (
|
||||
continue
|
||||
}
|
||||
if err != sql.ErrNoRows {
|
||||
return fmt.Errorf("check migration %d: %w", version, err)
|
||||
return nil, fmt.Errorf("check migration %d: %w", version, err)
|
||||
}
|
||||
|
||||
body, err := migrationFS.ReadFile("migrations/" + name)
|
||||
if err != nil {
|
||||
return fmt.Errorf("read migration %s: %w", name, err)
|
||||
return nil, fmt.Errorf("read migration %s: %w", name, err)
|
||||
}
|
||||
sqlText := strings.TrimSpace(string(body))
|
||||
tx, err := db.Begin()
|
||||
if err != nil {
|
||||
return fmt.Errorf("begin migration %d: %w", version, err)
|
||||
}
|
||||
if sqlText != "" {
|
||||
if _, err := tx.Exec(sqlText); err != nil {
|
||||
_ = tx.Rollback()
|
||||
return fmt.Errorf("apply migration %d: %w", version, err)
|
||||
}
|
||||
}
|
||||
if _, err := tx.Exec(
|
||||
`INSERT INTO schema_migrations(version, applied_at) VALUES(?, ?)`,
|
||||
version,
|
||||
time.Now().UnixMilli(),
|
||||
); err != nil {
|
||||
pending = append(pending, migrationFile{
|
||||
version: version,
|
||||
name: name,
|
||||
body: strings.TrimSpace(string(body)),
|
||||
})
|
||||
}
|
||||
return pending, nil
|
||||
}
|
||||
|
||||
func backupBeforeMigrate(db *sql.DB, dataDir string) error {
|
||||
backupDir := filepath.Join(dataDir, "backup")
|
||||
if err := os.MkdirAll(backupDir, 0o755); err != nil {
|
||||
return fmt.Errorf("mkdir backup: %w", err)
|
||||
}
|
||||
stamp := time.Now().UTC().Format("20060102T150405")
|
||||
backupPath := filepath.Join(backupDir, "pre-migrate-"+stamp+".db")
|
||||
// SQLite VACUUM INTO 需要字面量路径;统一用斜杠,并对单引号转义。
|
||||
quoted := strings.ReplaceAll(filepath.ToSlash(backupPath), "'", "''")
|
||||
if _, err := db.Exec("VACUUM INTO '" + quoted + "'"); err != nil {
|
||||
return fmt.Errorf("vacuum into backup %s: %w", backupPath, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func applyMigration(db *sql.DB, m migrationFile) error {
|
||||
tx, err := db.Begin()
|
||||
if err != nil {
|
||||
return fmt.Errorf("begin migration %d: %w", m.version, err)
|
||||
}
|
||||
if m.body != "" {
|
||||
if _, err := tx.Exec(m.body); err != nil {
|
||||
_ = tx.Rollback()
|
||||
return fmt.Errorf("record migration %d: %w", version, err)
|
||||
}
|
||||
if err := tx.Commit(); err != nil {
|
||||
return fmt.Errorf("commit migration %d: %w", version, err)
|
||||
return fmt.Errorf("apply migration %d (%s): %w", m.version, m.name, err)
|
||||
}
|
||||
}
|
||||
if _, err := tx.Exec(
|
||||
`INSERT INTO schema_migrations(version, applied_at) VALUES(?, ?)`,
|
||||
m.version,
|
||||
time.Now().UnixMilli(),
|
||||
); err != nil {
|
||||
_ = tx.Rollback()
|
||||
return fmt.Errorf("record migration %d: %w", m.version, err)
|
||||
}
|
||||
if err := tx.Commit(); err != nil {
|
||||
return fmt.Errorf("commit migration %d: %w", m.version, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,131 @@
|
||||
-- 完整业务表(T0.1 的 0001 仅为 SELECT 1 占位且可能已记入 schema_migrations,故放在 0002)。
|
||||
CREATE TABLE endpoints (
|
||||
id TEXT PRIMARY KEY,
|
||||
name TEXT NOT NULL DEFAULT '',
|
||||
remark TEXT NOT NULL DEFAULT '',
|
||||
source TEXT NOT NULL DEFAULT 'admin',
|
||||
login_hash TEXT NOT NULL,
|
||||
talk_hash TEXT,
|
||||
talk_version INTEGER NOT NULL DEFAULT 0,
|
||||
default_delay_ms INTEGER NOT NULL DEFAULT 0,
|
||||
enabled INTEGER NOT NULL DEFAULT 1,
|
||||
created_at INTEGER NOT NULL,
|
||||
online_since INTEGER,
|
||||
offline_since INTEGER,
|
||||
session_hash TEXT,
|
||||
session_issued_at INTEGER,
|
||||
session_used_at INTEGER
|
||||
);
|
||||
|
||||
CREATE TABLE settings (
|
||||
key TEXT PRIMARY KEY,
|
||||
value TEXT NOT NULL,
|
||||
updated_at INTEGER NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE talk_grants (
|
||||
sender_id TEXT NOT NULL,
|
||||
target_id TEXT NOT NULL,
|
||||
target_talk_version INTEGER NOT NULL,
|
||||
kind TEXT NOT NULL,
|
||||
created_at INTEGER NOT NULL,
|
||||
PRIMARY KEY (sender_id, target_id)
|
||||
);
|
||||
|
||||
CREATE TABLE groups (
|
||||
id TEXT PRIMARY KEY,
|
||||
name TEXT NOT NULL,
|
||||
owner_id TEXT NOT NULL,
|
||||
created_at INTEGER NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE group_members (
|
||||
group_id TEXT NOT NULL,
|
||||
endpoint_id TEXT NOT NULL,
|
||||
joined_at INTEGER NOT NULL,
|
||||
PRIMARY KEY (group_id, endpoint_id)
|
||||
);
|
||||
|
||||
CREATE TABLE messages (
|
||||
seq INTEGER PRIMARY KEY,
|
||||
id TEXT NOT NULL,
|
||||
sender_id TEXT NOT NULL,
|
||||
dest_kind TEXT NOT NULL,
|
||||
dest_id TEXT NOT NULL,
|
||||
meta TEXT NOT NULL DEFAULT '{}',
|
||||
content_type TEXT NOT NULL,
|
||||
body_enc TEXT NOT NULL,
|
||||
send_at INTEGER NOT NULL,
|
||||
keep INTEGER NOT NULL,
|
||||
ttl_seconds INTEGER NOT NULL DEFAULT 0,
|
||||
receipt INTEGER NOT NULL,
|
||||
state TEXT NOT NULL,
|
||||
reason TEXT NOT NULL DEFAULT '',
|
||||
created_at INTEGER NOT NULL,
|
||||
UNIQUE (sender_id, id)
|
||||
);
|
||||
|
||||
CREATE TABLE message_bodies (
|
||||
seq INTEGER PRIMARY KEY REFERENCES messages(seq) ON DELETE CASCADE,
|
||||
body BLOB NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE deliveries (
|
||||
seq INTEGER NOT NULL REFERENCES messages(seq) ON DELETE CASCADE,
|
||||
endpoint_id TEXT NOT NULL,
|
||||
send_at INTEGER NOT NULL,
|
||||
keep INTEGER NOT NULL,
|
||||
state TEXT NOT NULL,
|
||||
reason TEXT NOT NULL DEFAULT '',
|
||||
expire_at INTEGER,
|
||||
pushed_conn TEXT,
|
||||
pushed_at INTEGER,
|
||||
attempts INTEGER NOT NULL DEFAULT 0,
|
||||
updated_at INTEGER NOT NULL,
|
||||
PRIMARY KEY (seq, endpoint_id)
|
||||
);
|
||||
|
||||
CREATE TABLE receipts (
|
||||
receipt_id INTEGER PRIMARY KEY,
|
||||
sender_id TEXT NOT NULL,
|
||||
msg_id TEXT NOT NULL,
|
||||
endpoint_id TEXT NOT NULL DEFAULT '',
|
||||
state TEXT NOT NULL,
|
||||
reason TEXT NOT NULL DEFAULT '',
|
||||
created_at INTEGER NOT NULL,
|
||||
acked INTEGER NOT NULL DEFAULT 0
|
||||
);
|
||||
|
||||
CREATE TABLE send_keys (
|
||||
sender_id TEXT NOT NULL,
|
||||
msg_id TEXT NOT NULL,
|
||||
request_sha256 BLOB NOT NULL,
|
||||
created_at INTEGER NOT NULL,
|
||||
PRIMARY KEY (sender_id, msg_id)
|
||||
);
|
||||
|
||||
CREATE TABLE admin_sessions (
|
||||
token_hash TEXT PRIMARY KEY,
|
||||
created_at INTEGER NOT NULL,
|
||||
expires_at INTEGER NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE api_tokens (
|
||||
id TEXT PRIMARY KEY,
|
||||
name TEXT NOT NULL,
|
||||
token_hash TEXT NOT NULL UNIQUE,
|
||||
enabled INTEGER NOT NULL DEFAULT 1,
|
||||
created_at INTEGER NOT NULL,
|
||||
last_used_at INTEGER
|
||||
);
|
||||
|
||||
CREATE INDEX idx_messages_due ON messages(state, send_at);
|
||||
CREATE INDEX idx_messages_dest ON messages(dest_kind, dest_id, state);
|
||||
CREATE INDEX idx_messages_created ON messages(created_at);
|
||||
CREATE INDEX idx_messages_sender ON messages(sender_id, created_at);
|
||||
CREATE INDEX idx_messages_sender_state ON messages(sender_id, state);
|
||||
CREATE INDEX idx_deliveries_outbox ON deliveries(endpoint_id, state, send_at, seq);
|
||||
CREATE INDEX idx_deliveries_expire ON deliveries(state, expire_at);
|
||||
CREATE INDEX idx_receipts_outbox ON receipts(sender_id, acked, receipt_id);
|
||||
CREATE INDEX idx_talk_grants_target ON talk_grants(target_id);
|
||||
CREATE INDEX idx_group_members_endpoint ON group_members(endpoint_id);
|
||||
@@ -0,0 +1,62 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// ErrQueueClosed 表示写入队列已关闭。
|
||||
var ErrQueueClosed = errors.New("store: write queue closed")
|
||||
|
||||
// WriteFunc 在单个写事务中执行的操作。
|
||||
type WriteFunc func(tx *sql.Tx) error
|
||||
|
||||
// Queue 写入队列:提交一个写操作并拿到结果。
|
||||
//
|
||||
// 本任务(T0.3)实现为互斥串行的一操作一事务,不做合并;DEVELOPMENT 7.2
|
||||
// 要求的写 goroutine 合并提交(最多 256 个或凑满 2ms、SAVEPOINT 隔离失败)留给 P2。
|
||||
type Queue struct {
|
||||
db *sql.DB
|
||||
|
||||
mu sync.Mutex
|
||||
closed bool
|
||||
}
|
||||
|
||||
// NewQueue 创建简单写入队列(一操作一事务)。
|
||||
func NewQueue(db *sql.DB) *Queue {
|
||||
return &Queue{db: db}
|
||||
}
|
||||
|
||||
// Do 提交写操作并等待提交结果。
|
||||
func (q *Queue) Do(ctx context.Context, fn WriteFunc) error {
|
||||
if fn == nil {
|
||||
return errors.New("store: nil write func")
|
||||
}
|
||||
q.mu.Lock()
|
||||
defer q.mu.Unlock()
|
||||
if q.closed {
|
||||
return ErrQueueClosed
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
tx, err := q.db.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := fn(tx); err != nil {
|
||||
_ = tx.Rollback()
|
||||
return err
|
||||
}
|
||||
return tx.Commit()
|
||||
}
|
||||
|
||||
// Close 关闭队列,之后 Do 返回 ErrQueueClosed。
|
||||
func (q *Queue) Close() error {
|
||||
q.mu.Lock()
|
||||
defer q.mu.Unlock()
|
||||
q.closed = true
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,161 @@
|
||||
package harness
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/cookiejar"
|
||||
"net/url"
|
||||
"os"
|
||||
"os/exec"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// ErrAdminInitUnsupported 表示当前二进制还没有 admin init 子命令。
|
||||
var ErrAdminInitUnsupported = errors.New("admin init unsupported")
|
||||
|
||||
// AdminInitializer 在启动 serve 前初始化管理员密码。
|
||||
// P1 实现 admin init 后,默认实现会解析终端输出中的密码。
|
||||
type AdminInitializer interface {
|
||||
Init(binPath, configPath string) (password string, err error)
|
||||
}
|
||||
|
||||
// CLIAdminInit 调用 `nixmsg admin init`;若命令不存在则返回 ErrAdminInitUnsupported。
|
||||
type CLIAdminInit struct{}
|
||||
|
||||
// Init 执行 admin init。
|
||||
func (CLIAdminInit) Init(binPath, configPath string) (string, error) {
|
||||
cmd := exec.Command(binPath, "admin", "init")
|
||||
cmd.Env = append(os.Environ(), "NIXMSG_CONFIG="+configPath)
|
||||
out, err := cmd.CombinedOutput()
|
||||
text := string(out)
|
||||
if err != nil {
|
||||
if isAdminInitUnsupported(text) {
|
||||
return "", ErrAdminInitUnsupported
|
||||
}
|
||||
return "", fmt.Errorf("admin init: %w\n%s", err, text)
|
||||
}
|
||||
pass := parseAdminPassword(text)
|
||||
if pass == "" {
|
||||
return "", fmt.Errorf("admin init succeeded but password not found in output:\n%s", text)
|
||||
}
|
||||
return pass, nil
|
||||
}
|
||||
|
||||
func isAdminInitUnsupported(text string) bool {
|
||||
lower := strings.ToLower(text)
|
||||
return strings.Contains(lower, "unknown command: admin") ||
|
||||
strings.Contains(lower, "unknown command") && strings.Contains(lower, "admin")
|
||||
}
|
||||
|
||||
func parseAdminPassword(out string) string {
|
||||
// 约定:P1 实现后密码单独占一行,或出现在 "password:" 之后。
|
||||
lines := strings.Split(out, "\n")
|
||||
for _, line := range lines {
|
||||
line = strings.TrimSpace(line)
|
||||
if line == "" {
|
||||
continue
|
||||
}
|
||||
lower := strings.ToLower(line)
|
||||
switch {
|
||||
case strings.HasPrefix(lower, "password:"):
|
||||
return strings.TrimSpace(line[len("password:"):])
|
||||
case strings.HasPrefix(lower, "admin password:"):
|
||||
return strings.TrimSpace(line[len("admin password:"):])
|
||||
}
|
||||
}
|
||||
for i := len(lines) - 1; i >= 0; i-- {
|
||||
line := strings.TrimSpace(lines[i])
|
||||
if len(line) >= 12 && !strings.Contains(strings.ToLower(line), "error") {
|
||||
fields := strings.Fields(line)
|
||||
if len(fields) == 1 {
|
||||
return fields[0]
|
||||
}
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// AdminClient 管理接口 HTTP 客户端。
|
||||
// Cookie 会话模式下,改变状态的方法会自动加 X-Nixmsg-Request: 1。
|
||||
// 设置 APIToken 后走 Bearer,不再要求该头。
|
||||
type AdminClient struct {
|
||||
BaseURL string
|
||||
HTTP *http.Client
|
||||
APIToken string
|
||||
}
|
||||
|
||||
// NewAdminClient 创建带 CookieJar 的管理客户端。baseURL 形如 http://127.0.0.1:12345。
|
||||
func NewAdminClient(baseURL string) (*AdminClient, error) {
|
||||
jar, err := cookiejar.New(nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &AdminClient{
|
||||
BaseURL: strings.TrimRight(baseURL, "/"),
|
||||
HTTP: &http.Client{
|
||||
Timeout: 30 * time.Second,
|
||||
Jar: jar,
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
// SetSessionCookie 手动写入管理员会话 Cookie(测试辅助)。
|
||||
func (c *AdminClient) SetSessionCookie(value string) error {
|
||||
u, err := url.Parse(c.BaseURL)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
c.HTTP.Jar.SetCookies(u, []*http.Cookie{{
|
||||
Name: "nixmsg_admin",
|
||||
Value: value,
|
||||
Path: "/",
|
||||
}})
|
||||
return nil
|
||||
}
|
||||
|
||||
// Do 发送请求。method 为 POST/PUT/PATCH/DELETE 且未使用 API 令牌时,自动加 X-Nixmsg-Request: 1。
|
||||
func (c *AdminClient) Do(method, path string, body []byte, contentType string) (*http.Response, error) {
|
||||
if !strings.HasPrefix(path, "/") {
|
||||
path = "/" + path
|
||||
}
|
||||
var rdr io.Reader
|
||||
if body != nil {
|
||||
rdr = bytes.NewReader(body)
|
||||
}
|
||||
req, err := http.NewRequest(method, c.BaseURL+path, rdr)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if contentType != "" {
|
||||
req.Header.Set("Content-Type", contentType)
|
||||
}
|
||||
if c.APIToken != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+c.APIToken)
|
||||
} else if isMutatingMethod(method) {
|
||||
req.Header.Set("X-Nixmsg-Request", "1")
|
||||
}
|
||||
return c.HTTP.Do(req)
|
||||
}
|
||||
|
||||
// Get JSON GET。
|
||||
func (c *AdminClient) Get(path string) (*http.Response, error) {
|
||||
return c.Do(http.MethodGet, path, nil, "")
|
||||
}
|
||||
|
||||
// PostJSON POST application/json。
|
||||
func (c *AdminClient) PostJSON(path string, body []byte) (*http.Response, error) {
|
||||
return c.Do(http.MethodPost, path, body, "application/json")
|
||||
}
|
||||
|
||||
func isMutatingMethod(method string) bool {
|
||||
switch strings.ToUpper(method) {
|
||||
case http.MethodPost, http.MethodPut, http.MethodPatch, http.MethodDelete:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,55 @@
|
||||
package harness_test
|
||||
|
||||
import (
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/test/harness"
|
||||
)
|
||||
|
||||
func TestAdminClientMutatingHeader(t *testing.T) {
|
||||
t.Parallel()
|
||||
var sawRequestHdr, sawAuth bool
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Header.Get("X-Nixmsg-Request") == "1" {
|
||||
sawRequestHdr = true
|
||||
}
|
||||
if r.Header.Get("Authorization") != "" {
|
||||
sawAuth = true
|
||||
}
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write([]byte(`{}`))
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
c, err := harness.NewAdminClient(srv.URL)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
resp, err := c.PostJSON("/api/admin/login", []byte(`{"password":"x"}`))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, _ = io.Copy(io.Discard, resp.Body)
|
||||
_ = resp.Body.Close()
|
||||
if !sawRequestHdr {
|
||||
t.Fatal("expected X-Nixmsg-Request on POST")
|
||||
}
|
||||
|
||||
sawRequestHdr = false
|
||||
c.APIToken = "nxm_test"
|
||||
resp, err = c.PostJSON("/api/admin/endpoints", []byte(`{}`))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, _ = io.Copy(io.Discard, resp.Body)
|
||||
_ = resp.Body.Close()
|
||||
if sawRequestHdr {
|
||||
t.Fatal("API token mode should not set X-Nixmsg-Request")
|
||||
}
|
||||
if !sawAuth {
|
||||
t.Fatal("expected Authorization bearer")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,63 @@
|
||||
package harness
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"sync"
|
||||
)
|
||||
|
||||
var (
|
||||
binOnce sync.Once
|
||||
binPath string
|
||||
binErr error
|
||||
)
|
||||
|
||||
// Binary 返回已编译好的 nixmsg 可执行文件路径(整个进程只编译一次)。
|
||||
func Binary() (string, error) {
|
||||
binOnce.Do(func() {
|
||||
root, err := moduleRoot()
|
||||
if err != nil {
|
||||
binErr = err
|
||||
return
|
||||
}
|
||||
dir, err := os.MkdirTemp("", "nixmsg-harness-bin-*")
|
||||
if err != nil {
|
||||
binErr = err
|
||||
return
|
||||
}
|
||||
name := "nixmsg"
|
||||
if runtime.GOOS == "windows" {
|
||||
name += ".exe"
|
||||
}
|
||||
out := filepath.Join(dir, name)
|
||||
cmd := exec.Command("go", "build", "-o", out, "./cmd/nixmsg")
|
||||
cmd.Dir = root
|
||||
cmd.Env = append(os.Environ(), "CGO_ENABLED=0")
|
||||
if outBytes, runErr := cmd.CombinedOutput(); runErr != nil {
|
||||
binErr = fmt.Errorf("build nixmsg: %w\n%s", runErr, outBytes)
|
||||
return
|
||||
}
|
||||
binPath = out
|
||||
})
|
||||
return binPath, binErr
|
||||
}
|
||||
|
||||
func moduleRoot() (string, error) {
|
||||
dir, err := os.Getwd()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
for {
|
||||
if _, statErr := os.Stat(filepath.Join(dir, "go.mod")); statErr == nil {
|
||||
return dir, nil
|
||||
}
|
||||
parent := filepath.Dir(dir)
|
||||
if parent == dir {
|
||||
return "", fmt.Errorf("go.mod not found from %s", dir)
|
||||
}
|
||||
dir = parent
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"os/signal"
|
||||
"syscall"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/test/harness"
|
||||
)
|
||||
|
||||
func main() {
|
||||
srv, err := harness.Start(harness.Options{KeepDir: true})
|
||||
if err != nil {
|
||||
fmt.Fprintln(os.Stderr, err)
|
||||
os.Exit(1)
|
||||
}
|
||||
info := srv.ConnectionInfo()
|
||||
enc := json.NewEncoder(os.Stdout)
|
||||
enc.SetEscapeHTML(false)
|
||||
if err := enc.Encode(info); err != nil {
|
||||
_ = srv.Stop()
|
||||
_ = os.RemoveAll(srv.DataDir)
|
||||
fmt.Fprintln(os.Stderr, err)
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
sigCh := make(chan os.Signal, 1)
|
||||
signal.Notify(sigCh, os.Interrupt, syscall.SIGTERM)
|
||||
<-sigCh
|
||||
|
||||
_ = srv.Stop()
|
||||
_ = os.RemoveAll(srv.DataDir)
|
||||
}
|
||||
@@ -0,0 +1,116 @@
|
||||
package harness_test
|
||||
|
||||
import (
|
||||
"io"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/test/harness"
|
||||
)
|
||||
|
||||
func TestHealthzAndCleanup(t *testing.T) {
|
||||
srv, err := harness.Start(harness.Options{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
dataDir := srv.DataDir
|
||||
httpBase := srv.HTTPBase
|
||||
|
||||
resp, err := http.Get(httpBase + "/healthz")
|
||||
if err != nil {
|
||||
_ = srv.Stop()
|
||||
t.Fatalf("healthz: %v", err)
|
||||
}
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
_ = resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
_ = srv.Stop()
|
||||
t.Fatalf("status=%d body=%s", resp.StatusCode, body)
|
||||
}
|
||||
if string(body) != "ok" {
|
||||
_ = srv.Stop()
|
||||
t.Fatalf("body=%q", body)
|
||||
}
|
||||
|
||||
if err = srv.Stop(); err != nil {
|
||||
t.Fatalf("stop: %v", err)
|
||||
}
|
||||
|
||||
deadline := time.Now().Add(5 * time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
if _, statErr := os.Stat(dataDir); os.IsNotExist(statErr) {
|
||||
break
|
||||
}
|
||||
time.Sleep(20 * time.Millisecond)
|
||||
}
|
||||
if _, statErr := os.Stat(dataDir); !os.IsNotExist(statErr) {
|
||||
t.Fatalf("data dir still exists: %s", dataDir)
|
||||
}
|
||||
|
||||
_, err = http.Get(httpBase + "/healthz")
|
||||
if err == nil {
|
||||
t.Fatal("healthz still reachable after stop")
|
||||
}
|
||||
}
|
||||
|
||||
func TestParallelIsolation(t *testing.T) {
|
||||
var (
|
||||
wg sync.WaitGroup
|
||||
mu sync.Mutex
|
||||
addrs []string
|
||||
dataDirs []string
|
||||
)
|
||||
runOne := func(i int) {
|
||||
defer wg.Done()
|
||||
srv, err := harness.Start(harness.Options{})
|
||||
if err != nil {
|
||||
t.Errorf("start %d: %v", i, err)
|
||||
return
|
||||
}
|
||||
defer func() { _ = srv.Stop() }()
|
||||
|
||||
resp, err := http.Get(srv.HTTPBase + "/healthz")
|
||||
if err != nil {
|
||||
t.Errorf("healthz %d: %v", i, err)
|
||||
return
|
||||
}
|
||||
_ = resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
t.Errorf("healthz %d status=%d", i, resp.StatusCode)
|
||||
return
|
||||
}
|
||||
|
||||
mu.Lock()
|
||||
addrs = append(addrs, srv.Addr)
|
||||
dataDirs = append(dataDirs, srv.DataDir)
|
||||
mu.Unlock()
|
||||
}
|
||||
|
||||
const n = 2
|
||||
wg.Add(n)
|
||||
for i := 0; i < n; i++ {
|
||||
go runOne(i)
|
||||
}
|
||||
wg.Wait()
|
||||
if t.Failed() {
|
||||
return
|
||||
}
|
||||
if len(addrs) != n || len(dataDirs) != n {
|
||||
t.Fatalf("got %d addrs %d dirs", len(addrs), len(dataDirs))
|
||||
}
|
||||
if addrs[0] == addrs[1] {
|
||||
t.Fatalf("same listen addr: %s", addrs[0])
|
||||
}
|
||||
if dataDirs[0] == dataDirs[1] {
|
||||
t.Fatalf("same data dir: %s", dataDirs[0])
|
||||
}
|
||||
for _, d := range dataDirs {
|
||||
if filepath.Clean(d) == "" {
|
||||
t.Fatal("empty data dir")
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,34 @@
|
||||
// Package harness 提供集成测试启动器。
|
||||
package harness
|
||||
|
||||
// Info 是命令行启动器打印到 stdout 的连接信息。
|
||||
type Info struct {
|
||||
Listen string `json:"listen"`
|
||||
HTTPBase string `json:"http_base"`
|
||||
AdminListen string `json:"admin_listen,omitempty"`
|
||||
AdminHTTPBase string `json:"admin_http_base"`
|
||||
MQTTWS string `json:"mqtt_ws"`
|
||||
MQTTTCP string `json:"mqtt_tcp"`
|
||||
AdminPassword string `json:"admin_password"`
|
||||
DataDir string `json:"data_dir"`
|
||||
ConfigPath string `json:"config_path"`
|
||||
}
|
||||
|
||||
// ConnectionInfo 组装给外部语言测试用的 JSON 结构。
|
||||
func (s *Server) ConnectionInfo() Info {
|
||||
adminBase := s.AdminHTTPBase
|
||||
if adminBase == "" {
|
||||
adminBase = s.HTTPBase
|
||||
}
|
||||
return Info{
|
||||
Listen: s.Addr,
|
||||
HTTPBase: s.HTTPBase,
|
||||
AdminListen: s.AdminAddr,
|
||||
AdminHTTPBase: adminBase,
|
||||
MQTTWS: "ws://" + s.Addr + "/mqtt",
|
||||
MQTTTCP: s.Addr,
|
||||
AdminPassword: s.AdminPassword,
|
||||
DataDir: s.DataDir,
|
||||
ConfigPath: s.ConfigPath,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,321 @@
|
||||
package harness
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"crypto/rand"
|
||||
"crypto/sha1"
|
||||
"encoding/base64"
|
||||
"encoding/binary"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// MQTTClient 测试用 MQTT 传输层:能连上 WebSocket /mqtt 或裸 TCP,收发原始 MQTT 控制包字节。
|
||||
// 不实现业务握手(hello)与主题约定;完整帧协议留给各线集成测试自行组合。
|
||||
type MQTTClient interface {
|
||||
Send(packet []byte) error
|
||||
Recv() ([]byte, error)
|
||||
Close() error
|
||||
LocalAddr() net.Addr
|
||||
RemoteAddr() net.Addr
|
||||
}
|
||||
|
||||
// DialMQTTTCP 连接裸 MQTT TCP(与 listen 同一地址)。
|
||||
func DialMQTTTCP(addr string, timeout time.Duration) (MQTTClient, error) {
|
||||
if timeout <= 0 {
|
||||
timeout = 5 * time.Second
|
||||
}
|
||||
conn, err := net.DialTimeout("tcp", addr, timeout)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
_ = conn.SetDeadline(time.Now().Add(timeout))
|
||||
return &tcpMQTT{conn: conn, r: bufio.NewReader(conn)}, nil
|
||||
}
|
||||
|
||||
// DialMQTTWebSocket 连接 ws(s)://host/mqtt,子协议 mqtt。
|
||||
func DialMQTTWebSocket(httpBase string, timeout time.Duration) (MQTTClient, error) {
|
||||
if timeout <= 0 {
|
||||
timeout = 5 * time.Second
|
||||
}
|
||||
base := strings.TrimRight(httpBase, "/")
|
||||
u, err := url.Parse(base)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
switch u.Scheme {
|
||||
case "http":
|
||||
u.Scheme = "ws"
|
||||
case "https":
|
||||
u.Scheme = "wss"
|
||||
case "ws", "wss":
|
||||
case "":
|
||||
u.Scheme = "ws"
|
||||
default:
|
||||
return nil, fmt.Errorf("unsupported scheme %q", u.Scheme)
|
||||
}
|
||||
u.Path = "/mqtt"
|
||||
u.RawQuery = ""
|
||||
u.Fragment = ""
|
||||
|
||||
key := make([]byte, 16)
|
||||
if _, err = rand.Read(key); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
secKey := base64.StdEncoding.EncodeToString(key)
|
||||
|
||||
httpURL := *u
|
||||
if u.Scheme == "ws" {
|
||||
httpURL.Scheme = "http"
|
||||
} else {
|
||||
httpURL.Scheme = "https"
|
||||
}
|
||||
|
||||
req, err := http.NewRequest(http.MethodGet, httpURL.String(), nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
req.Header.Set("Connection", "Upgrade")
|
||||
req.Header.Set("Upgrade", "websocket")
|
||||
req.Header.Set("Sec-WebSocket-Version", "13")
|
||||
req.Header.Set("Sec-WebSocket-Key", secKey)
|
||||
req.Header.Set("Sec-WebSocket-Protocol", "mqtt")
|
||||
|
||||
host := u.Hostname()
|
||||
port := u.Port()
|
||||
if port == "" {
|
||||
if u.Scheme == "wss" {
|
||||
port = "443"
|
||||
} else {
|
||||
port = "80"
|
||||
}
|
||||
}
|
||||
dialer := net.Dialer{Timeout: timeout}
|
||||
raw, err := dialer.Dial("tcp", net.JoinHostPort(host, port))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
_ = raw.SetDeadline(time.Now().Add(timeout))
|
||||
|
||||
if err = req.Write(raw); err != nil {
|
||||
_ = raw.Close()
|
||||
return nil, err
|
||||
}
|
||||
br := bufio.NewReader(raw)
|
||||
resp, err := http.ReadResponse(br, req)
|
||||
if err != nil {
|
||||
_ = raw.Close()
|
||||
return nil, err
|
||||
}
|
||||
if resp.StatusCode != http.StatusSwitchingProtocols {
|
||||
body, _ := io.ReadAll(io.LimitReader(resp.Body, 512))
|
||||
_ = resp.Body.Close()
|
||||
_ = raw.Close()
|
||||
return nil, fmt.Errorf("websocket upgrade status %d: %s", resp.StatusCode, body)
|
||||
}
|
||||
if resp.Header.Get("Sec-WebSocket-Accept") != wsAcceptKey(secKey) {
|
||||
_ = raw.Close()
|
||||
return nil, fmt.Errorf("bad Sec-WebSocket-Accept")
|
||||
}
|
||||
if proto := resp.Header.Get("Sec-WebSocket-Protocol"); proto != "" && proto != "mqtt" {
|
||||
_ = raw.Close()
|
||||
return nil, fmt.Errorf("unexpected subprotocol %q", proto)
|
||||
}
|
||||
return &wsMQTT{conn: raw, r: br}, nil
|
||||
}
|
||||
|
||||
func wsAcceptKey(secKey string) string {
|
||||
const guid = "258EAFA5-E914-47DA-95CA-C5AB0DC85B11"
|
||||
sum := sha1.Sum([]byte(secKey + guid))
|
||||
return base64.StdEncoding.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
type tcpMQTT struct {
|
||||
conn net.Conn
|
||||
r *bufio.Reader
|
||||
}
|
||||
|
||||
func (c *tcpMQTT) Send(packet []byte) error {
|
||||
_, err := c.conn.Write(packet)
|
||||
return err
|
||||
}
|
||||
|
||||
func (c *tcpMQTT) Recv() ([]byte, error) {
|
||||
return readMQTTPacket(c.r)
|
||||
}
|
||||
|
||||
func (c *tcpMQTT) Close() error {
|
||||
return c.conn.Close()
|
||||
}
|
||||
func (c *tcpMQTT) LocalAddr() net.Addr {
|
||||
return c.conn.LocalAddr()
|
||||
}
|
||||
func (c *tcpMQTT) RemoteAddr() net.Addr {
|
||||
return c.conn.RemoteAddr()
|
||||
}
|
||||
|
||||
type wsMQTT struct {
|
||||
conn net.Conn
|
||||
r *bufio.Reader
|
||||
}
|
||||
|
||||
func (c *wsMQTT) Send(packet []byte) error {
|
||||
return writeWSClientBinary(c.conn, packet)
|
||||
}
|
||||
|
||||
func (c *wsMQTT) Recv() ([]byte, error) {
|
||||
for {
|
||||
payload, opcode, err := readWSFrame(c.r)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
switch opcode {
|
||||
case 0x2:
|
||||
return payload, nil
|
||||
case 0x8:
|
||||
return nil, io.EOF
|
||||
case 0x9:
|
||||
_ = writeWSClientControl(c.conn, 0xA, payload)
|
||||
case 0xA:
|
||||
continue
|
||||
default:
|
||||
continue
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (c *wsMQTT) Close() error {
|
||||
return c.conn.Close()
|
||||
}
|
||||
func (c *wsMQTT) LocalAddr() net.Addr {
|
||||
return c.conn.LocalAddr()
|
||||
}
|
||||
func (c *wsMQTT) RemoteAddr() net.Addr {
|
||||
return c.conn.RemoteAddr()
|
||||
}
|
||||
|
||||
func readMQTTPacket(r *bufio.Reader) ([]byte, error) {
|
||||
first, err := r.ReadByte()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
remaining, remBytes, err := readMQTTRemaining(r)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
buf := make([]byte, 1+len(remBytes)+remaining)
|
||||
buf[0] = first
|
||||
copy(buf[1:], remBytes)
|
||||
if remaining > 0 {
|
||||
if _, err := io.ReadFull(r, buf[1+len(remBytes):]); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return buf, nil
|
||||
}
|
||||
|
||||
func readMQTTRemaining(r *bufio.Reader) (value int, raw []byte, err error) {
|
||||
multiplier := 1
|
||||
for i := 0; i < 4; i++ {
|
||||
encoded, readErr := r.ReadByte()
|
||||
if readErr != nil {
|
||||
return 0, nil, readErr
|
||||
}
|
||||
raw = append(raw, encoded)
|
||||
value += int(encoded&127) * multiplier
|
||||
if encoded&128 == 0 {
|
||||
return value, raw, nil
|
||||
}
|
||||
multiplier *= 128
|
||||
}
|
||||
return 0, nil, fmt.Errorf("mqtt remaining length overflow")
|
||||
}
|
||||
|
||||
func writeWSClientBinary(w io.Writer, payload []byte) error {
|
||||
return writeWSClientFrame(w, 0x2, payload)
|
||||
}
|
||||
|
||||
func writeWSClientControl(w io.Writer, opcode byte, payload []byte) error {
|
||||
return writeWSClientFrame(w, opcode, payload)
|
||||
}
|
||||
|
||||
func writeWSClientFrame(w io.Writer, opcode byte, payload []byte) error {
|
||||
mask := make([]byte, 4)
|
||||
if _, err := rand.Read(mask); err != nil {
|
||||
return err
|
||||
}
|
||||
header := []byte{0x80 | (opcode & 0x0f)}
|
||||
n := len(payload)
|
||||
switch {
|
||||
case n < 126:
|
||||
header = append(header, 0x80|byte(n))
|
||||
case n <= 65535:
|
||||
header = append(header, 0x80|126, byte(n>>8), byte(n))
|
||||
default:
|
||||
var ext [8]byte
|
||||
binary.BigEndian.PutUint64(ext[:], uint64(n))
|
||||
header = append(header, 0x80|127)
|
||||
header = append(header, ext[:]...)
|
||||
}
|
||||
header = append(header, mask...)
|
||||
masked := make([]byte, n)
|
||||
for i := 0; i < n; i++ {
|
||||
masked[i] = payload[i] ^ mask[i%4]
|
||||
}
|
||||
if _, err := w.Write(header); err != nil {
|
||||
return err
|
||||
}
|
||||
_, err := w.Write(masked)
|
||||
return err
|
||||
}
|
||||
|
||||
func readWSFrame(r *bufio.Reader) (payload []byte, opcode byte, err error) {
|
||||
b0, err := r.ReadByte()
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
opcode = b0 & 0x0f
|
||||
b1, err := r.ReadByte()
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
masked := b1&0x80 != 0
|
||||
n := int(b1 & 0x7f)
|
||||
switch n {
|
||||
case 126:
|
||||
var ext [2]byte
|
||||
if _, err := io.ReadFull(r, ext[:]); err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
n = int(binary.BigEndian.Uint16(ext[:]))
|
||||
case 127:
|
||||
var ext [8]byte
|
||||
if _, err := io.ReadFull(r, ext[:]); err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
n = int(binary.BigEndian.Uint64(ext[:]))
|
||||
}
|
||||
var maskKey [4]byte
|
||||
if masked {
|
||||
if _, err := io.ReadFull(r, maskKey[:]); err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
}
|
||||
payload = make([]byte, n)
|
||||
if n > 0 {
|
||||
if _, err := io.ReadFull(r, payload); err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
if masked {
|
||||
for i := 0; i < n; i++ {
|
||||
payload[i] ^= maskKey[i%4]
|
||||
}
|
||||
}
|
||||
}
|
||||
return payload, opcode, nil
|
||||
}
|
||||
@@ -0,0 +1,162 @@
|
||||
package harness
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Options 控制测试进程启动。
|
||||
type Options struct {
|
||||
// Listen 写入配置的 listen,默认 127.0.0.1:0。
|
||||
Listen string
|
||||
// AdminListen 可选;非空时写入 admin_listen,并读取 admin.addr。
|
||||
AdminListen string
|
||||
// AdminInit 可选;nil 时用 CLIAdminInit。不支持时忽略并继续启动。
|
||||
AdminInit AdminInitializer
|
||||
// KeepDir 为 true 时 Stop 不删除数据目录(命令行启动器在清理前可读)。
|
||||
KeepDir bool
|
||||
}
|
||||
|
||||
// Server 是一次集成测试用的真实 nixmsg 进程。
|
||||
type Server struct {
|
||||
BinPath string
|
||||
ConfigPath string
|
||||
DataDir string
|
||||
Addr string // 来自 listen.addr,形如 127.0.0.1:12345
|
||||
AdminAddr string // 来自 admin.addr(若有)
|
||||
AdminPassword string
|
||||
HTTPBase string
|
||||
AdminHTTPBase string
|
||||
|
||||
cmd *exec.Cmd
|
||||
keepDir bool
|
||||
stopped bool
|
||||
}
|
||||
|
||||
// Start 编译(如需)并启动服务:临时目录、listen 端口 0、读 listen.addr。
|
||||
func Start(opts Options) (*Server, error) {
|
||||
bin, err := Binary()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
dataDir, err := os.MkdirTemp("", "nixmsg-harness-*")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
listen := opts.Listen
|
||||
if listen == "" {
|
||||
listen = "127.0.0.1:0"
|
||||
}
|
||||
cfgPath := filepath.Join(dataDir, "config.yaml")
|
||||
cfg := fmt.Sprintf("listen: %q\ndata_dir: %q\n", listen, filepath.ToSlash(dataDir))
|
||||
if opts.AdminListen != "" {
|
||||
cfg += fmt.Sprintf("admin_listen: %q\n", opts.AdminListen)
|
||||
}
|
||||
if err = os.WriteFile(cfgPath, []byte(cfg), 0o644); err != nil {
|
||||
_ = os.RemoveAll(dataDir)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
adminInit := opts.AdminInit
|
||||
if adminInit == nil {
|
||||
adminInit = CLIAdminInit{}
|
||||
}
|
||||
password, initErr := adminInit.Init(bin, cfgPath)
|
||||
if initErr != nil && !errors.Is(initErr, ErrAdminInitUnsupported) {
|
||||
_ = os.RemoveAll(dataDir)
|
||||
return nil, initErr
|
||||
}
|
||||
if errors.Is(initErr, ErrAdminInitUnsupported) {
|
||||
password = ""
|
||||
}
|
||||
|
||||
cmd := exec.Command(bin, "serve")
|
||||
cmd.Env = append(os.Environ(), "NIXMSG_CONFIG="+cfgPath)
|
||||
cmd.Stdout = os.Stderr
|
||||
cmd.Stderr = os.Stderr
|
||||
if err = cmd.Start(); err != nil {
|
||||
_ = os.RemoveAll(dataDir)
|
||||
return nil, fmt.Errorf("start serve: %w", err)
|
||||
}
|
||||
|
||||
s := &Server{
|
||||
BinPath: bin,
|
||||
ConfigPath: cfgPath,
|
||||
DataDir: dataDir,
|
||||
AdminPassword: password,
|
||||
cmd: cmd,
|
||||
keepDir: opts.KeepDir,
|
||||
}
|
||||
|
||||
addr, err := waitAddrFile(filepath.Join(dataDir, "listen.addr"), 10*time.Second)
|
||||
if err != nil {
|
||||
_ = s.Stop()
|
||||
return nil, fmt.Errorf("wait listen.addr: %w", err)
|
||||
}
|
||||
s.Addr = addr
|
||||
s.HTTPBase = "http://" + addr
|
||||
|
||||
if opts.AdminListen != "" {
|
||||
adminAddr, adminErr := waitAddrFile(filepath.Join(dataDir, "admin.addr"), 10*time.Second)
|
||||
if adminErr == nil {
|
||||
s.AdminAddr = adminAddr
|
||||
s.AdminHTTPBase = "http://" + adminAddr
|
||||
}
|
||||
}
|
||||
if s.AdminHTTPBase == "" {
|
||||
s.AdminHTTPBase = s.HTTPBase
|
||||
}
|
||||
return s, nil
|
||||
}
|
||||
|
||||
// AdminClient 返回指向管理接口的 HTTP 客户端。
|
||||
func (s *Server) AdminClient() (*AdminClient, error) {
|
||||
base := s.AdminHTTPBase
|
||||
if base == "" {
|
||||
base = s.HTTPBase
|
||||
}
|
||||
return NewAdminClient(base)
|
||||
}
|
||||
|
||||
// Stop 结束进程并删除临时目录(除非 KeepDir)。
|
||||
func (s *Server) Stop() error {
|
||||
if s == nil || s.stopped {
|
||||
return nil
|
||||
}
|
||||
s.stopped = true
|
||||
var stopErr error
|
||||
if s.cmd != nil && s.cmd.Process != nil {
|
||||
_ = s.cmd.Process.Kill()
|
||||
_, _ = s.cmd.Process.Wait()
|
||||
}
|
||||
if !s.keepDir && s.DataDir != "" {
|
||||
if err := os.RemoveAll(s.DataDir); err != nil {
|
||||
stopErr = err
|
||||
}
|
||||
}
|
||||
return stopErr
|
||||
}
|
||||
|
||||
func waitAddrFile(path string, timeout time.Duration) (string, error) {
|
||||
deadline := time.Now().Add(timeout)
|
||||
var lastErr error
|
||||
for time.Now().Before(deadline) {
|
||||
b, err := os.ReadFile(path)
|
||||
if err == nil {
|
||||
addr := strings.TrimSpace(string(b))
|
||||
if addr != "" {
|
||||
return addr, nil
|
||||
}
|
||||
lastErr = fmt.Errorf("empty addr file")
|
||||
} else {
|
||||
lastErr = err
|
||||
}
|
||||
time.Sleep(20 * time.Millisecond)
|
||||
}
|
||||
return "", lastErr
|
||||
}
|
||||
Reference in New Issue
Block a user