diff --git a/cmd/nixmsg/admin.go b/cmd/nixmsg/admin.go new file mode 100644 index 0000000..a88a979 --- /dev/null +++ b/cmd/nixmsg/admin.go @@ -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 ") + } + 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 = 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 = 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 +} diff --git a/cmd/nixmsg/backup.go b/cmd/nixmsg/backup.go new file mode 100644 index 0000000..585fa77 --- /dev/null +++ b/cmd/nixmsg/backup.go @@ -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 ") + } + out = args[i+1] + i++ + default: + return "", fmt.Errorf("unknown argument: %s", args[i]) + } + } + if out == "" { + return "", errors.New("usage: nixmsg backup --out ") + } + return out, nil +} diff --git a/cmd/nixmsg/checkconfig.go b/cmd/nixmsg/checkconfig.go new file mode 100644 index 0000000..4b95b0f --- /dev/null +++ b/cmd/nixmsg/checkconfig.go @@ -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 +} diff --git a/cmd/nixmsg/commands_test.go b/cmd/nixmsg/commands_test.go new file mode 100644 index 0000000..56831d7 --- /dev/null +++ b/cmd/nixmsg/commands_test.go @@ -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) + } +} diff --git a/cmd/nixmsg/healthcheck.go b/cmd/nixmsg/healthcheck.go new file mode 100644 index 0000000..e61c8c5 --- /dev/null +++ b/cmd/nixmsg/healthcheck.go @@ -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 +} diff --git a/cmd/nixmsg/main.go b/cmd/nixmsg/main.go index b335e86..a76f919 100644 --- a/cmd/nixmsg/main.go +++ b/cmd/nixmsg/main.go @@ -10,16 +10,28 @@ func main() { fmt.Fprintln(os.Stderr, "usage: nixmsg ") os.Exit(2) } + var err error switch os.Args[1] { case "version": cmdVersion(os.Args[2:]) case "serve": - if err := cmdServe(os.Args[2:]); err != nil { - fmt.Fprintln(os.Stderr, err) - os.Exit(1) - } + err = cmdServe(os.Args[2:]) + // P-WIRE-BEGIN + 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: fmt.Fprintf(os.Stderr, "unknown command: %s\n", os.Args[1]) os.Exit(2) } + if err != nil { + fmt.Fprintln(os.Stderr, err) + os.Exit(1) + } } diff --git a/cmd/nixmsg/serve.go b/cmd/nixmsg/serve.go index 7ad7383..f625b0d 100644 --- a/cmd/nixmsg/serve.go +++ b/cmd/nixmsg/serve.go @@ -4,11 +4,13 @@ import ( "context" "errors" "fmt" + "log/slog" "net" "net/http" "os" "os/signal" "path/filepath" + "strings" "syscall" "time" @@ -23,6 +25,12 @@ func cmdServe(_ []string) error { if err != nil { 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) defer stop() return runServe(ctx, cfg) @@ -41,6 +49,16 @@ func runServe(ctx context.Context, cfg config.Config) error { } 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 后生效。 if recoverErr := deps.Messages.RecoverOnStart(ctx); recoverErr != nil { 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.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() @@ -87,9 +116,17 @@ func runServe(ctx context.Context, cfg config.Config) error { select { 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() _ = 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 if serveErr != nil && !errors.Is(serveErr, http.ErrServerClosed) { return serveErr @@ -107,3 +144,20 @@ func writeListenAddr(dataDir, addr string) error { path := filepath.Join(dataDir, "listen.addr") 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 diff --git a/cmd/nixmsg/serve_test.go b/cmd/nixmsg/serve_test.go index 8524111..b91134c 100644 --- a/cmd/nixmsg/serve_test.go +++ b/cmd/nixmsg/serve_test.go @@ -10,23 +10,51 @@ import ( "testing" "time" + "git.asio.asia/nixevol/NixMsg/internal/auth" "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) { t.Parallel() dataDir := t.TempDir() - cfgPath := filepath.Join(dataDir, "config.yaml") - cfgYAML := []byte("listen: \"127.0.0.1:0\"\ndata_dir: \"" + filepath.ToSlash(dataDir) + "\"\n") - if err := os.WriteFile(cfgPath, cfgYAML, 0o644); err != nil { - t.Fatal(err) - } + cfgPath := writeTestConfig(t, dataDir) + initAdminForTest(t, dataDir) cfg, err := config.Load(cfgPath) if err != nil { t.Fatal(err) } + if vErr := cfg.Validate(); vErr != nil { + t.Fatal(vErr) + } ctx, cancel := context.WithCancel(context.Background()) defer cancel() @@ -37,7 +65,7 @@ func TestServeHealthzAndListenAddr(t *testing.T) { }() var addr string - deadline := time.Now().Add(5 * time.Second) + deadline := time.Now().Add(10 * time.Second) for time.Now().Before(deadline) { b, readErr := os.ReadFile(filepath.Join(dataDir, "listen.addr")) if readErr == nil { @@ -65,6 +93,15 @@ func TestServeHealthzAndListenAddr(t *testing.T) { 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 { t.Fatalf("db missing: %v", err) } @@ -75,7 +112,23 @@ func TestServeHealthzAndListenAddr(t *testing.T) { if err != nil { t.Fatalf("serve exit: %v", err) } - case <-time.After(5 * time.Second): + case <-time.After(10 * time.Second): 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) + } +} diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index 49b1f99..12ef3e0 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -152,7 +152,35 @@ ## 平台 P -暂无。 +### P1 2026-09-30 + +1. **admin set-password 传参方式** + - 原条款:DEVELOPMENT 11.2 仅列命令名,未规定密码如何传入。 + - 实际做法:支持 `--password `、位置参数,或从 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 diff --git a/go.mod b/go.mod index 5768d0c..c544a92 100644 --- a/go.mod +++ b/go.mod @@ -4,6 +4,7 @@ go 1.27 require ( go.yaml.in/yaml/v3 v3.0.5 + golang.org/x/crypto v0.57.0 modernc.org/sqlite v1.60.1 ) diff --git a/go.sum b/go.sum index ad2cb52..e8dd51d 100644 --- a/go.sum +++ b/go.sum @@ -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= 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= +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/go.mod h1:Ek9pY8RKWXwsWvd3rQiHYtMqkjSUV+s1Rj7j4H5Ur6o= golang.org/x/sync v0.23.0 h1:KameEIfc1IkluZyXWLn39Wd4tURc6GbCiISGiZm2bQk= diff --git a/internal/auth/phc.go b/internal/auth/phc.go new file mode 100644 index 0000000..dd9176d --- /dev/null +++ b/internal/auth/phc.go @@ -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 +} diff --git a/internal/config/config.go b/internal/config/config.go index acecf0a..b762374 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -2,11 +2,16 @@ package config import ( "fmt" + "net" "os" + "strings" "go.yaml.in/yaml/v3" ) +// MaxBodyBytesCap 是 max_body_bytes 的硬上限(DEVELOPMENT 11.1)。 +const MaxBodyBytesCap = 262144 + // Config 对应 DEVELOPMENT 第 11.1 节的配置文件。 type Config struct { Listen string `yaml:"listen"` @@ -65,7 +70,7 @@ func Default() Config { TrustedProxies: nil, DataDir: "./data", Limits: LimitsConfig{ - MaxBodyBytes: 262144, + MaxBodyBytes: MaxBodyBytesCap, MaxMetaBytes: 4096, MaxFrameBytes: 786432, MaxTTLSeconds: 2592000, @@ -101,16 +106,148 @@ func Load(path string) (Config, error) { if err := yaml.Unmarshal(data, &cfg); err != nil { return Config{}, fmt.Errorf("parse config: %w", err) } + applyEmptyDefaults(&cfg) + return cfg, nil +} + +func applyEmptyDefaults(cfg *Config) { + def := Default() if cfg.Listen == "" { - cfg.Listen = Default().Listen + cfg.Listen = def.Listen } if cfg.DataDir == "" { - cfg.DataDir = Default().DataDir + cfg.DataDir = def.DataDir } 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。 diff --git a/internal/config/config_test.go b/internal/config/config_test.go new file mode 100644 index 0000000..7d1b6b3 --- /dev/null +++ b/internal/config/config_test.go @@ -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) + } +} diff --git a/internal/store/admin.go b/internal/store/admin.go new file mode 100644 index 0000000..bcf90b7 --- /dev/null +++ b/internal/store/admin.go @@ -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 +} diff --git a/internal/store/queue.go b/internal/store/queue.go index 2f7f487..4286d1a 100644 --- a/internal/store/queue.go +++ b/internal/store/queue.go @@ -5,28 +5,34 @@ import ( "database/sql" "errors" "sync" + "time" ) // ErrQueueClosed 表示写入队列已关闭。 var ErrQueueClosed = errors.New("store: write queue closed") +// ErrBusy 表示写库基础设施失败(可映射为协议 busy;/readyz 应失败)。 +var ErrBusy = errors.New("store: busy") + // WriteFunc 在单个写事务中执行的操作。 type WriteFunc func(tx *sql.Tx) error // Queue 写入队列:提交一个写操作并拿到结果。 // -// 本任务(T0.3)实现为互斥串行的一操作一事务,不做合并;DEVELOPMENT 7.2 -// 要求的写 goroutine 合并提交(最多 256 个或凑满 2ms、SAVEPOINT 隔离失败)留给 P2。 +// P1 仍为一操作一事务;P2 换成合并提交(最多 256 / 2ms + SAVEPOINT)。 type Queue struct { db *sql.DB - mu sync.Mutex - closed bool + mu sync.Mutex + closed bool + inflight int + ready bool + lastWriteErr error } // NewQueue 创建简单写入队列(一操作一事务)。 func NewQueue(db *sql.DB) *Queue { - return &Queue{db: db} + return &Queue{db: db, ready: true} } // Do 提交写操作并等待提交结果。 @@ -35,22 +41,83 @@ func (q *Queue) Do(ctx context.Context, fn WriteFunc) error { return errors.New("store: nil write func") } q.mu.Lock() - defer q.mu.Unlock() if q.closed { + q.mu.Unlock() return ErrQueueClosed } + q.inflight++ + q.mu.Unlock() + + defer func() { + q.mu.Lock() + q.inflight-- + q.mu.Unlock() + }() + if err := ctx.Err(); err != nil { return err } tx, err := q.db.BeginTx(ctx, nil) if err != nil { - return err + q.markBusy(err) + return errors.Join(ErrBusy, err) } if err := fn(tx); err != nil { _ = tx.Rollback() 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。 diff --git a/internal/store/ready.go b/internal/store/ready.go new file mode 100644 index 0000000..ed7778f --- /dev/null +++ b/internal/store/ready.go @@ -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 +}