Compare commits

..
15 changed files with 1519 additions and 56 deletions
+1 -5
View File
@@ -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)
+47
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
}
+131
View File
@@ -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);
+62
View File
@@ -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
}
+161
View File
@@ -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
}
}
+55
View File
@@ -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")
}
}
+63
View File
@@ -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
}
}
+35
View File
@@ -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)
}
+116
View File
@@ -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")
}
}
}
+34
View File
@@ -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,
}
}
+321
View File
@@ -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
}
+162
View File
@@ -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
}