feat: 实现配置校验与运维命令及优雅停机

This commit is contained in:
Nixevol
2026-09-30 06:47:16 +08:00
parent 22c56d1f35
commit d777a1e3d3
17 changed files with 1053 additions and 26 deletions
+170
View File
@@ -0,0 +1,170 @@
package main
import (
"context"
"crypto/rand"
"errors"
"fmt"
"io"
"os"
"strings"
"git.asio.asia/nixevol/NixMsg/internal/auth"
"git.asio.asia/nixevol/NixMsg/internal/config"
"git.asio.asia/nixevol/NixMsg/internal/store"
)
const (
adminPasswordLen = 20
minAdminPasswordLen = 12
adminPasswordChars = "ABCDEFGHJKLMNPQRSTUVWXYZabcdefghijkmnopqrstuvwxyz23456789"
)
func cmdAdmin(args []string) error {
if len(args) < 1 {
return errors.New("usage: nixmsg admin <init|set-password>")
}
switch args[0] {
case "init":
return cmdAdminInit(args[1:])
case "set-password":
return cmdAdminSetPassword(args[1:])
default:
return fmt.Errorf("unknown admin command: %s", args[0])
}
}
func cmdAdminInit(_ []string) error {
cfg, err := loadAndValidateConfig()
if err != nil {
return err
}
ctx := context.Background()
db, err := store.Open(cfg.DataDir, cfg.SQLiteSynchronous)
if err != nil {
return err
}
defer func() { _ = db.Close() }()
ok, err := store.HasAdminPassword(ctx, db.Write)
if err != nil {
return err
}
if ok {
return errors.New("admin already initialized; use admin set-password")
}
password, err := generateAdminPassword(adminPasswordLen)
if err != nil {
return err
}
phc, err := auth.HashPassword(password)
if err != nil {
return err
}
if err := store.SetAdminPasswordHash(ctx, db.Write, phc); err != nil {
return err
}
// 密码只打印到终端一次,不进日志。
fmt.Printf("admin password: %s\n", password)
return nil
}
func cmdAdminSetPassword(args []string) error {
password, err := parseSetPasswordArgs(args, os.Stdin)
if err != nil {
return err
}
if vErr := validateAdminPassword(password); vErr != nil {
return vErr
}
cfg, err := loadAndValidateConfig()
if err != nil {
return err
}
ctx := context.Background()
db, err := store.Open(cfg.DataDir, cfg.SQLiteSynchronous)
if err != nil {
return err
}
defer func() { _ = db.Close() }()
phc, err := auth.HashPassword(password)
if err != nil {
return err
}
if err := store.SetAdminPasswordHash(ctx, db.Write, phc); err != nil {
return err
}
fmt.Fprintln(os.Stderr, "admin password updated")
return nil
}
func parseSetPasswordArgs(args []string, stdin io.Reader) (string, error) {
var password string
for i := 0; i < len(args); i++ {
switch args[i] {
case "--password":
if i+1 >= len(args) {
return "", errors.New("usage: nixmsg admin set-password --password <password>")
}
password = args[i+1]
i++
default:
if strings.HasPrefix(args[i], "-") {
return "", fmt.Errorf("unknown flag: %s", args[i])
}
if password != "" {
return "", errors.New("usage: nixmsg admin set-password --password <password>")
}
password = args[i]
}
}
if password == "" {
b, err := io.ReadAll(io.LimitReader(stdin, 4096))
if err != nil {
return "", err
}
password = strings.TrimSpace(string(b))
}
if password == "" {
return "", errors.New("password required (pass --password or stdin)")
}
return password, nil
}
func validateAdminPassword(password string) error {
if len(password) < minAdminPasswordLen {
return fmt.Errorf("admin password must be at least %d characters", minAdminPasswordLen)
}
return nil
}
func generateAdminPassword(n int) (string, error) {
buf := make([]byte, n)
charset := []byte(adminPasswordChars)
for i := 0; i < n; {
var b [1]byte
if _, err := rand.Read(b[:]); err != nil {
return "", err
}
if int(b[0]) >= 256-(256%len(charset)) {
continue
}
buf[i] = charset[int(b[0])%len(charset)]
i++
}
return string(buf), nil
}
func loadAndValidateConfig() (config.Config, error) {
cfgPath := config.PathFromEnv()
cfg, err := config.Load(cfgPath)
if err != nil {
return config.Config{}, err
}
if err := cfg.Validate(); err != nil {
return config.Config{}, err
}
return cfg, nil
}
+63
View File
@@ -0,0 +1,63 @@
package main
import (
"context"
"errors"
"fmt"
"os"
"path/filepath"
"git.asio.asia/nixevol/NixMsg/internal/store"
)
func cmdBackup(args []string) error {
outPath, err := parseBackupArgs(args)
if err != nil {
return err
}
cfg, err := loadAndValidateConfig()
if err != nil {
return err
}
if dir := filepath.Dir(outPath); dir != "" && dir != "." {
if mkErr := os.MkdirAll(dir, 0o755); mkErr != nil {
return fmt.Errorf("mkdir backup dir: %w", mkErr)
}
}
absOut, err := filepath.Abs(outPath)
if err != nil {
return err
}
ctx := context.Background()
// 对运行中的库:单独打开写连接执行 VACUUM INTO(可与 serve 并存,WAL 下安全)。
write, err := store.OpenWriter(cfg.DataDir, cfg.SQLiteSynchronous)
if err != nil {
return err
}
defer func() { _ = write.Close() }()
if err := store.VacuumInto(ctx, write, filepath.ToSlash(absOut)); err != nil {
return err
}
fmt.Fprintf(os.Stderr, "backup written to %s\n", absOut)
return nil
}
func parseBackupArgs(args []string) (string, error) {
var out string
for i := 0; i < len(args); i++ {
switch args[i] {
case "--out":
if i+1 >= len(args) {
return "", errors.New("usage: nixmsg backup --out <file.db>")
}
out = args[i+1]
i++
default:
return "", fmt.Errorf("unknown argument: %s", args[i])
}
}
if out == "" {
return "", errors.New("usage: nixmsg backup --out <file.db>")
}
return out, nil
}
+24
View File
@@ -0,0 +1,24 @@
package main
import (
"fmt"
"os"
"git.asio.asia/nixevol/NixMsg/internal/config"
)
func cmdCheckConfig(args []string) error {
path := config.PathFromEnv()
if len(args) > 0 {
path = args[0]
}
cfg, err := config.Load(path)
if err != nil {
return err
}
if err := cfg.Validate(); err != nil {
return err
}
fmt.Fprintln(os.Stderr, "config ok")
return nil
}
+153
View File
@@ -0,0 +1,153 @@
package main
import (
"bytes"
"context"
"os"
"path/filepath"
"strings"
"testing"
"time"
"git.asio.asia/nixevol/NixMsg/internal/store"
)
func TestCheckConfigRejectsLargeBody(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "bad.yaml")
body := "listen: \":0\"\ndata_dir: \"" + filepath.ToSlash(dir) + "\"\nlimits:\n max_body_bytes: 262145\n"
if err := os.WriteFile(path, []byte(body), 0o644); err != nil {
t.Fatal(err)
}
t.Setenv("NIXMSG_CONFIG", path)
err := cmdCheckConfig(nil)
if err == nil || !strings.Contains(err.Error(), "max_body_bytes") {
t.Fatalf("want max_body_bytes error, got %v", err)
}
}
func TestCheckConfigOK(t *testing.T) {
dir := t.TempDir()
path := writeTestConfig(t, dir)
t.Setenv("NIXMSG_CONFIG", path)
if err := cmdCheckConfig(nil); err != nil {
t.Fatal(err)
}
}
func TestAdminInitOnceAndSetPassword(t *testing.T) {
dir := t.TempDir()
path := writeTestConfig(t, dir)
t.Setenv("NIXMSG_CONFIG", path)
var out bytes.Buffer
old := os.Stdout
r, w, err := os.Pipe()
if err != nil {
t.Fatal(err)
}
os.Stdout = w
errInit := cmdAdminInit(nil)
_ = w.Close()
os.Stdout = old
_, _ = out.ReadFrom(r)
_ = r.Close()
if errInit != nil {
t.Fatal(errInit)
}
text := out.String()
if !strings.Contains(text, "admin password:") {
t.Fatalf("password not printed: %q", text)
}
pass := strings.TrimSpace(strings.TrimPrefix(strings.TrimSpace(text), "admin password:"))
if len(pass) != 20 {
t.Fatalf("password len=%d value=%q", len(pass), pass)
}
if err2 := cmdAdminInit(nil); err2 == nil || !strings.Contains(err2.Error(), "already initialized") {
t.Fatalf("want already initialized, got %v", err2)
}
if err2 := cmdAdminSetPassword([]string{"--password", "short"}); err2 == nil {
t.Fatal("expected short password error")
}
if err2 := cmdAdminSetPassword([]string{"--password", "long-enough-password"}); err2 != nil {
t.Fatal(err2)
}
db, openErr := store.Open(dir, "FULL")
if openErr != nil {
t.Fatal(openErr)
}
defer func() { _ = db.Close() }()
ok, hasErr := store.HasAdminPassword(context.Background(), db.Write)
if hasErr != nil || !ok {
t.Fatalf("has admin: ok=%v err=%v", ok, hasErr)
}
}
func TestBackupVacuumInto(t *testing.T) {
dir := t.TempDir()
path := writeTestConfig(t, dir)
t.Setenv("NIXMSG_CONFIG", path)
initAdminForTest(t, dir)
out := filepath.Join(dir, "backup", "copy.db")
if err := os.MkdirAll(filepath.Dir(out), 0o755); err != nil {
t.Fatal(err)
}
if err := cmdBackup([]string{"--out", out}); err != nil {
t.Fatal(err)
}
st, err := os.Stat(out)
if err != nil {
t.Fatal(err)
}
if st.Size() == 0 {
t.Fatal("backup empty")
}
}
func TestHealthcheck(t *testing.T) {
dir := t.TempDir()
cfgPath := writeTestConfig(t, dir)
initAdminForTest(t, dir)
t.Setenv("NIXMSG_CONFIG", cfgPath)
cfg, err := loadAndValidateConfig()
if err != nil {
t.Fatal(err)
}
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
errCh := make(chan error, 1)
go func() { errCh <- runServe(ctx, cfg) }()
deadline := time.Now().Add(10 * time.Second)
for time.Now().Before(deadline) {
if b, readErr := os.ReadFile(filepath.Join(dir, "listen.addr")); readErr == nil && strings.TrimSpace(string(b)) != "" {
break
}
time.Sleep(20 * time.Millisecond)
}
if err := cmdHealthcheck(nil); err != nil {
cancel()
<-errCh
t.Fatal(err)
}
cancel()
<-errCh
}
func TestParseSetPasswordArgs(t *testing.T) {
t.Parallel()
pass, err := parseSetPasswordArgs([]string{"--password", "abcdefghijkl"}, nil)
if err != nil || pass != "abcdefghijkl" {
t.Fatalf("pass=%q err=%v", pass, err)
}
pass, err = parseSetPasswordArgs([]string{"twelvechars!!"}, nil)
if err != nil || pass != "twelvechars!!" {
t.Fatalf("pass=%q err=%v", pass, err)
}
}
+38
View File
@@ -0,0 +1,38 @@
package main
import (
"fmt"
"io"
"net/http"
"os"
"path/filepath"
"strings"
"time"
)
func cmdHealthcheck(_ []string) error {
cfg, err := loadAndValidateConfig()
if err != nil {
return err
}
addrBytes, err := os.ReadFile(filepath.Join(cfg.DataDir, "listen.addr"))
if err != nil {
return fmt.Errorf("read listen.addr: %w", err)
}
addr := strings.TrimSpace(string(addrBytes))
if addr == "" {
return fmt.Errorf("listen.addr is empty")
}
url := "http://" + addr + "/healthz"
client := &http.Client{Timeout: 5 * time.Second}
resp, err := client.Get(url)
if err != nil {
return err
}
defer func() { _ = resp.Body.Close() }()
body, _ := io.ReadAll(io.LimitReader(resp.Body, 1024))
if resp.StatusCode != http.StatusOK {
return fmt.Errorf("healthz status %d: %s", resp.StatusCode, strings.TrimSpace(string(body)))
}
return nil
}
+16 -4
View File
@@ -10,16 +10,28 @@ func main() {
fmt.Fprintln(os.Stderr, "usage: nixmsg <command>") fmt.Fprintln(os.Stderr, "usage: nixmsg <command>")
os.Exit(2) os.Exit(2)
} }
var err error
switch os.Args[1] { switch os.Args[1] {
case "version": case "version":
cmdVersion(os.Args[2:]) cmdVersion(os.Args[2:])
case "serve": case "serve":
if err := cmdServe(os.Args[2:]); err != nil { err = cmdServe(os.Args[2:])
fmt.Fprintln(os.Stderr, err) // P-WIRE-BEGIN
os.Exit(1) case "admin":
} err = cmdAdmin(os.Args[2:])
case "backup":
err = cmdBackup(os.Args[2:])
case "check-config":
err = cmdCheckConfig(os.Args[2:])
case "healthcheck":
err = cmdHealthcheck(os.Args[2:])
// P-WIRE-END
default: default:
fmt.Fprintf(os.Stderr, "unknown command: %s\n", os.Args[1]) fmt.Fprintf(os.Stderr, "unknown command: %s\n", os.Args[1])
os.Exit(2) os.Exit(2)
} }
if err != nil {
fmt.Fprintln(os.Stderr, err)
os.Exit(1)
}
} }
+55 -1
View File
@@ -4,11 +4,13 @@ import (
"context" "context"
"errors" "errors"
"fmt" "fmt"
"log/slog"
"net" "net"
"net/http" "net/http"
"os" "os"
"os/signal" "os/signal"
"path/filepath" "path/filepath"
"strings"
"syscall" "syscall"
"time" "time"
@@ -23,6 +25,12 @@ func cmdServe(_ []string) error {
if err != nil { if err != nil {
return err return err
} }
// P-WIRE-BEGIN
if err := cfg.Validate(); err != nil {
return err
}
setupJSONLogger(cfg.Log)
// P-WIRE-END
ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
defer stop() defer stop()
return runServe(ctx, cfg) return runServe(ctx, cfg)
@@ -41,6 +49,16 @@ func runServe(ctx context.Context, cfg config.Config) error {
} }
defer func() { _ = db.Close() }() defer func() { _ = db.Close() }()
// P-WIRE-BEGIN
ok, err := store.HasAdminPassword(ctx, db.Write)
if err != nil {
return err
}
if !ok {
return errors.New("admin password not initialized; run: nixmsg admin init")
}
// P-WIRE-END
// 启动恢复入口已挂上(假实现为空操作);M 线替换 message.Service 后生效。 // 启动恢复入口已挂上(假实现为空操作);M 线替换 message.Service 后生效。
if recoverErr := deps.Messages.RecoverOnStart(ctx); recoverErr != nil { if recoverErr := deps.Messages.RecoverOnStart(ctx); recoverErr != nil {
return fmt.Errorf("message recover: %w", recoverErr) return fmt.Errorf("message recover: %w", recoverErr)
@@ -62,6 +80,17 @@ func runServe(ctx context.Context, cfg config.Config) error {
w.WriteHeader(http.StatusOK) w.WriteHeader(http.StatusOK)
_, _ = w.Write([]byte("ok")) _, _ = w.Write([]byte("ok"))
}) })
// P-WIRE-BEGIN
mux.HandleFunc("GET /readyz", func(w http.ResponseWriter, r *http.Request) {
if readyErr := db.Ready(r.Context()); readyErr != nil {
slog.Error("readyz failed", "err", readyErr)
http.Error(w, "not ready", http.StatusServiceUnavailable)
return
}
w.WriteHeader(http.StatusOK)
_, _ = w.Write([]byte("ok"))
})
// P-WIRE-END
// 确保前端资源被链接进二进制;完整静态托管由后续任务完善。 // 确保前端资源被链接进二进制;完整静态托管由后续任务完善。
_ = web.Dist() _ = web.Dist()
@@ -87,9 +116,17 @@ func runServe(ctx context.Context, cfg config.Config) error {
select { select {
case <-ctx.Done(): case <-ctx.Done():
shutdownCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second) // P-WIRE-BEGIN
// 先停止接受新连接,再等写队列最多 10 秒,然后断开并退出。
shutdownCtx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
defer cancel() defer cancel()
_ = srv.Shutdown(shutdownCtx) _ = srv.Shutdown(shutdownCtx)
drainCtx, drainCancel := context.WithTimeout(context.Background(), 10*time.Second)
defer drainCancel()
if drainErr := db.Queue.Drain(drainCtx); drainErr != nil && !errors.Is(drainErr, context.DeadlineExceeded) {
slog.Error("write queue drain", "err", drainErr)
}
// P-WIRE-END
serveErr := <-errCh serveErr := <-errCh
if serveErr != nil && !errors.Is(serveErr, http.ErrServerClosed) { if serveErr != nil && !errors.Is(serveErr, http.ErrServerClosed) {
return serveErr return serveErr
@@ -107,3 +144,20 @@ func writeListenAddr(dataDir, addr string) error {
path := filepath.Join(dataDir, "listen.addr") path := filepath.Join(dataDir, "listen.addr")
return os.WriteFile(path, []byte(addr+"\n"), 0o644) return os.WriteFile(path, []byte(addr+"\n"), 0o644)
} }
// P-WIRE-BEGIN
func setupJSONLogger(cfg config.LogConfig) {
level := slog.LevelInfo
switch strings.ToLower(strings.TrimSpace(cfg.Level)) {
case "debug":
level = slog.LevelDebug
case "warn", "warning":
level = slog.LevelWarn
case "error":
level = slog.LevelError
}
h := slog.NewJSONHandler(os.Stdout, &slog.HandlerOptions{Level: level})
slog.SetDefault(slog.New(h))
}
// P-WIRE-END
+60 -7
View File
@@ -10,23 +10,51 @@ import (
"testing" "testing"
"time" "time"
"git.asio.asia/nixevol/NixMsg/internal/auth"
"git.asio.asia/nixevol/NixMsg/internal/config" "git.asio.asia/nixevol/NixMsg/internal/config"
"git.asio.asia/nixevol/NixMsg/internal/store"
) )
func writeTestConfig(t *testing.T, dataDir string) string {
t.Helper()
cfgPath := filepath.Join(dataDir, "config.yaml")
cfgYAML := []byte("listen: \"127.0.0.1:0\"\ndata_dir: \"" + filepath.ToSlash(dataDir) + "\"\nlog:\n level: error\n")
if err := os.WriteFile(cfgPath, cfgYAML, 0o644); err != nil {
t.Fatal(err)
}
return cfgPath
}
func initAdminForTest(t *testing.T, dataDir string) {
t.Helper()
db, err := store.Open(dataDir, "FULL")
if err != nil {
t.Fatal(err)
}
defer func() { _ = db.Close() }()
phc, err := auth.HashPassword("test-admin-password-xx")
if err != nil {
t.Fatal(err)
}
if err := store.SetAdminPasswordHash(context.Background(), db.Write, phc); err != nil {
t.Fatal(err)
}
}
func TestServeHealthzAndListenAddr(t *testing.T) { func TestServeHealthzAndListenAddr(t *testing.T) {
t.Parallel() t.Parallel()
dataDir := t.TempDir() dataDir := t.TempDir()
cfgPath := filepath.Join(dataDir, "config.yaml") cfgPath := writeTestConfig(t, dataDir)
cfgYAML := []byte("listen: \"127.0.0.1:0\"\ndata_dir: \"" + filepath.ToSlash(dataDir) + "\"\n") initAdminForTest(t, dataDir)
if err := os.WriteFile(cfgPath, cfgYAML, 0o644); err != nil {
t.Fatal(err)
}
cfg, err := config.Load(cfgPath) cfg, err := config.Load(cfgPath)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
if vErr := cfg.Validate(); vErr != nil {
t.Fatal(vErr)
}
ctx, cancel := context.WithCancel(context.Background()) ctx, cancel := context.WithCancel(context.Background())
defer cancel() defer cancel()
@@ -37,7 +65,7 @@ func TestServeHealthzAndListenAddr(t *testing.T) {
}() }()
var addr string var addr string
deadline := time.Now().Add(5 * time.Second) deadline := time.Now().Add(10 * time.Second)
for time.Now().Before(deadline) { for time.Now().Before(deadline) {
b, readErr := os.ReadFile(filepath.Join(dataDir, "listen.addr")) b, readErr := os.ReadFile(filepath.Join(dataDir, "listen.addr"))
if readErr == nil { if readErr == nil {
@@ -65,6 +93,15 @@ func TestServeHealthzAndListenAddr(t *testing.T) {
t.Fatalf("unexpected body: %q", body) t.Fatalf("unexpected body: %q", body)
} }
ready, err := http.Get("http://" + addr + "/readyz")
if err != nil {
t.Fatalf("readyz: %v", err)
}
defer func() { _ = ready.Body.Close() }()
if ready.StatusCode != http.StatusOK {
t.Fatalf("readyz status=%d", ready.StatusCode)
}
if _, err := os.Stat(filepath.Join(dataDir, "nixmsg.db")); err != nil { if _, err := os.Stat(filepath.Join(dataDir, "nixmsg.db")); err != nil {
t.Fatalf("db missing: %v", err) t.Fatalf("db missing: %v", err)
} }
@@ -75,7 +112,23 @@ func TestServeHealthzAndListenAddr(t *testing.T) {
if err != nil { if err != nil {
t.Fatalf("serve exit: %v", err) t.Fatalf("serve exit: %v", err)
} }
case <-time.After(5 * time.Second): case <-time.After(10 * time.Second):
t.Fatal("serve did not stop") t.Fatal("serve did not stop")
} }
} }
func TestServeRejectsWithoutAdmin(t *testing.T) {
t.Parallel()
dataDir := t.TempDir()
cfgPath := writeTestConfig(t, dataDir)
cfg, err := config.Load(cfgPath)
if err != nil {
t.Fatal(err)
}
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
err = runServe(ctx, cfg)
if err == nil || !strings.Contains(err.Error(), "admin password") {
t.Fatalf("want admin init error, got %v", err)
}
}
+29 -1
View File
@@ -152,7 +152,35 @@
## 平台 P ## 平台 P
暂无。 ### P1 2026-09-30
1. **admin set-password 传参方式**
- 原条款:DEVELOPMENT 11.2 仅列命令名,未规定密码如何传入。
- 实际做法:支持 `--password <pwd>`、位置参数,或从 stdin 读一行;最短 12 位。
- 原因:自动化测试与 Docker 非交互环境需要非交互传参。
- 备选方案:仅交互式 prompt。
- 影响:运维文档需写明推荐用 `--password` 或管道,勿把密码写进 shell 历史时可改用 stdin。
2. **admin init 密码字符集**
- 原条款:生成 20 位密码,未规定字符集。
- 实际做法:从去掉易混字符(0/O/1/I/l)的字母数字中均匀抽样 20 位。
- 原因:终端抄写友好。
- 备选方案:全 ASCII 可打印字符。
- 影响:熵略低于全字符集,对 20 位仍足够。
3. **PHC 哈希辅助先放在 auth,池在 P3**
- 原条款:argon2 并发池属 P3;admin init 需落库哈希属 P1。
- 实际做法:P1 在 `internal/auth` 提供 `HashPassword`/`VerifyPassword`(PHC),admin 命令直接调用;P3 再用同参数实现带并发上限的 `HashPool`。
- 原因:避免 P1 在 cmd 内复制算法,又不等到 P3 才做 init。
- 备选方案:P1 cmd 内联 argon2;或 P1/P3 合并提交。
- 影响:P3 合入后 admin 可改为走 HashPool(非必须,单次 init 无并发压力)。
4. **优雅停机顺序**
- 原条款:停接受 → 写队列最多 10 秒 → 断开连接退出。
- 实际做法:`http.Server.Shutdown`(先停接受并等待进行中的 HTTP)后,再 `Queue.Drain` 最多 10 秒;本期尚无长连接表,断开连接由 Shutdown 覆盖。
- 原因:当前 main 仅有 HTTP 健康检查;N 线接入后应在 Drain 前后显式踢连接。
- 备选方案:先关 listener、Drain、再 Shutdown。
- 影响:有 MQTT 长连接后需 N/P 联调停机路径。
## 连接 N ## 连接 N
+1
View File
@@ -4,6 +4,7 @@ go 1.27
require ( require (
go.yaml.in/yaml/v3 v3.0.5 go.yaml.in/yaml/v3 v3.0.5
golang.org/x/crypto v0.57.0
modernc.org/sqlite v1.60.1 modernc.org/sqlite v1.60.1
) )
+2
View File
@@ -14,6 +14,8 @@ github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo= github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo=
go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw= go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw=
go.yaml.in/yaml/v3 v3.0.5/go.mod h1:HVTZu1O7/Vkt2N+BFy8Zza+lnLsABggaTM2ZpNIGuKg= go.yaml.in/yaml/v3 v3.0.5/go.mod h1:HVTZu1O7/Vkt2N+BFy8Zza+lnLsABggaTM2ZpNIGuKg=
golang.org/x/crypto v0.57.0 h1:3ZVCjf8Ggz7zneR/EHRVx68Ctf+2pmIMP2UFhh9cC6M=
golang.org/x/crypto v0.57.0/go.mod h1:Fdz0i5U6CoizGwLda9DttjSk6qlZo25zYNtR+ycvuZA=
golang.org/x/mod v0.41.0 h1:qJmnOUb4YB+FsEuM3HcWucdZASCPGhsX6uljO6pog0c= golang.org/x/mod v0.41.0 h1:qJmnOUb4YB+FsEuM3HcWucdZASCPGhsX6uljO6pog0c=
golang.org/x/mod v0.41.0/go.mod h1:Ek9pY8RKWXwsWvd3rQiHYtMqkjSUV+s1Rj7j4H5Ur6o= golang.org/x/mod v0.41.0/go.mod h1:Ek9pY8RKWXwsWvd3rQiHYtMqkjSUV+s1Rj7j4H5Ur6o=
golang.org/x/sync v0.23.0 h1:KameEIfc1IkluZyXWLn39Wd4tURc6GbCiISGiZm2bQk= golang.org/x/sync v0.23.0 h1:KameEIfc1IkluZyXWLn39Wd4tURc6GbCiISGiZm2bQk=
+90
View File
@@ -0,0 +1,90 @@
package auth
import (
"crypto/rand"
"crypto/subtle"
"encoding/base64"
"errors"
"fmt"
"strings"
"golang.org/x/crypto/argon2"
)
// Argon2 参数(DEVELOPMENT 第 12 节 / OWASP 最低配置)。
const (
ArgonMemoryKiB = 19 * 1024 // 19 MiB
ArgonTime = 2
ArgonThreads = 1
ArgonKeyLen = 32
ArgonSaltLen = 16
)
// ErrInvalidPHC 表示 PHC 字符串无法解析或参数不支持。
var ErrInvalidPHC = errors.New("auth: invalid phc")
// HashPassword 用 argon2id 生成 PHC 字符串(不含并发池;池见 HashPool)。
func HashPassword(password string) (string, error) {
salt := make([]byte, ArgonSaltLen)
if _, err := rand.Read(salt); err != nil {
return "", err
}
hash := argon2.IDKey([]byte(password), salt, ArgonTime, ArgonMemoryKiB, ArgonThreads, ArgonKeyLen)
return encodePHC(salt, hash), nil
}
// VerifyPassword 常量时间比较密码与 PHC;不匹配时 ok=false 且 err=nil。
func VerifyPassword(password, phc string) (bool, error) {
salt, hash, err := decodePHC(phc)
if err != nil {
return false, err
}
got := argon2.IDKey([]byte(password), salt, ArgonTime, ArgonMemoryKiB, ArgonThreads, ArgonKeyLen)
if subtle.ConstantTimeCompare(got, hash) == 1 {
return true, nil
}
return false, nil
}
func encodePHC(salt, hash []byte) string {
return fmt.Sprintf(
"$argon2id$v=%d$m=%d,t=%d,p=%d$%s$%s",
argon2.Version,
ArgonMemoryKiB,
ArgonTime,
ArgonThreads,
base64.RawStdEncoding.EncodeToString(salt),
base64.RawStdEncoding.EncodeToString(hash),
)
}
func decodePHC(phc string) (salt, hash []byte, err error) {
// $argon2id$v=19$m=19456,t=2,p=1$salt$hash
parts := strings.Split(phc, "$")
if len(parts) != 6 || parts[1] != "argon2id" {
return nil, nil, ErrInvalidPHC
}
var version int
if _, scanErr := fmt.Sscanf(parts[2], "v=%d", &version); scanErr != nil || version != argon2.Version {
return nil, nil, ErrInvalidPHC
}
var m, t, p int
if _, scanErr := fmt.Sscanf(parts[3], "m=%d,t=%d,p=%d", &m, &t, &p); scanErr != nil {
return nil, nil, ErrInvalidPHC
}
if m != ArgonMemoryKiB || t != ArgonTime || p != ArgonThreads {
return nil, nil, ErrInvalidPHC
}
salt, err = base64.RawStdEncoding.DecodeString(parts[4])
if err != nil {
return nil, nil, ErrInvalidPHC
}
hash, err = base64.RawStdEncoding.DecodeString(parts[5])
if err != nil {
return nil, nil, ErrInvalidPHC
}
if len(salt) == 0 || len(hash) == 0 {
return nil, nil, ErrInvalidPHC
}
return salt, hash, nil
}
+142 -5
View File
@@ -2,11 +2,16 @@ package config
import ( import (
"fmt" "fmt"
"net"
"os" "os"
"strings"
"go.yaml.in/yaml/v3" "go.yaml.in/yaml/v3"
) )
// MaxBodyBytesCap 是 max_body_bytes 的硬上限(DEVELOPMENT 11.1)。
const MaxBodyBytesCap = 262144
// Config 对应 DEVELOPMENT 第 11.1 节的配置文件。 // Config 对应 DEVELOPMENT 第 11.1 节的配置文件。
type Config struct { type Config struct {
Listen string `yaml:"listen"` Listen string `yaml:"listen"`
@@ -65,7 +70,7 @@ func Default() Config {
TrustedProxies: nil, TrustedProxies: nil,
DataDir: "./data", DataDir: "./data",
Limits: LimitsConfig{ Limits: LimitsConfig{
MaxBodyBytes: 262144, MaxBodyBytes: MaxBodyBytesCap,
MaxMetaBytes: 4096, MaxMetaBytes: 4096,
MaxFrameBytes: 786432, MaxFrameBytes: 786432,
MaxTTLSeconds: 2592000, MaxTTLSeconds: 2592000,
@@ -101,16 +106,148 @@ func Load(path string) (Config, error) {
if err := yaml.Unmarshal(data, &cfg); err != nil { if err := yaml.Unmarshal(data, &cfg); err != nil {
return Config{}, fmt.Errorf("parse config: %w", err) return Config{}, fmt.Errorf("parse config: %w", err)
} }
applyEmptyDefaults(&cfg)
return cfg, nil
}
func applyEmptyDefaults(cfg *Config) {
def := Default()
if cfg.Listen == "" { if cfg.Listen == "" {
cfg.Listen = Default().Listen cfg.Listen = def.Listen
} }
if cfg.DataDir == "" { if cfg.DataDir == "" {
cfg.DataDir = Default().DataDir cfg.DataDir = def.DataDir
} }
if cfg.SQLiteSynchronous == "" { if cfg.SQLiteSynchronous == "" {
cfg.SQLiteSynchronous = Default().SQLiteSynchronous cfg.SQLiteSynchronous = def.SQLiteSynchronous
} }
return cfg, nil if cfg.Log.Level == "" {
cfg.Log.Level = def.Log.Level
}
if cfg.Limits.MaxBodyBytes == 0 {
cfg.Limits.MaxBodyBytes = def.Limits.MaxBodyBytes
}
if cfg.Limits.MaxMetaBytes == 0 {
cfg.Limits.MaxMetaBytes = def.Limits.MaxMetaBytes
}
if cfg.Limits.MaxFrameBytes == 0 {
cfg.Limits.MaxFrameBytes = def.Limits.MaxFrameBytes
}
if cfg.Limits.MaxTTLSeconds == 0 {
cfg.Limits.MaxTTLSeconds = def.Limits.MaxTTLSeconds
}
if cfg.Limits.MaxScheduleSeconds == 0 {
cfg.Limits.MaxScheduleSeconds = def.Limits.MaxScheduleSeconds
}
if cfg.Limits.MaxGroupMembers == 0 {
cfg.Limits.MaxGroupMembers = def.Limits.MaxGroupMembers
}
if cfg.Limits.GraceSeconds == 0 {
cfg.Limits.GraceSeconds = def.Limits.GraceSeconds
}
if cfg.Limits.AckTimeoutSeconds == 0 {
cfg.Limits.AckTimeoutSeconds = def.Limits.AckTimeoutSeconds
}
if cfg.Limits.DeliveryWindow == 0 {
cfg.Limits.DeliveryWindow = def.Limits.DeliveryWindow
}
if cfg.Limits.ReceiptWindow == 0 {
cfg.Limits.ReceiptWindow = def.Limits.ReceiptWindow
}
if cfg.Limits.RequestsPerSecond == 0 {
cfg.Limits.RequestsPerSecond = def.Limits.RequestsPerSecond
}
}
// Validate 校验 DEVELOPMENT 11.1 全部字段;拒绝 max_body_bytes > 262144。
func (c Config) Validate() error {
var errs []string
if strings.TrimSpace(c.Listen) == "" {
errs = append(errs, "listen is required")
}
if strings.TrimSpace(c.DataDir) == "" {
errs = append(errs, "data_dir is required")
}
sync := strings.ToUpper(strings.TrimSpace(c.SQLiteSynchronous))
if sync != "FULL" && sync != "NORMAL" {
errs = append(errs, "sqlite_synchronous must be FULL or NORMAL")
}
switch strings.ToLower(strings.TrimSpace(c.Log.Level)) {
case "debug", "info", "warn", "warning", "error":
default:
errs = append(errs, "log.level must be debug, info, warn, or error")
}
if c.Limits.MaxBodyBytes <= 0 {
errs = append(errs, "limits.max_body_bytes must be > 0")
}
if c.Limits.MaxBodyBytes > MaxBodyBytesCap {
errs = append(errs, fmt.Sprintf("limits.max_body_bytes must be <= %d", MaxBodyBytesCap))
}
if c.Limits.MaxMetaBytes <= 0 {
errs = append(errs, "limits.max_meta_bytes must be > 0")
}
if c.Limits.MaxFrameBytes <= 0 {
errs = append(errs, "limits.max_frame_bytes must be > 0")
}
if c.Limits.MaxFrameBytes < c.Limits.MaxBodyBytes {
errs = append(errs, "limits.max_frame_bytes must be >= limits.max_body_bytes")
}
if c.Limits.MaxTTLSeconds <= 0 {
errs = append(errs, "limits.max_ttl_seconds must be > 0")
}
if c.Limits.MaxScheduleSeconds <= 0 {
errs = append(errs, "limits.max_schedule_seconds must be > 0")
}
if c.Limits.MaxGroupMembers <= 0 {
errs = append(errs, "limits.max_group_members must be > 0")
}
if c.Limits.GraceSeconds < 0 {
errs = append(errs, "limits.grace_seconds must be >= 0")
}
if c.Limits.AckTimeoutSeconds <= 0 {
errs = append(errs, "limits.ack_timeout_seconds must be > 0")
}
if c.Limits.DeliveryWindow <= 0 {
errs = append(errs, "limits.delivery_window must be > 0")
}
if c.Limits.ReceiptWindow <= 0 {
errs = append(errs, "limits.receipt_window must be > 0")
}
if c.Limits.RequestsPerSecond <= 0 {
errs = append(errs, "limits.requests_per_second must be > 0")
}
if c.Limits.MaxPendingPerSender < 0 {
errs = append(errs, "limits.max_pending_per_sender must be >= 0")
}
if c.Limits.MaxPendingPerReceiver < 0 {
errs = append(errs, "limits.max_pending_per_receiver must be >= 0")
}
if c.SessionIdleDays < 0 {
errs = append(errs, "session_idle_days must be >= 0")
}
if c.RecordRetentionDays < 0 {
errs = append(errs, "record_retention_days must be >= 0")
}
if c.IdempotencyHours < 0 {
errs = append(errs, "idempotency_hours must be >= 0")
}
if c.ReceiptRetentionDays < 0 {
errs = append(errs, "receipt_retention_days must be >= 0")
}
cert := strings.TrimSpace(c.TLS.CertFile)
key := strings.TrimSpace(c.TLS.KeyFile)
if (cert == "") != (key == "") {
errs = append(errs, "tls.cert_file and tls.key_file must both be set or both empty")
}
for _, cidr := range c.TrustedProxies {
if _, _, err := net.ParseCIDR(strings.TrimSpace(cidr)); err != nil {
errs = append(errs, fmt.Sprintf("trusted_proxies entry %q is not a valid CIDR", cidr))
}
}
if len(errs) > 0 {
return fmt.Errorf("invalid config: %s", strings.Join(errs, "; "))
}
return nil
} }
// PathFromEnv 返回 NIXMSG_CONFIG 或默认 ./config.yaml。 // PathFromEnv 返回 NIXMSG_CONFIG 或默认 ./config.yaml。
+62
View File
@@ -0,0 +1,62 @@
package config
import (
"os"
"path/filepath"
"strings"
"testing"
)
func TestValidateDefaultOK(t *testing.T) {
t.Parallel()
cfg := Default()
if err := cfg.Validate(); err != nil {
t.Fatal(err)
}
}
func TestValidateRejectsMaxBodyBytesTooLarge(t *testing.T) {
t.Parallel()
cfg := Default()
cfg.Limits.MaxBodyBytes = MaxBodyBytesCap + 1
err := cfg.Validate()
if err == nil {
t.Fatal("expected error")
}
if !strings.Contains(err.Error(), "max_body_bytes") {
t.Fatalf("unexpected error: %v", err)
}
}
func TestLoadAndValidate(t *testing.T) {
t.Parallel()
dir := t.TempDir()
path := filepath.Join(dir, "config.yaml")
body := "listen: \"127.0.0.1:0\"\ndata_dir: \"" + filepath.ToSlash(dir) + "\"\n"
if err := os.WriteFile(path, []byte(body), 0o644); err != nil {
t.Fatal(err)
}
cfg, err := Load(path)
if err != nil {
t.Fatal(err)
}
if err := cfg.Validate(); err != nil {
t.Fatal(err)
}
if cfg.Listen != "127.0.0.1:0" {
t.Fatalf("listen=%q", cfg.Listen)
}
}
func TestValidateTrustedProxies(t *testing.T) {
t.Parallel()
cfg := Default()
cfg.TrustedProxies = []string{"not-a-cidr"}
if err := cfg.Validate(); err == nil {
t.Fatal("expected error")
}
cfg.TrustedProxies = []string{"127.0.0.1/32"}
if err := cfg.Validate(); err != nil {
t.Fatal(err)
}
}
+47
View File
@@ -0,0 +1,47 @@
package store
import (
"context"
"database/sql"
"fmt"
"strings"
"time"
)
const settingAdminPasswordHash = "admin_password_hash"
// HasAdminPassword 检查 settings 中是否已有管理员密码哈希。
func HasAdminPassword(ctx context.Context, db *sql.DB) (bool, error) {
var value string
err := db.QueryRowContext(ctx, `SELECT value FROM settings WHERE key = ?`, settingAdminPasswordHash).Scan(&value)
if err == sql.ErrNoRows {
return false, nil
}
if err != nil {
return false, err
}
return value != "", nil
}
// SetAdminPasswordHash 写入或覆盖管理员密码哈希。
func SetAdminPasswordHash(ctx context.Context, db *sql.DB, phc string) error {
now := time.Now().UnixMilli()
_, err := db.ExecContext(ctx, `
INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?)
ON CONFLICT(key) DO UPDATE SET value = excluded.value, updated_at = excluded.updated_at
`, settingAdminPasswordHash, phc, now)
return err
}
// VacuumInto 对打开的写连接执行 VACUUM INTO(可用于运行中备份)。
func VacuumInto(ctx context.Context, db *sql.DB, outPath string) error {
if outPath == "" {
return fmt.Errorf("store: vacuum into path empty")
}
escaped := strings.ReplaceAll(outPath, "'", "''")
_, err := db.ExecContext(ctx, "VACUUM INTO '"+escaped+"'")
if err != nil {
return fmt.Errorf("vacuum into: %w", err)
}
return nil
}
+75 -8
View File
@@ -5,28 +5,34 @@ import (
"database/sql" "database/sql"
"errors" "errors"
"sync" "sync"
"time"
) )
// ErrQueueClosed 表示写入队列已关闭。 // ErrQueueClosed 表示写入队列已关闭。
var ErrQueueClosed = errors.New("store: write queue closed") var ErrQueueClosed = errors.New("store: write queue closed")
// ErrBusy 表示写库基础设施失败(可映射为协议 busy;/readyz 应失败)。
var ErrBusy = errors.New("store: busy")
// WriteFunc 在单个写事务中执行的操作。 // WriteFunc 在单个写事务中执行的操作。
type WriteFunc func(tx *sql.Tx) error type WriteFunc func(tx *sql.Tx) error
// Queue 写入队列:提交一个写操作并拿到结果。 // Queue 写入队列:提交一个写操作并拿到结果。
// //
// 本任务(T0.3)实现为互斥串行的一操作一事务,不做合并;DEVELOPMENT 7.2 // P1 仍为一操作一事务;P2 换成合并提交(最多 256 / 2ms + SAVEPOINT)。
// 要求的写 goroutine 合并提交(最多 256 个或凑满 2ms、SAVEPOINT 隔离失败)留给 P2。
type Queue struct { type Queue struct {
db *sql.DB db *sql.DB
mu sync.Mutex mu sync.Mutex
closed bool closed bool
inflight int
ready bool
lastWriteErr error
} }
// NewQueue 创建简单写入队列(一操作一事务)。 // NewQueue 创建简单写入队列(一操作一事务)。
func NewQueue(db *sql.DB) *Queue { func NewQueue(db *sql.DB) *Queue {
return &Queue{db: db} return &Queue{db: db, ready: true}
} }
// Do 提交写操作并等待提交结果。 // Do 提交写操作并等待提交结果。
@@ -35,22 +41,83 @@ func (q *Queue) Do(ctx context.Context, fn WriteFunc) error {
return errors.New("store: nil write func") return errors.New("store: nil write func")
} }
q.mu.Lock() q.mu.Lock()
defer q.mu.Unlock()
if q.closed { if q.closed {
q.mu.Unlock()
return ErrQueueClosed return ErrQueueClosed
} }
q.inflight++
q.mu.Unlock()
defer func() {
q.mu.Lock()
q.inflight--
q.mu.Unlock()
}()
if err := ctx.Err(); err != nil { if err := ctx.Err(); err != nil {
return err return err
} }
tx, err := q.db.BeginTx(ctx, nil) tx, err := q.db.BeginTx(ctx, nil)
if err != nil { if err != nil {
return err q.markBusy(err)
return errors.Join(ErrBusy, err)
} }
if err := fn(tx); err != nil { if err := fn(tx); err != nil {
_ = tx.Rollback() _ = tx.Rollback()
return err return err
} }
return tx.Commit() if err := tx.Commit(); err != nil {
q.markBusy(err)
return errors.Join(ErrBusy, err)
}
return nil
}
func (q *Queue) markBusy(err error) {
q.mu.Lock()
defer q.mu.Unlock()
q.ready = false
q.lastWriteErr = err
}
// IsReady 写库是否仍可用(写失败后为 false)。
func (q *Queue) IsReady() bool {
q.mu.Lock()
defer q.mu.Unlock()
return q.ready
}
// LastWriteError 返回最近一次基础设施写失败。
func (q *Queue) LastWriteError() error {
q.mu.Lock()
defer q.mu.Unlock()
return q.lastWriteErr
}
// Len 返回进行中的写操作数。
func (q *Queue) Len() int {
q.mu.Lock()
defer q.mu.Unlock()
return q.inflight
}
// Drain 等待已进行中的写操作完成,或 ctx 取消。
func (q *Queue) Drain(ctx context.Context) error {
ticker := time.NewTicker(5 * time.Millisecond)
defer ticker.Stop()
for {
q.mu.Lock()
n := q.inflight
q.mu.Unlock()
if n == 0 {
return nil
}
select {
case <-ctx.Done():
return ctx.Err()
case <-ticker.C:
}
}
} }
// Close 关闭队列,之后 Do 返回 ErrQueueClosed。 // Close 关闭队列,之后 Do 返回 ErrQueueClosed。
+26
View File
@@ -0,0 +1,26 @@
package store
import (
"context"
"errors"
)
// Ready 检查读写库可用且写入队列未因磁盘类失败进入 busy。
func (d *DB) Ready(ctx context.Context) error {
if d == nil || d.Write == nil || d.Read == nil {
return errors.New("store: not open")
}
if err := d.Write.PingContext(ctx); err != nil {
return err
}
if err := d.Read.PingContext(ctx); err != nil {
return err
}
if d.Queue != nil && !d.Queue.IsReady() {
if err := d.Queue.LastWriteError(); err != nil {
return errors.Join(ErrBusy, err)
}
return ErrBusy
}
return nil
}