feat: 实现配置校验与运维命令及优雅停机
This commit is contained in:
@@ -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
@@ -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。
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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。
|
||||
|
||||
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user