fix: 修复写队列 busy 恢复、关闭安全、备份与配置构建问题
This commit is contained in:
+4
-2
@@ -5,6 +5,8 @@ vars:
|
||||
BIN_NAME: nixmsg
|
||||
EMBED_TAG: embeddist
|
||||
EXE: '{{if eq OS "windows"}}.exe{{end}}'
|
||||
VERSION:
|
||||
sh: git describe --tags --always --dirty
|
||||
|
||||
includes:
|
||||
'*':
|
||||
@@ -44,7 +46,7 @@ tasks:
|
||||
desc: 构建嵌入前端的单个可执行文件
|
||||
deps: [web:build]
|
||||
cmds:
|
||||
- go build -tags {{.EMBED_TAG}} -o {{.BIN_DIR}}/{{.BIN_NAME}}{{.EXE}} ./cmd/nixmsg
|
||||
- go build -tags {{.EMBED_TAG}} -ldflags "-X main.Version={{.VERSION}}" -o {{.BIN_DIR}}/{{.BIN_NAME}}{{.EXE}} ./cmd/nixmsg
|
||||
env:
|
||||
CGO_ENABLED: "0"
|
||||
|
||||
@@ -74,4 +76,4 @@ tasks:
|
||||
docker:
|
||||
desc: 构建开发用 Docker 镜像
|
||||
cmds:
|
||||
- docker build -t nixmsg:dev -f deploy/Dockerfile .
|
||||
- docker build --build-arg VERSION={{.VERSION}} -t nixmsg:dev -f deploy/Dockerfile .
|
||||
|
||||
+1
-1
@@ -93,7 +93,7 @@ func cmdAdminSetPassword(args []string) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := store.SetAdminPasswordHash(ctx, db.Write, phc); err != nil {
|
||||
if err := db.ResetAdminPassword(ctx, phc); err != nil {
|
||||
return err
|
||||
}
|
||||
fmt.Fprintln(os.Stderr, "admin password updated")
|
||||
|
||||
+24
-3
@@ -19,6 +19,25 @@ func cmdBackup(args []string) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
dataAbs, err := filepath.Abs(cfg.DataDir)
|
||||
if err != nil {
|
||||
return fmt.Errorf("resolve data_dir: %w", err)
|
||||
}
|
||||
srcAbs, err := filepath.Abs(filepath.Join(dataAbs, store.DBFileName))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
st, err := os.Stat(srcAbs)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return fmt.Errorf("database not found: %s", srcAbs)
|
||||
}
|
||||
return fmt.Errorf("stat database %s: %w", srcAbs, err)
|
||||
}
|
||||
if st.IsDir() {
|
||||
return fmt.Errorf("database path is a directory: %s", srcAbs)
|
||||
}
|
||||
|
||||
if dir := filepath.Dir(outPath); dir != "" && dir != "." {
|
||||
if mkErr := os.MkdirAll(dir, 0o755); mkErr != nil {
|
||||
return fmt.Errorf("mkdir backup dir: %w", mkErr)
|
||||
@@ -29,8 +48,7 @@ func cmdBackup(args []string) error {
|
||||
return err
|
||||
}
|
||||
ctx := context.Background()
|
||||
// 对运行中的库:单独打开写连接执行 VACUUM INTO(可与 serve 并存,WAL 下安全)。
|
||||
write, err := store.OpenWriter(cfg.DataDir, cfg.SQLiteSynchronous)
|
||||
write, err := store.OpenExistingWriter(dataAbs, cfg.SQLiteSynchronous)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -38,7 +56,10 @@ func cmdBackup(args []string) error {
|
||||
if err := store.VacuumInto(ctx, write, filepath.ToSlash(absOut)); err != nil {
|
||||
return err
|
||||
}
|
||||
fmt.Fprintf(os.Stderr, "backup written to %s\n", absOut)
|
||||
if chErr := os.Chmod(absOut, 0o600); chErr != nil {
|
||||
return fmt.Errorf("chmod backup: %w", chErr)
|
||||
}
|
||||
fmt.Fprintf(os.Stderr, "backup written from %s to %s\n", srcAbs, absOut)
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
@@ -3,8 +3,13 @@ package main
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/cookiejar"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
@@ -108,6 +113,92 @@ func TestBackupVacuumInto(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestBackupMissingDB(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := writeTestConfig(t, dir)
|
||||
t.Setenv("NIXMSG_CONFIG", path)
|
||||
|
||||
out := filepath.Join(dir, "copy.db")
|
||||
err := cmdBackup([]string{"--out", out})
|
||||
if err == nil || !strings.Contains(err.Error(), "database not found") {
|
||||
t.Fatalf("want database not found, got %v", err)
|
||||
}
|
||||
if _, statErr := os.Stat(filepath.Join(dir, store.DBFileName)); !os.IsNotExist(statErr) {
|
||||
t.Fatalf("backup must not create %s: %v", store.DBFileName, statErr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAdminSetPasswordClearsSessions(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) }()
|
||||
addr := waitListenAddr(t, dir, 15*time.Second)
|
||||
base := "http://" + addr
|
||||
|
||||
jar, err := cookiejar.New(nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
client := &http.Client{Jar: jar, Timeout: 10 * time.Second}
|
||||
loginBody := `{"username":"admin","password":"test-admin-password-xx"}`
|
||||
resp, err := client.Post(base+"/api/admin/login", "application/json", strings.NewReader(loginBody))
|
||||
if err != nil {
|
||||
cancel()
|
||||
<-errCh
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, _ = io.ReadAll(resp.Body)
|
||||
_ = resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
cancel()
|
||||
<-errCh
|
||||
t.Fatalf("login status=%d", resp.StatusCode)
|
||||
}
|
||||
|
||||
me, err := client.Get(base + "/api/admin/me")
|
||||
if err != nil {
|
||||
cancel()
|
||||
<-errCh
|
||||
t.Fatal(err)
|
||||
}
|
||||
_ = me.Body.Close()
|
||||
if me.StatusCode != http.StatusOK {
|
||||
cancel()
|
||||
<-errCh
|
||||
t.Fatalf("me before set-password status=%d", me.StatusCode)
|
||||
}
|
||||
|
||||
if err := cmdAdminSetPassword([]string{"--password", "long-enough-password"}); err != nil {
|
||||
cancel()
|
||||
<-errCh
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
me2, err := client.Get(base + "/api/admin/me")
|
||||
if err != nil {
|
||||
cancel()
|
||||
<-errCh
|
||||
t.Fatal(err)
|
||||
}
|
||||
body, _ := io.ReadAll(me2.Body)
|
||||
_ = me2.Body.Close()
|
||||
cancel()
|
||||
<-errCh
|
||||
if me2.StatusCode != http.StatusUnauthorized {
|
||||
t.Fatalf("want 401 after set-password, got %d body=%s", me2.StatusCode, body)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHealthcheck(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
cfgPath := writeTestConfig(t, dir)
|
||||
@@ -151,3 +242,48 @@ func TestParseSetPasswordArgs(t *testing.T) {
|
||||
t.Fatalf("pass=%q err=%v", pass, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCmdVersionPrintsInjected(t *testing.T) {
|
||||
old := Version
|
||||
Version = "d05-injected"
|
||||
defer func() { Version = old }()
|
||||
|
||||
r, w, err := os.Pipe()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
oldOut := os.Stdout
|
||||
os.Stdout = w
|
||||
cmdVersion(nil)
|
||||
_ = w.Close()
|
||||
os.Stdout = oldOut
|
||||
var buf bytes.Buffer
|
||||
_, _ = buf.ReadFrom(r)
|
||||
_ = r.Close()
|
||||
if !strings.Contains(buf.String(), "d05-injected") {
|
||||
t.Fatalf("version output=%q", buf.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestGoBuildInjectsVersion(t *testing.T) {
|
||||
_, thisFile, _, ok := runtime.Caller(0)
|
||||
if !ok {
|
||||
t.Fatal("runtime.Caller")
|
||||
}
|
||||
pkgDir := filepath.Dir(thisFile)
|
||||
out := filepath.Join(t.TempDir(), "nixmsg-d05-version.exe")
|
||||
build := exec.Command("go", "build", "-ldflags", "-X main.Version=d05-ldflags", "-o", out)
|
||||
build.Dir = pkgDir
|
||||
build.Env = append(os.Environ(), "CGO_ENABLED=0")
|
||||
if b, err := build.CombinedOutput(); err != nil {
|
||||
t.Fatalf("go build: %v\n%s", err, b)
|
||||
}
|
||||
run := exec.Command(out, "version")
|
||||
got, err := run.CombinedOutput()
|
||||
if err != nil {
|
||||
t.Fatalf("version: %v\n%s", err, got)
|
||||
}
|
||||
if !strings.Contains(string(got), "d05-ldflags") {
|
||||
t.Fatalf("version output=%q", got)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"crypto/tls"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
@@ -25,8 +26,15 @@ func cmdHealthcheck(_ []string) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
url := "http://" + addr + "/healthz"
|
||||
client := &http.Client{Timeout: 3 * time.Second}
|
||||
scheme := "http"
|
||||
if healthcheckUseHTTPS(cfg) {
|
||||
scheme = "https"
|
||||
client.Transport = &http.Transport{
|
||||
TLSClientConfig: &tls.Config{InsecureSkipVerify: true}, // 本机 HEALTHCHECK,自签证书可接受
|
||||
}
|
||||
}
|
||||
url := scheme + "://" + addr + "/healthz"
|
||||
resp, err := client.Get(url)
|
||||
if err != nil {
|
||||
return fmt.Errorf("healthcheck %s: %w", url, err)
|
||||
@@ -39,6 +47,12 @@ func cmdHealthcheck(_ []string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func healthcheckUseHTTPS(cfg config.Config) bool {
|
||||
cert := strings.TrimSpace(cfg.TLS.CertFile)
|
||||
key := strings.TrimSpace(cfg.TLS.KeyFile)
|
||||
return cert != "" && key != "" && !cfg.TLS.AllowPlaintext
|
||||
}
|
||||
|
||||
func resolveHealthAddr(dataDir, listen string) (string, error) {
|
||||
path := filepath.Join(dataDir, "listen.addr")
|
||||
if b, err := os.ReadFile(path); err == nil {
|
||||
|
||||
@@ -54,3 +54,34 @@ func TestCmdHealthcheckOK(t *testing.T) {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCmdHealthcheckHTTPS(t *testing.T) {
|
||||
srv := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write([]byte("ok"))
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
dir := t.TempDir()
|
||||
hostPort := srv.Listener.Addr().String()
|
||||
if err := os.WriteFile(filepath.Join(dir, "listen.addr"), []byte(hostPort+"\n"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
cert := filepath.Join(dir, "cert.pem")
|
||||
key := filepath.Join(dir, "key.pem")
|
||||
if err := os.WriteFile(cert, []byte("dummy"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(key, []byte("dummy"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
cfgPath := filepath.Join(dir, "config.yaml")
|
||||
body := "listen: \":0\"\ndata_dir: \"" + filepath.ToSlash(dir) + "\"\ntls:\n cert_file: \"" + filepath.ToSlash(cert) + "\"\n key_file: \"" + filepath.ToSlash(key) + "\"\n allow_plaintext: false\n"
|
||||
if err := os.WriteFile(cfgPath, []byte(body), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Setenv("NIXMSG_CONFIG", cfgPath)
|
||||
if err := cmdHealthcheck(nil); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
+2
-1
@@ -13,13 +13,14 @@ RUN pnpm build
|
||||
FROM golang:1.27-bookworm AS build
|
||||
ARG TARGETOS=linux
|
||||
ARG TARGETARCH=amd64
|
||||
ARG VERSION=dev
|
||||
WORKDIR /src
|
||||
COPY go.mod go.sum ./
|
||||
RUN go mod download
|
||||
COPY . .
|
||||
COPY --from=web /src/web/dist ./web/dist
|
||||
ENV CGO_ENABLED=0
|
||||
RUN GOOS=$TARGETOS GOARCH=$TARGETARCH go build -tags embeddist -o /out/nixmsg ./cmd/nixmsg
|
||||
RUN GOOS=$TARGETOS GOARCH=$TARGETARCH go build -tags embeddist -ldflags "-X main.Version=${VERSION}" -o /out/nixmsg ./cmd/nixmsg
|
||||
|
||||
FROM gcr.io/distroless/static:nonroot
|
||||
COPY --from=build /out/nixmsg /nixmsg
|
||||
|
||||
@@ -297,6 +297,49 @@
|
||||
- 备选方案:更短前缀或 HistogramVec。
|
||||
- 影响:仪表盘按上述名字配置。
|
||||
|
||||
### 复审修复 D-01 2026-09-30
|
||||
|
||||
1. **写队列 busy 后恢复 IsReady**
|
||||
- 原条款:DEVELOPMENT 4.3 `/readyz` 已能读写数据库;DEVIATIONS P2 只写失败时 `IsReady()=false`。
|
||||
- 实际做法:`runBatch` 提交成功后在锁内 `ready=true` 并清除 `lastWriteErr`。
|
||||
- 原因:短暂 busy_timeout / 磁盘瞬时错误后不应永久 503。
|
||||
- 备选方案:要求最近 30 秒无失败才恢复(防抖动)。
|
||||
- 影响:`/readyz` 在随后一次成功写入后恢复。
|
||||
|
||||
### 复审修复 D-02 2026-09-30
|
||||
|
||||
1. **写队列关闭安全与非事务执行**
|
||||
- 原条款:DEVELOPMENT 7.6 清理后 `PRAGMA wal_checkpoint(TRUNCATE)`;L-03 / C-03 依赖存储层能力。
|
||||
- 实际做法:`Close` 先置 `closed`,再在写锁内关闭数据通道,并发 `Do` 不会向已关闭 channel 发送;关闭后返回已有的 `ErrQueueClosed`。新增 `ExecOnWriter` / `Checkpoint` / `Optimize`,在写 goroutine 上、事务外执行。不改 `message` 包,不在此调用 checkpoint(留给 C-03)。
|
||||
- 原因:避免停机 panic,并为 C-03 提供非事务接口。
|
||||
- 备选方案:只关 stop 通道、永不 close 数据通道。
|
||||
- 影响:L-03 可安全 Close;C-03 可调用 `DB.Checkpoint`。
|
||||
|
||||
### 复审修复 D-04 2026-09-30
|
||||
|
||||
1. **backup 空库与恢复步骤**
|
||||
- 原条款:PRD F22;OPS 第 3 节。
|
||||
- 实际做法:备份前检查 `nixmsg.db` 绝对路径,不存在则报错且不创建数据目录/空库;成功打印源库绝对路径;输出文件 chmod 0600。OPS 示例 `data_dir` 改为绝对路径;恢复改为同时移走 `-wal`/`-shm`。未做 `nixmsg restore` 命令(允许文件清单未含 restore.go)。
|
||||
- 原因:cron 相对路径会在错误位置新建空库并当成功备份。
|
||||
- 备选方案:增加 `restore --from` 命令。
|
||||
- 影响:空目录 backup 失败;运维按 OPS 恢复时不会叠旧 WAL。
|
||||
|
||||
### 复审修复 D-05 2026-09-30
|
||||
|
||||
1. **迁移备份按版本命名**
|
||||
- 原条款:DEVELOPMENT 7.7 迁移前 VACUUM INTO。
|
||||
- 实际做法:`pre-migrate-v{当前}-to-v{目标}.db`,已存在则复用;复制前尽力检查剩余磁盘空间。
|
||||
- 原因:迁移稳定失败时 `restart: unless-stopped` 会写满磁盘。
|
||||
- 备选方案:按秒时间戳并在启动失败时删除本次备份。
|
||||
- 影响:同一版本区间反复失败只保留一份备份。
|
||||
|
||||
2. **构建注入 Version;grace_seconds: 0 按 0 生效**
|
||||
- 原条款:DEVELOPMENT 6.1 `server_version`;11.1 `grace_seconds` Validate `>= 0`。
|
||||
- 实际做法:Taskfile / q.yml / Dockerfile 用 `-ldflags -X main.Version=…`(默认 `git describe`)。去掉 `applyEmptyDefaults` 对数值 0 的回填;显式 `grace_seconds: 0` 表示无宽限(不在 Validate 里报错)。`max_frame_bytes` 上限 786432。`admin set-password` 清空全部 `admin_sessions`。TLS 仅明文关闭时 `healthcheck` 走 HTTPS(本机跳过证书校验)。K-05 是 SDK 打包,与本次 ldflags 无冲突。
|
||||
- 原因:显式 0 被改回 60 会让运维误以为关掉了宽限;hello 的 max_frame 不能超过 broker 包长。
|
||||
- 备选方案:`grace_seconds: 0` 在 Validate 报错,强制至少 1 秒。
|
||||
- 影响:未写该字段仍为默认 60;命令行改密后旧 Cookie 一律 401。
|
||||
|
||||
## 连接 N
|
||||
|
||||
### N1 / N2 2026-09-30
|
||||
|
||||
+12
-5
@@ -42,23 +42,30 @@ Docker 约定:配置 `/etc/nixmsg/config.yaml`,数据 `/data`,证书 `/cer
|
||||
|
||||
程序不做定时备份,用 cron 或 1Panel 计划任务调用:
|
||||
|
||||
配置里的 `data_dir` 必须写成绝对路径。cron 的工作目录通常是家目录,相对路径会指到错误位置,甚至在空目录里新建空库并当成功备份。
|
||||
|
||||
```bash
|
||||
# 宿主机(服务可在运行中)
|
||||
NIXMSG_CONFIG=/path/to/config.yaml nixmsg backup --out /path/to/data/backup/manual-$(date +%Y%m%d).db
|
||||
# 宿主机(服务可在运行中;NIXMSG_CONFIG 与 data_dir 都用绝对路径)
|
||||
NIXMSG_CONFIG=/opt/nixmsg/config.yaml nixmsg backup --out /opt/nixmsg/data/backup/manual-$(date +%Y%m%d).db
|
||||
|
||||
# Docker Compose
|
||||
docker compose -f deploy/docker-compose.yml exec nixmsg /nixmsg backup --out /data/backup/manual.db
|
||||
```
|
||||
|
||||
备份文件含当时未送完的正文,按敏感数据保管。旧备份自行清理。
|
||||
备份文件含当时未送完的正文,按敏感数据保管(文件权限 0600)。旧备份自行清理。库文件不存在时 backup 报错并打印解析后的绝对路径,不会新建空库。
|
||||
|
||||
恢复:停服务,用备份文件替换 `data/nixmsg.db`(或拷到新 `data_dir`),再启动;勿在半迁移状态硬切。
|
||||
恢复:
|
||||
|
||||
1. 停服务。
|
||||
2. 把 `nixmsg.db`、`nixmsg.db-wal`、`nixmsg.db-shm` 三个文件一起移走(不要只替换 `.db`,残留 WAL 会叠到恢复库上)。
|
||||
3. 把备份文件复制成 `nixmsg.db`。
|
||||
4. 启动。勿在半迁移状态硬切。
|
||||
|
||||
## 4. 升级与自动迁移
|
||||
|
||||
1. 换上新二进制或拉新镜像。
|
||||
2. 启动时若有未应用的嵌入迁移版本,会先 `VACUUM INTO` 到
|
||||
`<data_dir>/backup/pre-migrate-<UTC时间>.db`,再执行迁移。
|
||||
`<data_dir>/backup/pre-migrate-v{当前版本}-to-v{目标版本}.db`,再执行迁移。已有同名备份则复用,避免迁移反复失败时写满磁盘。
|
||||
3. 迁移失败则进程退出,不带半新半旧库继续服务;运维可从备份恢复后排查。
|
||||
4. 空库首次建表不会产生迁移前备份。
|
||||
|
||||
|
||||
@@ -8,6 +8,7 @@ require (
|
||||
github.com/prometheus/client_golang v1.24.1
|
||||
go.yaml.in/yaml/v3 v3.0.5
|
||||
golang.org/x/crypto v0.57.0
|
||||
golang.org/x/sys v0.48.0
|
||||
modernc.org/sqlite v1.60.1
|
||||
)
|
||||
|
||||
@@ -25,7 +26,6 @@ require (
|
||||
github.com/prometheus/procfs v0.21.1 // indirect
|
||||
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect
|
||||
github.com/rs/xid v1.4.0 // indirect
|
||||
golang.org/x/sys v0.48.0 // indirect
|
||||
google.golang.org/protobuf v1.36.11 // indirect
|
||||
modernc.org/libc v1.77.1 // indirect
|
||||
modernc.org/mathutil v1.7.1 // indirect
|
||||
|
||||
@@ -12,6 +12,9 @@ import (
|
||||
// MaxBodyBytesCap 是 max_body_bytes 的硬上限(DEVELOPMENT 11.1)。
|
||||
const MaxBodyBytesCap = 262144
|
||||
|
||||
// MaxFrameBytesCap 是 max_frame_bytes 的硬上限,与 broker MaximumPacketSize 一致。
|
||||
const MaxFrameBytesCap = 786432
|
||||
|
||||
// Config 对应 DEVELOPMENT 第 11.1 节的配置文件。
|
||||
type Config struct {
|
||||
Listen string `yaml:"listen"`
|
||||
@@ -124,39 +127,8 @@ func applyEmptyDefaults(cfg *Config) {
|
||||
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
|
||||
}
|
||||
// 数值字段不在这里把 0 改回默认:Load 已先填 Default 再解析 YAML,
|
||||
// 未写的字段保留默认值;显式写 0 对 grace 等字段有意义,其余由 Validate 拒绝。
|
||||
}
|
||||
|
||||
// Validate 校验 DEVELOPMENT 11.1 全部字段;拒绝 max_body_bytes > 262144。
|
||||
@@ -189,6 +161,9 @@ func (c Config) Validate() error {
|
||||
if c.Limits.MaxFrameBytes <= 0 {
|
||||
errs = append(errs, "limits.max_frame_bytes must be > 0")
|
||||
}
|
||||
if c.Limits.MaxFrameBytes > MaxFrameBytesCap {
|
||||
errs = append(errs, fmt.Sprintf("limits.max_frame_bytes must be <= %d", MaxFrameBytesCap))
|
||||
}
|
||||
if c.Limits.MaxFrameBytes < c.Limits.MaxBodyBytes {
|
||||
errs = append(errs, "limits.max_frame_bytes must be >= limits.max_body_bytes")
|
||||
}
|
||||
|
||||
@@ -48,6 +48,62 @@ func TestLoadAndValidate(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadExplicitZeroGraceSeconds(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) + "\"\nlimits:\n grace_seconds: 0\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.Limits.GraceSeconds != 0 {
|
||||
t.Fatalf("grace_seconds=%d, want 0", cfg.Limits.GraceSeconds)
|
||||
}
|
||||
if cfg.Limits.MaxBodyBytes != MaxBodyBytesCap {
|
||||
t.Fatalf("omitted max_body_bytes=%d", cfg.Limits.MaxBodyBytes)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateRejectsMaxFrameBytesTooLarge(t *testing.T) {
|
||||
t.Parallel()
|
||||
cfg := Default()
|
||||
cfg.Limits.MaxFrameBytes = MaxFrameBytesCap + 1
|
||||
err := cfg.Validate()
|
||||
if err == nil {
|
||||
t.Fatal("expected error")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "max_frame_bytes") {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadExplicitZeroMaxBodyBytesRejected(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) + "\"\nlimits:\n max_body_bytes: 0\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 cfg.Limits.MaxBodyBytes != 0 {
|
||||
t.Fatalf("max_body_bytes=%d, want explicit 0", cfg.Limits.MaxBodyBytes)
|
||||
}
|
||||
if err := cfg.Validate(); err == nil || !strings.Contains(err.Error(), "max_body_bytes") {
|
||||
t.Fatalf("want max_body_bytes error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateTrustedProxies(t *testing.T) {
|
||||
t.Parallel()
|
||||
cfg := Default()
|
||||
|
||||
+28
-5
@@ -8,7 +8,13 @@ import (
|
||||
"time"
|
||||
)
|
||||
|
||||
const settingAdminPasswordHash = "admin_password_hash"
|
||||
const (
|
||||
settingAdminPasswordHash = "admin_password_hash"
|
||||
adminPasswordSQL = `
|
||||
INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?)
|
||||
ON CONFLICT(key) DO UPDATE SET value = excluded.value, updated_at = excluded.updated_at
|
||||
`
|
||||
)
|
||||
|
||||
// HasAdminPassword 检查 settings 中是否已有管理员密码哈希。
|
||||
func HasAdminPassword(ctx context.Context, db *sql.DB) (bool, error) {
|
||||
@@ -26,13 +32,30 @@ func HasAdminPassword(ctx context.Context, db *sql.DB) (bool, error) {
|
||||
// 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)
|
||||
_, err := db.ExecContext(ctx, adminPasswordSQL, settingAdminPasswordHash, phc, now)
|
||||
return err
|
||||
}
|
||||
|
||||
func setAdminPasswordHashTx(tx *sql.Tx, phc string) error {
|
||||
now := time.Now().UnixMilli()
|
||||
_, err := tx.Exec(adminPasswordSQL, settingAdminPasswordHash, phc, now)
|
||||
return err
|
||||
}
|
||||
|
||||
// ResetAdminPassword 更新管理员密码哈希并清空全部管理会话。
|
||||
func (d *DB) ResetAdminPassword(ctx context.Context, phc string) error {
|
||||
if d == nil || d.Queue == nil {
|
||||
return fmt.Errorf("store: not open")
|
||||
}
|
||||
return d.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
if err := setAdminPasswordHashTx(tx, phc); err != nil {
|
||||
return err
|
||||
}
|
||||
_, err := tx.Exec(`DELETE FROM admin_sessions`)
|
||||
return err
|
||||
})
|
||||
}
|
||||
|
||||
// VacuumInto 对打开的写连接执行 VACUUM INTO(可用于运行中备份)。
|
||||
func VacuumInto(ctx context.Context, db *sql.DB, outPath string) error {
|
||||
if outPath == "" {
|
||||
|
||||
+36
-3
@@ -1,6 +1,7 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"os"
|
||||
@@ -10,7 +11,8 @@ import (
|
||||
_ "modernc.org/sqlite"
|
||||
)
|
||||
|
||||
const dbFileName = "nixmsg.db"
|
||||
// DBFileName 是数据目录下的主库文件名。
|
||||
const DBFileName = "nixmsg.db"
|
||||
|
||||
// DB 持有读写连接与写入队列。
|
||||
type DB struct {
|
||||
@@ -25,7 +27,7 @@ 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)
|
||||
dbPath := filepath.Join(dataDir, DBFileName)
|
||||
existed, err := fileExists(dbPath)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -69,11 +71,42 @@ func (d *DB) Close() error {
|
||||
return first
|
||||
}
|
||||
|
||||
// Checkpoint 在写 goroutine 上执行 wal_checkpoint(TRUNCATE)。
|
||||
func (d *DB) Checkpoint(ctx context.Context) error {
|
||||
if d == nil || d.Queue == nil {
|
||||
return fmt.Errorf("store: not open")
|
||||
}
|
||||
return d.Queue.Checkpoint(ctx)
|
||||
}
|
||||
|
||||
// OpenWriter 按 DEVELOPMENT 7.7 打开写连接(带 _txlock=immediate,MaxOpenConns=1)。
|
||||
func OpenWriter(dataDir, synchronous string) (*sql.DB, error) {
|
||||
return openWriter(dataDir, synchronous, true)
|
||||
}
|
||||
|
||||
// OpenExistingWriter 打开已有库的写连接,不创建数据目录;库文件不存在时返回错误。
|
||||
func OpenExistingWriter(dataDir, synchronous string) (*sql.DB, error) {
|
||||
dbPath := filepath.Join(dataDir, DBFileName)
|
||||
ok, err := fileExists(dbPath)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !ok {
|
||||
abs, absErr := filepath.Abs(dbPath)
|
||||
if absErr != nil {
|
||||
abs = dbPath
|
||||
}
|
||||
return nil, fmt.Errorf("database not found: %s", abs)
|
||||
}
|
||||
return openWriter(dataDir, synchronous, false)
|
||||
}
|
||||
|
||||
func openWriter(dataDir, synchronous string, mkdir bool) (*sql.DB, error) {
|
||||
if mkdir {
|
||||
if err := os.MkdirAll(dataDir, 0o755); err != nil {
|
||||
return nil, fmt.Errorf("mkdir data_dir: %w", err)
|
||||
}
|
||||
}
|
||||
dsn, err := buildDSN(dataDir, synchronous, true)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -116,7 +149,7 @@ func buildDSN(dataDir, synchronous string, writer bool) (string, error) {
|
||||
if sync != "FULL" && sync != "NORMAL" {
|
||||
return "", fmt.Errorf("invalid sqlite_synchronous: %s", synchronous)
|
||||
}
|
||||
dbPath := filepath.ToSlash(filepath.Join(dataDir, dbFileName))
|
||||
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,
|
||||
|
||||
@@ -5,7 +5,6 @@ import (
|
||||
"database/sql"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
@@ -94,8 +93,10 @@ CREATE TABLE schema_migrations (
|
||||
t.Fatalf("backup dir missing: %v", err)
|
||||
}
|
||||
found := false
|
||||
var names []string
|
||||
for _, e := range entries {
|
||||
if strings.HasPrefix(e.Name(), "pre-migrate-") && strings.HasSuffix(e.Name(), ".db") {
|
||||
names = append(names, e.Name())
|
||||
if e.Name() == "pre-migrate-v1-to-v2.db" {
|
||||
found = true
|
||||
info, statErr := e.Info()
|
||||
if statErr != nil {
|
||||
@@ -107,7 +108,7 @@ CREATE TABLE schema_migrations (
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Fatal("expected pre-migrate-*.db backup")
|
||||
t.Fatalf("expected pre-migrate-v1-to-v2.db, got %v", names)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,9 @@
|
||||
//go:build !windows && !unix
|
||||
|
||||
package store
|
||||
|
||||
import "fmt"
|
||||
|
||||
func availableBytes(_ string) (uint64, error) {
|
||||
return 0, fmt.Errorf("disk space probe unsupported on this platform")
|
||||
}
|
||||
@@ -0,0 +1,13 @@
|
||||
//go:build unix
|
||||
|
||||
package store
|
||||
|
||||
import "golang.org/x/sys/unix"
|
||||
|
||||
func availableBytes(path string) (uint64, error) {
|
||||
var st unix.Statfs_t
|
||||
if err := unix.Statfs(path, &st); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return uint64(st.Bavail) * uint64(st.Bsize), nil
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
//go:build windows
|
||||
|
||||
package store
|
||||
|
||||
import "golang.org/x/sys/windows"
|
||||
|
||||
func availableBytes(path string) (uint64, error) {
|
||||
p, err := windows.UTF16PtrFromString(path)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
var free, total, totalFree uint64
|
||||
if err := windows.GetDiskFreeSpaceEx(p, &free, &total, &totalFree); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return free, nil
|
||||
}
|
||||
@@ -18,9 +18,22 @@ var migrationFS embed.FS
|
||||
|
||||
// Migrate 应用尚未执行的嵌入迁移。
|
||||
// 若 dbExisted 为 true(调用 Open/Migrate 前已有 nixmsg.db)且存在未应用版本,
|
||||
// 先 VACUUM INTO <data_dir>/backup/pre-migrate-<时间>.db,再迁移。
|
||||
// 先 VACUUM INTO <data_dir>/backup/pre-migrate-v{当前}-to-v{目标}.db,再迁移。
|
||||
// 备份文件已存在则复用,避免迁移稳定失败时反复全量复制。
|
||||
// 任一版本失败则返回错误,调用方不得继续带半新半旧库提供服务。
|
||||
func Migrate(db *sql.DB, dataDir string, dbExisted bool) error {
|
||||
if err := ensureMigrationsTable(db); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
pending, err := pendingMigrations(db)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return applyPending(db, dataDir, dbExisted, pending)
|
||||
}
|
||||
|
||||
func ensureMigrationsTable(db *sql.DB) error {
|
||||
if _, err := db.Exec(`
|
||||
CREATE TABLE IF NOT EXISTS schema_migrations (
|
||||
version INTEGER PRIMARY KEY,
|
||||
@@ -28,17 +41,21 @@ CREATE TABLE IF NOT EXISTS schema_migrations (
|
||||
)`); err != nil {
|
||||
return fmt.Errorf("ensure schema_migrations: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
pending, err := pendingMigrations(db)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
func applyPending(db *sql.DB, dataDir string, dbExisted bool, pending []migrationFile) error {
|
||||
if len(pending) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
if dbExisted {
|
||||
if err := backupBeforeMigrate(db, dataDir); err != nil {
|
||||
from, err := currentSchemaVersion(db)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
to := pending[len(pending)-1].version
|
||||
if err := backupBeforeMigrate(db, dataDir, from, to); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
@@ -98,13 +115,33 @@ func pendingMigrations(db *sql.DB) ([]migrationFile, error) {
|
||||
return pending, nil
|
||||
}
|
||||
|
||||
func backupBeforeMigrate(db *sql.DB, dataDir string) error {
|
||||
func currentSchemaVersion(db *sql.DB) (int, error) {
|
||||
var v sql.NullInt64
|
||||
if err := db.QueryRow(`SELECT MAX(version) FROM schema_migrations`).Scan(&v); err != nil {
|
||||
return 0, fmt.Errorf("current schema version: %w", err)
|
||||
}
|
||||
if !v.Valid {
|
||||
return 0, nil
|
||||
}
|
||||
return int(v.Int64), nil
|
||||
}
|
||||
|
||||
func backupBeforeMigrate(db *sql.DB, dataDir string, fromVer, toVer int) 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")
|
||||
backupPath := filepath.Join(backupDir, fmt.Sprintf("pre-migrate-v%d-to-v%d.db", fromVer, toVer))
|
||||
exists, err := fileExists(backupPath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if exists {
|
||||
return nil
|
||||
}
|
||||
if err := ensureDiskSpace(backupDir, backupNeedBytes(dataDir)); err != nil {
|
||||
return err
|
||||
}
|
||||
// SQLite VACUUM INTO 需要字面量路径;统一用斜杠,并对单引号转义。
|
||||
quoted := strings.ReplaceAll(filepath.ToSlash(backupPath), "'", "''")
|
||||
if _, err := db.Exec("VACUUM INTO '" + quoted + "'"); err != nil {
|
||||
@@ -113,6 +150,34 @@ func backupBeforeMigrate(db *sql.DB, dataDir string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func backupNeedBytes(dataDir string) int64 {
|
||||
var n int64
|
||||
for _, name := range []string{DBFileName, DBFileName + "-wal", DBFileName + "-shm"} {
|
||||
st, err := os.Stat(filepath.Join(dataDir, name))
|
||||
if err == nil {
|
||||
n += st.Size()
|
||||
}
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
func ensureDiskSpace(dir string, need int64) error {
|
||||
if need < 0 {
|
||||
need = 0
|
||||
}
|
||||
avail, err := availableBytes(dir)
|
||||
if err != nil {
|
||||
// 探测失败不阻断迁移,VACUUM INTO 自身会因空间不足报错。
|
||||
return nil
|
||||
}
|
||||
const margin = 1 << 20
|
||||
want := uint64(need) + margin
|
||||
if avail < want {
|
||||
return fmt.Errorf("insufficient disk space for migrate backup: have %d bytes, need ~%d", avail, want)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func applyMigration(db *sql.DB, m migrationFile) error {
|
||||
tx, err := db.Begin()
|
||||
if err != nil {
|
||||
|
||||
@@ -0,0 +1,73 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestFailingMigrationReusesBackup(t *testing.T) {
|
||||
t.Parallel()
|
||||
dir := t.TempDir()
|
||||
|
||||
db, err := Open(dir, "FULL")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_ = db.Close()
|
||||
|
||||
w, err := OpenWriter(dir, "FULL")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer func() { _ = w.Close() }()
|
||||
|
||||
pending := []migrationFile{{
|
||||
version: 9999,
|
||||
name: "9999_fail.sql",
|
||||
body: "THIS IS NOT VALID SQL",
|
||||
}}
|
||||
if err := applyPending(w, dir, true, pending); err == nil {
|
||||
t.Fatal("expected first failing migration to error")
|
||||
}
|
||||
backupDir := filepath.Join(dir, "backup")
|
||||
first, err := os.ReadDir(backupDir)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(first) != 1 || first[0].Name() != "pre-migrate-v2-to-v9999.db" {
|
||||
t.Fatalf("first backups=%v", dirNames(first))
|
||||
}
|
||||
|
||||
if err := applyPending(w, dir, true, pending); err == nil {
|
||||
t.Fatal("expected second failing migration to error")
|
||||
}
|
||||
second, err := os.ReadDir(backupDir)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(second) != 1 || second[0].Name() != first[0].Name() {
|
||||
t.Fatalf("second backups=%v first=%v", dirNames(second), dirNames(first))
|
||||
}
|
||||
info1, err := first[0].Info()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
info2, err := second[0].Info()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if info2.ModTime().After(info1.ModTime().Add(time.Second)) && info2.Size() != info1.Size() {
|
||||
// 复用同一文件即可;时钟精度下允许 mtime 相同。
|
||||
t.Logf("mtime first=%s second=%s", info1.ModTime(), info2.ModTime())
|
||||
}
|
||||
}
|
||||
|
||||
func dirNames(entries []os.DirEntry) []string {
|
||||
out := make([]string, 0, len(entries))
|
||||
for _, e := range entries {
|
||||
out = append(out, e.Name())
|
||||
}
|
||||
return out
|
||||
}
|
||||
+112
-26
@@ -28,6 +28,7 @@ type WriteFunc func(tx *sql.Tx) error
|
||||
type writeJob struct {
|
||||
ctx context.Context
|
||||
fn WriteFunc
|
||||
raw func(*sql.DB) error
|
||||
res chan error
|
||||
}
|
||||
|
||||
@@ -38,6 +39,7 @@ type Queue struct {
|
||||
ch chan writeJob
|
||||
done chan struct{}
|
||||
closed atomic.Bool
|
||||
sendMu sync.RWMutex
|
||||
|
||||
mu sync.Mutex
|
||||
ready bool
|
||||
@@ -65,34 +67,71 @@ func (q *Queue) Do(ctx context.Context, fn WriteFunc) error {
|
||||
if fn == nil {
|
||||
return errors.New("store: nil write func")
|
||||
}
|
||||
if q.closed.Load() {
|
||||
return ErrQueueClosed
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
job := writeJob{ctx: ctx, fn: fn, res: make(chan error, 1)}
|
||||
q.mu.Lock()
|
||||
q.pending++
|
||||
q.mu.Unlock()
|
||||
select {
|
||||
case q.ch <- job:
|
||||
case <-ctx.Done():
|
||||
q.mu.Lock()
|
||||
q.pending--
|
||||
q.mu.Unlock()
|
||||
return ctx.Err()
|
||||
case <-q.done:
|
||||
q.mu.Lock()
|
||||
q.pending--
|
||||
q.mu.Unlock()
|
||||
if err := q.enqueue(job); err != nil {
|
||||
return err
|
||||
}
|
||||
return q.waitResult(ctx, job)
|
||||
}
|
||||
|
||||
// ExecOnWriter 在写 goroutine 上、事务外执行 fn(如 PRAGMA wal_checkpoint)。
|
||||
func (q *Queue) ExecOnWriter(ctx context.Context, fn func(*sql.DB) error) error {
|
||||
if fn == nil {
|
||||
return errors.New("store: nil exec func")
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
job := writeJob{ctx: ctx, raw: fn, res: make(chan error, 1)}
|
||||
if err := q.enqueue(job); err != nil {
|
||||
return err
|
||||
}
|
||||
return q.waitResult(ctx, job)
|
||||
}
|
||||
|
||||
// Checkpoint 在写连接上执行 PRAGMA wal_checkpoint(TRUNCATE)。
|
||||
func (q *Queue) Checkpoint(ctx context.Context) error {
|
||||
return q.ExecOnWriter(ctx, func(db *sql.DB) error {
|
||||
_, err := db.ExecContext(ctx, `PRAGMA wal_checkpoint(TRUNCATE)`)
|
||||
return err
|
||||
})
|
||||
}
|
||||
|
||||
// Optimize 在写连接上执行 PRAGMA optimize。
|
||||
func (q *Queue) Optimize(ctx context.Context) error {
|
||||
return q.ExecOnWriter(ctx, func(db *sql.DB) error {
|
||||
_, err := db.ExecContext(ctx, `PRAGMA optimize`)
|
||||
return err
|
||||
})
|
||||
}
|
||||
|
||||
func (q *Queue) enqueue(job writeJob) error {
|
||||
q.sendMu.RLock()
|
||||
if q.closed.Load() {
|
||||
q.sendMu.RUnlock()
|
||||
return ErrQueueClosed
|
||||
}
|
||||
q.addPending(1)
|
||||
select {
|
||||
case q.ch <- job:
|
||||
q.sendMu.RUnlock()
|
||||
return nil
|
||||
case <-job.ctx.Done():
|
||||
q.addPending(-1)
|
||||
q.sendMu.RUnlock()
|
||||
return job.ctx.Err()
|
||||
}
|
||||
}
|
||||
|
||||
func (q *Queue) waitResult(ctx context.Context, job writeJob) error {
|
||||
select {
|
||||
case err := <-job.res:
|
||||
return err
|
||||
case <-ctx.Done():
|
||||
// 操作可能仍在队列中执行;结果通道仍会被写端关闭式填入。
|
||||
// 操作可能仍在队列中执行;结果通道仍会被写端填入。
|
||||
select {
|
||||
case err := <-job.res:
|
||||
if err != nil {
|
||||
@@ -112,8 +151,13 @@ func (q *Queue) loop() {
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if job.raw != nil {
|
||||
q.runRaw(job)
|
||||
continue
|
||||
}
|
||||
batch := []writeJob{job}
|
||||
timer := time.NewTimer(batchWait)
|
||||
ranRaw := false
|
||||
collect:
|
||||
for len(batch) < maxBatchOps {
|
||||
select {
|
||||
@@ -121,25 +165,47 @@ func (q *Queue) loop() {
|
||||
if !ok {
|
||||
break collect
|
||||
}
|
||||
if j.raw != nil {
|
||||
stopTimer(timer)
|
||||
q.runBatch(batch)
|
||||
q.runRaw(j)
|
||||
ranRaw = true
|
||||
break collect
|
||||
}
|
||||
batch = append(batch, j)
|
||||
case <-timer.C:
|
||||
break collect
|
||||
}
|
||||
}
|
||||
timer.Stop()
|
||||
if !ranRaw {
|
||||
stopTimer(timer)
|
||||
q.runBatch(batch)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func stopTimer(timer *time.Timer) {
|
||||
if !timer.Stop() {
|
||||
select {
|
||||
case <-timer.C:
|
||||
default:
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (q *Queue) runRaw(job writeJob) {
|
||||
defer q.addPending(-1)
|
||||
if err := job.ctx.Err(); err != nil {
|
||||
job.res <- err
|
||||
return
|
||||
}
|
||||
job.res <- job.raw(q.db)
|
||||
}
|
||||
|
||||
func (q *Queue) runBatch(batch []writeJob) {
|
||||
started := time.Now()
|
||||
defer func() {
|
||||
q.mu.Lock()
|
||||
q.pending -= len(batch)
|
||||
if q.pending < 0 {
|
||||
q.pending = 0
|
||||
}
|
||||
q.mu.Unlock()
|
||||
q.addPending(-len(batch))
|
||||
}()
|
||||
|
||||
// 过滤已取消的任务。
|
||||
@@ -233,6 +299,7 @@ func (q *Queue) runBatch(batch []writeJob) {
|
||||
}
|
||||
return
|
||||
}
|
||||
q.markReady()
|
||||
if q.OnBatchCommit != nil {
|
||||
q.OnBatchCommit(time.Since(started))
|
||||
}
|
||||
@@ -252,7 +319,23 @@ func (q *Queue) markBusy(err error) {
|
||||
q.lastWriteErr = err
|
||||
}
|
||||
|
||||
// IsReady 写库是否仍可用(写失败后为 false)。
|
||||
func (q *Queue) markReady() {
|
||||
q.mu.Lock()
|
||||
defer q.mu.Unlock()
|
||||
q.ready = true
|
||||
q.lastWriteErr = nil
|
||||
}
|
||||
|
||||
func (q *Queue) addPending(delta int) {
|
||||
q.mu.Lock()
|
||||
defer q.mu.Unlock()
|
||||
q.pending += delta
|
||||
if q.pending < 0 {
|
||||
q.pending = 0
|
||||
}
|
||||
}
|
||||
|
||||
// IsReady 写库是否仍可用(写失败后为 false;随后一次成功提交会恢复)。
|
||||
func (q *Queue) IsReady() bool {
|
||||
q.mu.Lock()
|
||||
defer q.mu.Unlock()
|
||||
@@ -293,11 +376,14 @@ func (q *Queue) Drain(ctx context.Context) error {
|
||||
}
|
||||
|
||||
// Close 关闭队列:不再接受新任务,并等待写 goroutine 处理完已入队任务后退出。
|
||||
// 在写锁内关闭数据通道,避免并发 Do 向已关闭 channel 发送而 panic。
|
||||
func (q *Queue) Close() error {
|
||||
if q.closed.Swap(true) {
|
||||
return nil
|
||||
}
|
||||
q.sendMu.Lock()
|
||||
close(q.ch)
|
||||
q.sendMu.Unlock()
|
||||
<-q.done
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -5,6 +5,9 @@ import (
|
||||
"database/sql"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
@@ -148,6 +151,119 @@ func TestWriteQueueNoBacklogAt200PerSec(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteQueueRecoversReadyAfterBusy(t *testing.T) {
|
||||
t.Parallel()
|
||||
dir := t.TempDir()
|
||||
db, err := Open(dir, "FULL")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer func() { _ = db.Close() }()
|
||||
|
||||
db.Queue.markBusy(errors.New("injected busy"))
|
||||
if db.Queue.IsReady() {
|
||||
t.Fatal("expected not ready after markBusy")
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
if err := db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
_, e := tx.Exec(
|
||||
`INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?)`,
|
||||
"after_busy", "1", time.Now().UnixMilli(),
|
||||
)
|
||||
return e
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !db.Queue.IsReady() {
|
||||
t.Fatalf("expected ready after successful Do, last=%v", db.Queue.LastWriteError())
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteQueueCloseConcurrentDoNoPanic(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()
|
||||
const n = 1000
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(n)
|
||||
for i := 0; i < n; i++ {
|
||||
i := i
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
_ = db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
_, e := tx.Exec(
|
||||
`INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?)
|
||||
ON CONFLICT(key) DO UPDATE SET value = excluded.value, updated_at = excluded.updated_at`,
|
||||
fmt.Sprintf("close_%d", i%50), "1", time.Now().UnixMilli(),
|
||||
)
|
||||
return e
|
||||
})
|
||||
}()
|
||||
}
|
||||
time.Sleep(2 * time.Millisecond)
|
||||
if err := db.Queue.Close(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
err = db.Queue.Do(ctx, func(tx *sql.Tx) error { return nil })
|
||||
if !errors.Is(err, ErrQueueClosed) {
|
||||
t.Fatalf("want ErrQueueClosed, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestQueueCheckpointTruncatesWAL(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()
|
||||
payload := strings.Repeat("x", 4096)
|
||||
for i := 0; i < 300; i++ {
|
||||
i := i
|
||||
if err := db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
_, e := tx.Exec(
|
||||
`INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?)`,
|
||||
fmt.Sprintf("wal_%d", i), payload, time.Now().UnixMilli(),
|
||||
)
|
||||
return e
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
walPath := filepath.Join(dir, DBFileName+"-wal")
|
||||
before, statErr := os.Stat(walPath)
|
||||
if err := db.Checkpoint(ctx); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if statErr != nil {
|
||||
return
|
||||
}
|
||||
after, err := os.Stat(walPath)
|
||||
if err != nil {
|
||||
// TRUNCATE 后 WAL 可能被删掉,视为回落成功。
|
||||
if os.IsNotExist(err) {
|
||||
return
|
||||
}
|
||||
t.Fatal(err)
|
||||
}
|
||||
if before.Size() > 0 && after.Size() >= before.Size() {
|
||||
t.Fatalf("WAL did not shrink: before=%d after=%d", before.Size(), after.Size())
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkWriteQueue200PerSec(b *testing.B) {
|
||||
dir := b.TempDir()
|
||||
db, err := Open(dir, "FULL")
|
||||
|
||||
+7
-5
@@ -4,6 +4,8 @@ vars:
|
||||
IMAGE: git.asio.asia/nixevol/nixmsg
|
||||
IMAGE_TAG: "0.1.0"
|
||||
BIN_DIR: bin
|
||||
VERSION:
|
||||
sh: git describe --tags --always --dirty
|
||||
|
||||
tasks:
|
||||
q:accept:
|
||||
@@ -46,7 +48,7 @@ tasks:
|
||||
q:docker-build:
|
||||
desc: 构建当前架构镜像(不推送)
|
||||
cmds:
|
||||
- docker build -t {{.IMAGE}}:{{.IMAGE_TAG}} -t {{.IMAGE}}:latest -f deploy/Dockerfile .
|
||||
- docker build --build-arg VERSION={{.VERSION}} -t {{.IMAGE}}:{{.IMAGE_TAG}} -t {{.IMAGE}}:latest -f deploy/Dockerfile .
|
||||
|
||||
q:docker-push:
|
||||
desc: 构建并推送当前架构镜像到 git.asio.asia/nixevol/nixmsg(正式推送在 Z3)
|
||||
@@ -58,7 +60,7 @@ tasks:
|
||||
q:docker-buildx:
|
||||
desc: 多架构 buildx 构建并推送 linux/amd64+arm64(需 QEMU;正式推送在 Z3)
|
||||
cmds:
|
||||
- docker buildx build --platform linux/amd64,linux/arm64 -t {{.IMAGE}}:{{.IMAGE_TAG}} -t {{.IMAGE}}:latest -f deploy/Dockerfile --push .
|
||||
- docker buildx build --platform linux/amd64,linux/arm64 --build-arg VERSION={{.VERSION}} -t {{.IMAGE}}:{{.IMAGE_TAG}} -t {{.IMAGE}}:latest -f deploy/Dockerfile --push .
|
||||
|
||||
q:release-bins:
|
||||
desc: 交叉编译三平台二进制到 bin/(CGO_ENABLED=0,不提交)
|
||||
@@ -66,17 +68,17 @@ tasks:
|
||||
cmds:
|
||||
- |
|
||||
{{if eq OS "windows"}}powershell -NoProfile -Command "New-Item -ItemType Directory -Force -Path '{{.BIN_DIR}}' | Out-Null"{{else}}mkdir -p {{.BIN_DIR}}{{end}}
|
||||
- cmd: go build -tags embeddist -o {{.BIN_DIR}}/nixmsg-linux-amd64 ./cmd/nixmsg
|
||||
- cmd: go build -tags embeddist -ldflags "-X main.Version={{.VERSION}}" -o {{.BIN_DIR}}/nixmsg-linux-amd64 ./cmd/nixmsg
|
||||
env:
|
||||
CGO_ENABLED: "0"
|
||||
GOOS: linux
|
||||
GOARCH: amd64
|
||||
- cmd: go build -tags embeddist -o {{.BIN_DIR}}/nixmsg-linux-arm64 ./cmd/nixmsg
|
||||
- cmd: go build -tags embeddist -ldflags "-X main.Version={{.VERSION}}" -o {{.BIN_DIR}}/nixmsg-linux-arm64 ./cmd/nixmsg
|
||||
env:
|
||||
CGO_ENABLED: "0"
|
||||
GOOS: linux
|
||||
GOARCH: arm64
|
||||
- cmd: go build -tags embeddist -o {{.BIN_DIR}}/nixmsg-windows-amd64.exe ./cmd/nixmsg
|
||||
- cmd: go build -tags embeddist -ldflags "-X main.Version={{.VERSION}}" -o {{.BIN_DIR}}/nixmsg-windows-amd64.exe ./cmd/nixmsg
|
||||
env:
|
||||
CGO_ENABLED: "0"
|
||||
GOOS: windows
|
||||
|
||||
Reference in New Issue
Block a user