后端基础层:模块定义、配置解析、SQLite 存储层、Argon2id 认证与 Markdown 渲染

- internal/config:serve/init/backup/gc 子命令参数解析(--dev 强制 loopback 守卫)
- internal/store:user_version 版本化迁移(仅追加式 + 越界拒启 + --allow-newer)、
  笔记/图片/引用/会话/设置 DAO、并集可见性查询、VACUUM INTO 在线备份、
  回收站 30 天 + 孤儿图 7 天宽限 gc
- internal/auth:Argon2id PHC 串(m=19456,t=2,p=1)、256bit token 与 SHA-256 摘要
- internal/markdown:goldmark(默认转义)+ bluemonday 双保险,摘要纯文本提取
- internal/middleware:安全头(CSP/HSTS/nosniff 等)、错误日志、
  有界令牌桶限流(per-IP 桶上限 + TTL 逐出)、Origin/Referer 同源校验
This commit is contained in:
2026-09-08 08:14:12 +08:00
parent 685f628e26
commit 247e88c4fb
17 changed files with 1786 additions and 0 deletions
+96
View File
@@ -0,0 +1,96 @@
// Package auth 提供 Argon2id 口令哈希(PHC 串)、随机 token 与 token 摘要。
package auth
import (
"crypto/rand"
"crypto/sha256"
"crypto/subtle"
"encoding/base64"
"encoding/hex"
"errors"
"fmt"
"strings"
"golang.org/x/crypto/argon2"
)
// OWASP 现行推荐 Argon2id 参数:m=19MiB、t=2、p=1(设计 §3.1/§7.3)。
const (
ArgonTime = 2
ArgonMemory = 19456
ArgonThreads = 1
SaltLen = 16
KeyLen = 32
// MinPasswordLength 新口令最小长度(§7.1 改密)。
MinPasswordLength = 12
)
// HashPassword 生成 PHC 格式串:
// $argon2id$v=19$m=19456,t=2,p=1$<b64salt>$<b64hash>(参数随哈希走)。
func HashPassword(password string) (string, error) {
salt := make([]byte, SaltLen)
if _, err := rand.Read(salt); err != nil {
return "", fmt.Errorf("生成盐失败: %w", err)
}
key := argon2.IDKey([]byte(password), salt, ArgonTime, ArgonMemory, ArgonThreads, KeyLen)
return fmt.Sprintf("$argon2id$v=19$m=%d,t=%d,p=%d$%s$%s",
ArgonMemory, ArgonTime, ArgonThreads,
base64.RawStdEncoding.EncodeToString(salt),
base64.RawStdEncoding.EncodeToString(key),
), nil
}
// VerifyPassword 按 PHC 串内参数重派生并常量时间比较。任何解析失败均返回 false。
func VerifyPassword(encoded, password string) bool {
parts := strings.Split(encoded, "$")
// ["", "argon2id", "v=19", "m=..,t=..,p=..", salt, hash]
if len(parts) != 6 || parts[0] != "" || parts[1] != "argon2id" {
return false
}
var version int
if _, err := fmt.Sscanf(parts[2], "v=%d", &version); err != nil || version != argon2.Version {
return false
}
var m uint32
var t uint32
var p uint8
if _, err := fmt.Sscanf(parts[3], "m=%d,t=%d,p=%d", &m, &t, &p); err != nil {
return false
}
salt, err := base64.RawStdEncoding.DecodeString(parts[4])
if err != nil {
return false
}
want, err := base64.RawStdEncoding.DecodeString(parts[5])
if err != nil {
return false
}
got := argon2.IDKey([]byte(password), salt, t, m, p, uint32(len(want)))
return subtle.ConstantTimeCompare(got, want) == 1
}
// ErrWeakPassword 新口令不满足强度要求。
var ErrWeakPassword = errors.New("密码长度至少 12 个字符")
// CheckPasswordStrength 校验新口令强度。
func CheckPasswordStrength(password string) error {
if len(password) < MinPasswordLength {
return ErrWeakPassword
}
return nil
}
// NewToken 生成 256bit 随机 token(base64url,43 字符)。
func NewToken() (string, error) {
b := make([]byte, 32)
if _, err := rand.Read(b); err != nil {
return "", err
}
return base64.RawURLEncoding.EncodeToString(b), nil
}
// HashToken token 明文的 SHA-256 十六进制摘要(库中只存摘要)。
func HashToken(token string) string {
sum := sha256.Sum256([]byte(token))
return hex.EncodeToString(sum[:])
}
+54
View File
@@ -0,0 +1,54 @@
package auth
import "testing"
func TestPasswordHashRoundtrip(t *testing.T) {
hash, err := HashPassword("correct-horse-12")
if err != nil {
t.Fatal(err)
}
// PHC 串格式与参数(§7.3-1)
if len(hash) < len("$argon2id$v=19$m=19456,t=2,p=1$") ||
hash[:len("$argon2id$v=19$m=19456,t=2,p=1$")] != "$argon2id$v=19$m=19456,t=2,p=1$" {
t.Errorf("PHC 前缀/参数错误: %s", hash)
}
if !VerifyPassword(hash, "correct-horse-12") {
t.Error("正确口令应通过校验")
}
if VerifyPassword(hash, "wrong-password") {
t.Error("错误口令不应通过")
}
// 盐随机:同口令两次哈希不同
hash2, _ := HashPassword("correct-horse-12")
if hash == hash2 {
t.Error("同口令两次哈希应不同(随机盐)")
}
// 非法串安全返回 false
for _, bad := range []string{"", "$argon2id$", "$bcrypt$x$y$z", "$argon2id$v=19$m=1,t=1,p=1$!!$!!"} {
if VerifyPassword(bad, "x") {
t.Errorf("非法串 %q 不应通过", bad)
}
}
}
func TestTokenAndStrength(t *testing.T) {
tok, err := NewToken()
if err != nil {
t.Fatal(err)
}
if len(tok) != 43 {
t.Errorf("256bit base64url token 应为 43 字符,实际 %d", len(tok))
}
if HashToken(tok) == tok {
t.Error("token 摘要不应等于明文")
}
if HashToken(tok) != HashToken(tok) {
t.Error("摘要应确定")
}
if err := CheckPasswordStrength("short12"); err == nil {
t.Error("弱口令应被拒绝")
}
if err := CheckPasswordStrength("strong-enough-12"); err != nil {
t.Errorf("合格口令不应被拒绝: %v", err)
}
}
+87
View File
@@ -0,0 +1,87 @@
// Package config 解析各子命令的命令行开关。
package config
import (
"flag"
"fmt"
"net"
"strings"
)
// Config 运行时配置,serve 与维护子命令共用。
type Config struct {
Addr string
DataDir string
LogLevel string
LogFormat string // text|json
BehindProxy bool
Dev bool
AllowNewer bool
}
// DBPath SQLite 数据库文件路径。
func (c *Config) DBPath() string { return c.DataDir + "/pure-note.db" }
func addCommonFlags(fs *flag.FlagSet, c *Config) {
fs.StringVar(&c.DataDir, "data-dir", "./data", "数据目录(SQLite 数据库所在)")
fs.BoolVar(&c.AllowNewer, "allow-newer", false, "允许在更新的数据库 schema 版本上运行(跳过版本上界守卫)")
}
func addServeFlags(fs *flag.FlagSet, c *Config) {
addCommonFlags(fs, c)
fs.StringVar(&c.Addr, "addr", ":8080", "HTTP 监听地址")
fs.StringVar(&c.LogLevel, "log-level", "info", "日志级别(debug|info|warn|error)")
fs.StringVar(&c.LogFormat, "log-format", "text", "日志格式(text|json)")
fs.BoolVar(&c.BehindProxy, "behind-proxy", false, "位于可信反向代理之后(取 X-Forwarded-For 最右条目作为客户端 IP)")
fs.BoolVar(&c.Dev, "dev", false, "开发模式:允许非 Secure Cookie,仅允许监听 loopback 地址")
}
// ParseServe 解析 `pure-note serve` 参数。
func ParseServe(args []string) (*Config, error) {
c := &Config{}
fs := flag.NewFlagSet("serve", flag.ContinueOnError)
addServeFlags(fs, c)
if err := fs.Parse(args); err != nil {
return nil, err
}
if c.Dev && !isLoopbackAddr(c.Addr) {
return nil, fmt.Errorf("--dev 仅允许监听 loopback 地址(如 127.0.0.1:8080),当前为 %q", c.Addr)
}
return c, nil
}
// ParseInit 解析 `pure-note init` 参数。
func ParseInit(args []string) (*Config, error) {
c := &Config{}
fs := flag.NewFlagSet("init", flag.ContinueOnError)
addCommonFlags(fs, c)
if err := fs.Parse(args); err != nil {
return nil, err
}
return c, nil
}
// ParseMaint 解析 `pure-note backup` / `pure-note gc` 参数(均只需数据目录)。
// 返回 flag 解析后的剩余位置参数(如 backup 的输出路径)。
func ParseMaint(cmd string, args []string) (*Config, []string, error) {
c := &Config{}
fs := flag.NewFlagSet(cmd, flag.ContinueOnError)
addCommonFlags(fs, c)
if err := fs.Parse(args); err != nil {
return nil, nil, err
}
return c, fs.Args(), nil
}
func isLoopbackAddr(addr string) bool {
host, _, err := net.SplitHostPort(addr)
if err != nil {
// 无端口:整体视为主机名
host = addr
}
if host == "" || strings.EqualFold(host, "localhost") {
return host != ""
}
ip := net.ParseIP(host)
return ip != nil && ip.IsLoopback()
}
+23
View File
@@ -0,0 +1,23 @@
package config
import "testing"
func TestDevLoopbackGuard(t *testing.T) {
// --dev 强制 loopback(§7.3-8)
if _, err := ParseServe([]string{"--dev", "--addr", ":8080"}); err == nil {
t.Error("--dev + 0.0.0.0 应拒绝启动")
}
if _, err := ParseServe([]string{"--dev", "--addr", "192.168.1.5:8080"}); err == nil {
t.Error("--dev + 局域网地址应拒绝启动")
}
if _, err := ParseServe([]string{"--dev", "--addr", "127.0.0.1:8080"}); err != nil {
t.Errorf("--dev + 127.0.0.1 应允许: %v", err)
}
if _, err := ParseServe([]string{"--dev", "--addr", "localhost:8080"}); err != nil {
t.Errorf("--dev + localhost 应允许: %v", err)
}
// 非 dev 不限制
if _, err := ParseServe([]string{"--addr", ":8080"}); err != nil {
t.Errorf("非 dev 任意地址应允许: %v", err)
}
}
+54
View File
@@ -0,0 +1,54 @@
// Package markdown 服务端 Markdown 渲染(goldmark + bluemonday),
// 仅用于 RSS / meta 摘要;库中始终只存 Markdown 原文(§7.5)。
package markdown
import (
"bytes"
"html"
"regexp"
"strings"
"unicode/utf8"
"github.com/microcosm-cc/bluemonday"
"github.com/yuin/goldmark"
"github.com/yuin/goldmark/extension"
)
var engine = goldmark.New(
goldmark.WithExtensions(extension.GFM),
// 默认 html.WithUnsafe=false:原始 HTML 与危险 URL 被转义(XSS 安全默认),
// 再经 bluemonday 白名单清洗双保险。
)
// ugcPolicy RSS description 用白名单清洗(双保险,与 goldmark 转义叠加)。
var ugcPolicy = bluemonday.UGCPolicy()
// strictPolicy 摘要提取用:剥掉全部标签。
var strictPolicy = bluemonday.StrictPolicy()
// Render 将 Markdown 渲染为清洗后的 HTML。
func Render(source string) string {
var buf bytes.Buffer
if err := engine.Convert([]byte(source), &buf); err != nil {
// goldmark 对任意输入不产生错误;兜底按纯文本转义。
return html.EscapeString(source)
}
return ugcPolicy.Sanitize(buf.String())
}
var wsRe = regexp.MustCompile(`\s+`)
// Summary 从 Markdown 源提取纯文本摘要:渲染→剥标签→反转义→压空白→截断。
func Summary(source string, limit int) string {
if limit <= 0 {
return ""
}
text := strictPolicy.Sanitize(Render(source))
text = html.UnescapeString(text)
text = strings.TrimSpace(wsRe.ReplaceAllString(text, " "))
if r := utf8.RuneCountInString(text); r > limit {
runes := []rune(text)
return string(runes[:limit]) + "…"
}
return text
}
+33
View File
@@ -0,0 +1,33 @@
package markdown
import (
"strings"
"testing"
)
// 服务端渲染防线(§9.1-T1):goldmark 转义 + bluemonday 白名单。
func TestRenderEscapesDangerousHTML(t *testing.T) {
out := Render("# 标题\n\n<script>alert(1)</script>\n\n[link](javascript:alert(1))\n\n**bold**")
if strings.Contains(out, "<script>") {
t.Errorf("script 应被转义: %s", out)
}
if strings.Contains(out, `href="javascript:`) {
t.Errorf("javascript: 协议应被剥除: %s", out)
}
if !strings.Contains(out, "<strong>bold</strong>") {
t.Errorf("正常 Markdown 应渲染: %s", out)
}
}
func TestSummary(t *testing.T) {
got := Summary("# 标题\n\n这是**正文**内容,包含 <b>HTML</b>。", 10)
if strings.ContainsAny(got, "<>#*") {
t.Errorf("摘要应为纯文本: %q", got)
}
if strings.Count(got, "")-1 > 11 { // 10 runes + 可能的省略号
t.Errorf("摘要应截断: %q", got)
}
if s := Summary("", 100); s != "" {
t.Errorf("空内容摘要应为空: %q", s)
}
}
+136
View File
@@ -0,0 +1,136 @@
// Package middleware 自写中间件链(§7.2):
// SecurityHeaders → 日志 → 全局限流 → Origin 校验 →(路由级)MaxBytes → Auth → CSRF。
package middleware
import (
"context"
"log/slog"
"net"
"net/http"
"strings"
"time"
)
// statusWriter 捕获响应状态码供日志使用。
type statusWriter struct {
http.ResponseWriter
status int
}
func (w *statusWriter) WriteHeader(code int) {
if w.status == 0 {
w.status = code
}
w.ResponseWriter.WriteHeader(code)
}
func (w *statusWriter) Write(b []byte) (int, error) {
if w.status == 0 {
w.status = http.StatusOK
}
return w.ResponseWriter.Write(b)
}
// SecurityHeaders 下发安全 HTTP 头(§9.2)。
func SecurityHeaders(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
h := w.Header()
h.Set("Content-Security-Policy",
"default-src 'self'; "+
"script-src 'self'; "+
// 'unsafe-inline' 仅因 CodeMirror 经 style-mod 运行时注入 <style>(§9.2)。
"style-src 'self' 'unsafe-inline'; "+
"img-src 'self' data:; "+
"font-src 'self'; "+
"connect-src 'self'; "+
"object-src 'none'; base-uri 'none'; "+
"frame-ancestors 'none'; form-action 'self'")
// 站点仅经 HTTPS 反代对外服务(§9.1-T11)
h.Set("Strict-Transport-Security", "max-age=31536000; includeSubDomains")
h.Set("Referrer-Policy", "strict-origin-when-cross-origin")
h.Set("X-Content-Type-Options", "nosniff")
h.Set("X-Frame-Options", "DENY")
h.Set("Permissions-Policy", "camera=(), microphone=(), geolocation=()")
next.ServeHTTP(w, r)
})
}
// RequestLogger 请求日志:仅记录错误(status ≥ 400)。
func RequestLogger(log *slog.Logger) func(http.Handler) http.Handler {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
sw := &statusWriter{ResponseWriter: w}
start := time.Now()
next.ServeHTTP(sw, r)
if sw.status >= 400 {
log.Warn("http",
"method", r.Method,
"path", r.URL.Path,
"status", sw.status,
"ip", ClientIP(r, trueBehindProxy(r)),
"duration_ms", time.Since(start).Milliseconds(),
)
}
})
}
}
// behindProxyKey 由 Server 注入到请求上下文,供日志取 IP 用。
type ctxKey string
const behindProxyKey ctxKey = "behind_proxy"
// BehindProxy 中间件:把「位于可信反代之后」标记注入请求上下文,
// 供日志与限流取客户端 IP 使用。
func BehindProxy(v bool) func(http.Handler) http.Handler {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
next.ServeHTTP(w, WithBehindProxy(r, v))
})
}
}
// WithBehindProxy 标记请求位于可信反代之后。
func WithBehindProxy(r *http.Request, v bool) *http.Request {
return r.WithContext(context.WithValue(r.Context(), behindProxyKey, v))
}
func trueBehindProxy(r *http.Request) bool {
v, _ := r.Context().Value(behindProxyKey).(bool)
return v
}
// ClientIP 提取客户端 IP:behindProxy 时取 X-Forwarded-For 最右条目
// (Caddy 追加语义,§14),否则取 RemoteAddr。
func ClientIP(r *http.Request, behindProxy bool) string {
if behindProxy {
if xff := r.Header.Get("X-Forwarded-For"); xff != "" {
parts := strings.Split(xff, ",")
return strings.TrimSpace(parts[len(parts)-1])
}
}
host, _, err := net.SplitHostPort(r.RemoteAddr)
if err != nil {
return r.RemoteAddr
}
return host
}
// MaxBytes 路由级请求体上限中间件(§7.2)。
func MaxBytes(n int64) func(http.Handler) http.Handler {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
r.Body = http.MaxBytesReader(w, r.Body, n)
next.ServeHTTP(w, r)
})
}
}
// NoStore 为 /api/admin/* 与 /api/auth/* 响应统一附加
// Cache-Control: no-store(防登出后 bfcache 回看,§9.2)。
func NoStore(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Cache-Control", "no-store")
next.ServeHTTP(w, r)
})
}
+157
View File
@@ -0,0 +1,157 @@
package middleware
import (
"net/http"
"net/url"
"strconv"
"strings"
"sync"
"time"
"golang.org/x/time/rate"
)
// OriginCheck 对所有非 GET/HEAD/OPTIONS 请求执行 Origin/Referer→Host 校验
// (§9.1-T2,含 /api/auth/*):Origin 优先,缺省回退 Referer;两者皆缺亦拒绝。
func OriginCheck(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet && r.Method != http.MethodHead && r.Method != http.MethodOptions {
if !sameOrigin(r) {
writeErr(w, http.StatusForbidden, "origin_forbidden", "跨站请求被拒绝(Origin/Referer 校验失败)")
return
}
}
next.ServeHTTP(w, r)
})
}
// sameOrigin 校验 Origin(或 Referer)的 host 与请求 Host 一致。
func sameOrigin(r *http.Request) bool {
raw := r.Header.Get("Origin")
if raw == "" {
raw = r.Header.Get("Referer")
}
if raw == "" {
return false
}
u, err := url.Parse(raw)
if err != nil || u.Host == "" {
return false
}
return strings.EqualFold(u.Host, r.Host)
}
func writeErr(w http.ResponseWriter, status int, code, msg string) {
w.Header().Set("Content-Type", "application/json; charset=utf-8")
w.WriteHeader(status)
// 与 httpapi 错误包络同构;此处避免包循环引用。
_, _ = w.Write([]byte(`{"error":{"code":"` + code + `","message":` + strconv.Quote(msg) + `}}`))
}
// Limiter 有界令牌桶集合:per-key 限流,桶上限 maxEntries,
// 超限时逐出最久未活跃且已耗尽的桶(防伪造 IP 撑爆内存,§7.2/§11)。
type Limiter struct {
mu sync.Mutex
buckets map[string]*bucket
rate rate.Limit
burst int
maxEntries int
ttl time.Duration
lastEvictAt time.Time
}
type bucket struct {
lim *rate.Limiter
lastSeen time.Time
}
// NewLimiter 创建限流器:每秒 rate 个、burst 容量、maxEntries 桶上限、ttl 空闲逐出。
func NewLimiter(perSec float64, burst, maxEntries int, ttl time.Duration) *Limiter {
return &Limiter{
buckets: make(map[string]*bucket),
rate: rate.Limit(perSec),
burst: burst,
maxEntries: maxEntries,
ttl: ttl,
}
}
// get 取(或建)key 的桶。
func (l *Limiter) get(key string) *rate.Limiter {
l.mu.Lock()
defer l.mu.Unlock()
now := time.Now()
// 惰性清理:桶数达上限或距上次清理超过 ttl 时执行
if len(l.buckets) >= l.maxEntries || now.Sub(l.lastEvictAt) > l.ttl {
l.evictLocked(now)
}
b, ok := l.buckets[key]
if !ok {
b = &bucket{lim: rate.NewLimiter(l.rate, l.burst)}
l.buckets[key] = b
}
b.lastSeen = now
return b.lim
}
// evictLocked 逐出过期桶;仍超限时逐出最久未活跃的桶。
func (l *Limiter) evictLocked(now time.Time) {
l.lastEvictAt = now
var oldestKey string
var oldest time.Time
for k, b := range l.buckets {
if now.Sub(b.lastSeen) > l.ttl {
delete(l.buckets, k)
continue
}
if oldestKey == "" || b.lastSeen.Before(oldest) {
oldestKey, oldest = k, b.lastSeen
}
}
for len(l.buckets) >= l.maxEntries && oldestKey != "" {
delete(l.buckets, oldestKey)
oldestKey = ""
for k, b := range l.buckets {
if oldestKey == "" || b.lastSeen.Before(oldest) {
oldestKey, oldest = k, b.lastSeen
}
}
}
}
// Allow 消费 1 个 token;无可用 token 返回 false(不阻塞)。
func (l *Limiter) Allow(key string) bool {
return l.get(key).Allow()
}
// Available 桶内剩余 token 是否 ≥1(登录防爆破的快速预检,不消费)。
func (l *Limiter) Available(key string) bool {
return l.get(key).Tokens() >= 1
}
// Remaining 桶内剩余 token 数(日志观测用,向下取整,不消费)。
func (l *Limiter) Remaining(key string) int {
return int(l.get(key).Tokens())
}
// RetryAfter 距下一个可用 token 的秒数(向上取整,至少 1)。
// 用 Reserve+Cancel 探测延迟,不真正消费 token。
func (l *Limiter) RetryAfter(key string) int {
l.mu.Lock()
b := l.buckets[key]
l.mu.Unlock()
if b == nil {
return 1
}
res := b.lim.Reserve()
d := res.Delay()
res.CancelAt(time.Now())
if d <= 0 {
return 1
}
s := int(d.Seconds())
if s < 1 {
s = 1
}
return s
}
+134
View File
@@ -0,0 +1,134 @@
package store
import (
"database/sql"
"errors"
"regexp"
)
// Image 图片元信息(不含 BLOB)。
type Image struct {
ID int64 `json:"id"`
SHA256 string `json:"sha256"`
MIME string `json:"mime"`
Size int `json:"size"`
CreatedAt int64 `json:"created_at"`
}
// UpsertImage 按 sha256 去重入库,返回图片 id。
func (s *Store) UpsertImage(sha256Hex, mime string, size int, data []byte, createdAt int64) (int64, error) {
if _, err := s.db.Exec(
`INSERT OR IGNORE INTO images (sha256, mime, size, data, created_at) VALUES (?, ?, ?, ?, ?)`,
sha256Hex, mime, size, data, createdAt); err != nil {
return 0, err
}
var id int64
err := s.db.QueryRow(`SELECT id FROM images WHERE sha256=?`, sha256Hex).Scan(&id)
return id, err
}
// GetImageMeta 取图片元信息。
func (s *Store) GetImageMeta(id int64) (*Image, error) {
row := s.db.QueryRow(`SELECT id, sha256, mime, size, created_at FROM images WHERE id=?`, id)
var img Image
err := row.Scan(&img.ID, &img.SHA256, &img.MIME, &img.Size, &img.CreatedAt)
if errors.Is(err, sql.ErrNoRows) {
return nil, ErrNotFound
}
return &img, err
}
// GetImageData 取图片字节。
func (s *Store) GetImageData(id int64) ([]byte, error) {
var data []byte
err := s.db.QueryRow(`SELECT data FROM images WHERE id=?`, id).Scan(&data)
if errors.Is(err, sql.ErrNoRows) {
return nil, ErrNotFound
}
return data, err
}
// ImageIsPublic 并集语义(§6.2):任一引用笔记 public 且未删除 ⇒ 匿名可读。
func (s *Store) ImageIsPublic(id int64) (bool, error) {
var one int
err := s.db.QueryRow(`
SELECT 1 FROM image_refs r JOIN notes n ON n.id = r.note_id
WHERE r.image_id = ? AND n.status='public' AND n.deleted_at IS NULL
LIMIT 1`, id).Scan(&one)
if errors.Is(err, sql.ErrNoRows) {
return false, nil
}
return err == nil, err
}
// imageRefRe 扫描 content 中对 /api/images/{id} 的引用(§6.2 引用维护)。
var imageRefRe = regexp.MustCompile(`/api/images/(\d+)`)
// RebuildImageRefs 在**同一事务**内重建笔记的图片引用:
// 删除该笔记全部 refs → 正则扫描 content → 重建(§6.2)。
func (s *Store) RebuildImageRefs(noteID int64, content string) error {
tx, err := s.db.Begin()
if err != nil {
return err
}
defer tx.Rollback()
if err := rebuildRefsTx(tx, noteID, content); err != nil {
return err
}
return tx.Commit()
}
func rebuildRefsTx(tx *sql.Tx, noteID int64, content string) error {
if _, err := tx.Exec(`DELETE FROM image_refs WHERE note_id=?`, noteID); err != nil {
return err
}
seen := map[int64]bool{}
for _, m := range imageRefRe.FindAllStringSubmatch(content, -1) {
var imgID int64
if _, err := fmtSscanInt(m[1], &imgID); err != nil {
continue
}
if seen[imgID] {
continue
}
seen[imgID] = true
// 图片不存在时忽略(引用悬空无害)
if _, err := tx.Exec(`INSERT OR IGNORE INTO image_refs (image_id, note_id) VALUES (?, ?)`, imgID, noteID); err != nil {
return err
}
}
return nil
}
// ListOrphanImages 零引用图片清单(GET /api/admin/images?orphan=1 与 gc 检视用)。
func (s *Store) ListOrphanImages() ([]Image, error) {
rows, err := s.db.Query(`
SELECT i.id, i.sha256, i.mime, i.size, i.created_at FROM images i
WHERE NOT EXISTS (SELECT 1 FROM image_refs r WHERE r.image_id = i.id)
ORDER BY i.created_at DESC`)
if err != nil {
return nil, err
}
defer rows.Close()
out := []Image{}
for rows.Next() {
var img Image
if err := rows.Scan(&img.ID, &img.SHA256, &img.MIME, &img.Size, &img.CreatedAt); err != nil {
return nil, err
}
out = append(out, img)
}
return out, rows.Err()
}
// DeleteImage 物理删除单张图片(refs 级联)。
func (s *Store) DeleteImage(id int64) error {
res, err := s.db.Exec(`DELETE FROM images WHERE id=?`, id)
if err != nil {
return err
}
if rows, _ := res.RowsAffected(); rows == 0 {
return ErrNotFound
}
return nil
}
+134
View File
@@ -0,0 +1,134 @@
package store
import (
"fmt"
"os"
"time"
)
// GCReport gc 结果报告(dry-run 与 commit 共用)。
type GCReport struct {
ExpiredNotes []ExpiredNote `json:"expired_notes"`
OrphanImages []Image `json:"orphan_images"`
ExpiredSessons int64 `json:"expired_sessions"`
Committed bool `json:"committed"`
}
// ExpiredNote 待物理删除的回收站笔记。
type ExpiredNote struct {
ID int64 `json:"id"`
Slug string `json:"slug"`
DeletedAt int64 `json:"deleted_at"`
}
// TrashTTL 回收站保留期(§6.3:30 天)。
const TrashTTL = 30 * 24 * time.Hour
// OrphanGrace 孤儿图片宽限期(§6.3:7 天,防误删编辑中刚上传的图)。
const OrphanGrace = 7 * 24 * time.Hour
// GC 清理:回收站过期笔记物理删除、孤儿图片清除、过期会话清理。
// dryRun=true 时只输出计划不改数据。
func (s *Store) GC(now time.Time, dryRun bool) (*GCReport, error) {
rep := &GCReport{Committed: !dryRun}
noteCutoff := now.Add(-TrashTTL).Unix()
rows, err := s.db.Query(
`SELECT id, slug, deleted_at FROM notes WHERE deleted_at IS NOT NULL AND deleted_at < ? ORDER BY deleted_at`,
noteCutoff)
if err != nil {
return nil, err
}
for rows.Next() {
var en ExpiredNote
if err := rows.Scan(&en.ID, &en.Slug, &en.DeletedAt); err != nil {
rows.Close()
return nil, err
}
rep.ExpiredNotes = append(rep.ExpiredNotes, en)
}
if err := rows.Err(); err != nil {
rows.Close()
return nil, err
}
rows.Close()
// 孤儿图片:0 引用且超过宽限期
imageCutoff := now.Add(-OrphanGrace).Unix()
orphanRows, err := s.db.Query(`
SELECT i.id, i.sha256, i.mime, i.size, i.created_at FROM images i
WHERE i.created_at < ? AND NOT EXISTS (SELECT 1 FROM image_refs r WHERE r.image_id = i.id)
ORDER BY i.created_at`, imageCutoff)
if err != nil {
return nil, err
}
for orphanRows.Next() {
var img Image
if err := orphanRows.Scan(&img.ID, &img.SHA256, &img.MIME, &img.Size, &img.CreatedAt); err != nil {
orphanRows.Close()
return nil, err
}
rep.OrphanImages = append(rep.OrphanImages, img)
}
if err := orphanRows.Err(); err != nil {
orphanRows.Close()
return nil, err
}
orphanRows.Close()
if dryRun {
// 会话计数(不改数据)
var cnt int64
if err := s.db.QueryRow(`SELECT COUNT(*) FROM sessions WHERE expires_at < ?`, now.Unix()).Scan(&cnt); err != nil {
return nil, err
}
rep.ExpiredSessons = cnt
return rep, nil
}
for _, en := range rep.ExpiredNotes {
if err := s.DeleteNoteForever(en.ID); err != nil {
return nil, fmt.Errorf("物理删除笔记 %d: %w", en.ID, err)
}
}
for _, img := range rep.OrphanImages {
if err := s.DeleteImage(img.ID); err != nil {
return nil, fmt.Errorf("删除孤儿图片 %d: %w", img.ID, err)
}
}
n, err := s.DeleteExpiredSessions(now.Unix())
if err != nil {
return nil, err
}
rep.ExpiredSessons = n
return rep, nil
}
// Backup 在线备份:VACUUM INTO 一致快照(WAL 下与 serve 并发安全,§10.3)。
func (s *Store) Backup(destPath string) error {
if _, err := os.Stat(destPath); err == nil {
return fmt.Errorf("目标文件已存在: %s", destPath)
}
// VACUUM INTO 不接受参数绑定,路径经单引号转义(无参数化通道时的最小注入面)。
if _, err := s.db.Exec("VACUUM INTO " + escapeSQLString(destPath)); err != nil {
return err
}
// 备份含全部私钥内容:强制 0600(§10.3)
if err := os.Chmod(destPath, 0o600); err != nil {
return err
}
return nil
}
func escapeSQLString(s string) string {
out := make([]byte, 0, len(s)+2)
out = append(out, '\'')
for i := 0; i < len(s); i++ {
if s[i] == '\'' {
out = append(out, '\'', '\'')
} else {
out = append(out, s[i])
}
}
return string(append(out, '\''))
}
+361
View File
@@ -0,0 +1,361 @@
package store
import (
"database/sql"
"encoding/json"
"errors"
"fmt"
"strings"
)
// Note 笔记实体。DeletedAt 非 nil 表示处于回收站(软删除)。
type Note struct {
ID int64 `json:"id"`
Slug string `json:"slug"`
Title string `json:"title"`
Summary string `json:"summary"`
Content string `json:"content,omitempty"`
Status string `json:"status"`
Tags []string `json:"tags"`
Pinned bool `json:"pinned"`
DeletedAt *int64 `json:"deleted_at,omitempty"`
CreatedAt int64 `json:"created_at"`
UpdatedAt int64 `json:"updated_at"`
}
// ErrNotFound 统一的「不存在」错误。
var ErrNotFound = errors.New("not found")
const noteColumns = "id, slug, title, summary, content, status, tags, pinned, deleted_at, created_at, updated_at"
func scanNote(scan func(dest ...any) error) (*Note, error) {
var n Note
var tagsJSON string
var pinned int
var content string
if err := scan(&n.ID, &n.Slug, &n.Title, &n.Summary, &content, &n.Status, &tagsJSON, &pinned, &n.DeletedAt, &n.CreatedAt, &n.UpdatedAt); err != nil {
return nil, err
}
n.Content = content
n.Pinned = pinned != 0
if err := json.Unmarshal([]byte(tagsJSON), &n.Tags); err != nil {
n.Tags = []string{}
}
if n.Tags == nil {
n.Tags = []string{}
}
return &n, nil
}
// CreateNote 新建笔记,返回 id。
func (s *Store) CreateNote(n *Note) (int64, error) {
tags, err := marshalTags(n.Tags)
if err != nil {
return 0, err
}
res, err := s.db.Exec(
`INSERT INTO notes (slug, title, summary, content, status, tags, pinned, created_at, updated_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)`,
n.Slug, n.Title, n.Summary, n.Content, n.Status, tags, boolToInt(n.Pinned), n.CreatedAt, n.UpdatedAt)
if err != nil {
return 0, err
}
return res.LastInsertId()
}
// UpdateNote 更新笔记内容字段(不含 slug 冲突裁决,见 httpapi 层)。
func (s *Store) UpdateNote(n *Note) error {
tags, err := marshalTags(n.Tags)
if err != nil {
return err
}
res, err := s.db.Exec(
`UPDATE notes SET title=?, summary=?, content=?, status=?, tags=?, pinned=?, updated_at=? WHERE id=? AND deleted_at IS NULL`,
n.Title, n.Summary, n.Content, n.Status, tags, boolToInt(n.Pinned), n.UpdatedAt, n.ID)
if err != nil {
return err
}
if rows, _ := res.RowsAffected(); rows == 0 {
return ErrNotFound
}
return nil
}
// UpdateNoteSlug 修改 slug(仅正常态笔记)。
func (s *Store) UpdateNoteSlug(id int64, slug string, updatedAt int64) error {
res, err := s.db.Exec(`UPDATE notes SET slug=?, updated_at=? WHERE id=? AND deleted_at IS NULL`, slug, updatedAt, id)
if err != nil {
return err
}
if rows, _ := res.RowsAffected(); rows == 0 {
return ErrNotFound
}
return nil
}
// SoftDeleteNote 置 deleted_at(回收站)。refs 不动(§6.2)。
func (s *Store) SoftDeleteNote(id int64, at int64) error {
res, err := s.db.Exec(`UPDATE notes SET deleted_at=? WHERE id=? AND deleted_at IS NULL`, at, id)
if err != nil {
return err
}
if rows, _ := res.RowsAffected(); rows == 0 {
return ErrNotFound
}
return nil
}
// RestoreNote 清空 deleted_at。
func (s *Store) RestoreNote(id int64) error {
res, err := s.db.Exec(`UPDATE notes SET deleted_at=NULL WHERE id=? AND deleted_at IS NOT NULL`, id)
if err != nil {
return err
}
if rows, _ := res.RowsAffected(); rows == 0 {
return ErrNotFound
}
return nil
}
// GetNoteByID 按 id 取(含回收站)。
func (s *Store) GetNoteByID(id int64) (*Note, error) {
row := s.db.QueryRow(`SELECT `+noteColumns+` FROM notes WHERE id=?`, id)
n, err := scanNote(row.Scan)
if errors.Is(err, sql.ErrNoRows) {
return nil, ErrNotFound
}
return n, err
}
// GetNoteBySlug 按 slug 取(含回收站;可见性裁决在调用方)。
func (s *Store) GetNoteBySlug(slug string) (*Note, error) {
row := s.db.QueryRow(`SELECT `+noteColumns+` FROM notes WHERE slug=?`, slug)
n, err := scanNote(row.Scan)
if errors.Is(err, sql.ErrNoRows) {
return nil, ErrNotFound
}
return n, err
}
// SlugExists slug 是否已被占用(**含回收站**,§6.2:回收期内 slug 仍被占用)。
// excludeID 用于更新场景排除自身;传 0 表示不排除。
func (s *Store) SlugExists(slug string, excludeID int64) (bool, error) {
var one int
err := s.db.QueryRow(`SELECT 1 FROM notes WHERE slug=? AND id<>? LIMIT 1`, slug, excludeID).Scan(&one)
if errors.Is(err, sql.ErrNoRows) {
return false, nil
}
return err == nil, err
}
// PublicNoteFilter 可见性过滤的统一谓词(§9.1-T10 单一可信点)。
const PublicNoteFilter = "status='public' AND deleted_at IS NULL"
// ListPublicNotes 公开笔记分页列表(不含全文)。tag 非空时按标签过滤。
func (s *Store) ListPublicNotes(page, pageSize int, tag string) ([]Note, int, error) {
where := PublicNoteFilter
args := []any{}
if tag != "" {
// json_each 参数化过滤,无拼接注入面。
where += ` AND EXISTS (SELECT 1 FROM json_each(notes.tags) je WHERE je.value = ?)`
args = append(args, tag)
}
var total int
if err := s.db.QueryRow(`SELECT COUNT(*) FROM notes n WHERE `+where, args...).Scan(&total); err != nil {
return nil, 0, err
}
q := `SELECT n.id, n.slug, n.title, n.summary, '' AS content, n.status, n.tags, n.pinned, n.deleted_at, n.created_at, n.updated_at
FROM notes n WHERE ` + where + `
ORDER BY n.pinned DESC, n.updated_at DESC, n.id DESC LIMIT ? OFFSET ?`
rows, err := s.db.Query(q, append(args, pageSize, (page-1)*pageSize)...)
if err != nil {
return nil, 0, err
}
defer rows.Close()
var out []Note
for rows.Next() {
n, err := scanNote(rows.Scan)
if err != nil {
return nil, 0, err
}
out = append(out, *n)
}
return out, total, rows.Err()
}
// AdjacentPublicNote 返回公开序列中与 n 相邻的上一篇/下一篇
// (按列表序 updated_at DESC, id DESC)。仅返回 slug 与 title。
func (s *Store) AdjacentPublicNote(n *Note, dir string) (slug, title string, err error) {
var q string
switch dir {
case "prev": // 列表中更早的一篇
q = `SELECT slug, title FROM notes WHERE ` + PublicNoteFilter + `
AND (updated_at > ? OR (updated_at = ? AND id < ?))
ORDER BY updated_at ASC, id DESC LIMIT 1`
case "next": // 列表中更新的一篇
q = `SELECT slug, title FROM notes WHERE ` + PublicNoteFilter + `
AND (updated_at < ? OR (updated_at = ? AND id > ?))
ORDER BY updated_at DESC, id ASC LIMIT 1`
default:
return "", "", fmt.Errorf("dir 必须为 prev|next")
}
row := s.db.QueryRow(q, n.UpdatedAt, n.UpdatedAt, n.ID)
if err := row.Scan(&slug, &title); err != nil {
if errors.Is(err, sql.ErrNoRows) {
return "", "", nil
}
return "", "", err
}
return slug, title, nil
}
// ListPublicNotesFull 公开笔记列表(**含全文**,RSS/渲染用),上限 limit 条。
func (s *Store) ListPublicNotesFull(limit int) ([]Note, error) {
q := `SELECT ` + noteColumns + ` FROM notes WHERE ` + PublicNoteFilter + `
ORDER BY updated_at DESC, id DESC LIMIT ?`
rows, err := s.db.Query(q, limit)
if err != nil {
return nil, err
}
defer rows.Close()
out := []Note{}
for rows.Next() {
n, err := scanNote(rows.Scan)
if err != nil {
return nil, err
}
out = append(out, *n)
}
return out, rows.Err()
}
// ListAdminNotes 全部正常笔记(含私有,不含回收站),不含全文。
func (s *Store) ListAdminNotes() ([]Note, error) {
q := `SELECT id, slug, title, summary, '' AS content, status, tags, pinned, deleted_at, created_at, updated_at
FROM notes WHERE deleted_at IS NULL
ORDER BY pinned DESC, updated_at DESC, id DESC`
return s.queryNotes(q)
}
// ListTrash 回收站列表(软删除中),不含全文。
func (s *Store) ListTrash() ([]Note, error) {
q := `SELECT id, slug, title, summary, '' AS content, status, tags, pinned, deleted_at, created_at, updated_at
FROM notes WHERE deleted_at IS NOT NULL
ORDER BY deleted_at DESC`
return s.queryNotes(q)
}
// ListExpiredTrash 过期待物理删除的笔记 id(gc 用)。
func (s *Store) ListExpiredTrash(before int64) ([]int64, error) {
rows, err := s.db.Query(`SELECT id FROM notes WHERE deleted_at IS NOT NULL AND deleted_at < ?`, before)
if err != nil {
return nil, err
}
defer rows.Close()
var ids []int64
for rows.Next() {
var id int64
if err := rows.Scan(&id); err != nil {
return nil, err
}
ids = append(ids, id)
}
return ids, rows.Err()
}
// DeleteNoteForever 物理删除(refs 级联,需 foreign_keys=1)。
func (s *Store) DeleteNoteForever(id int64) error {
res, err := s.db.Exec(`DELETE FROM notes WHERE id=?`, id)
if err != nil {
return err
}
if rows, _ := res.RowsAffected(); rows == 0 {
return ErrNotFound
}
return nil
}
// TagCount 标签聚合项。
type TagCount struct {
Name string `json:"name"`
Count int `json:"count"`
}
// PublicTags 仅统计公开且未删除笔记的标签(§7.1 /api/tags)。
func (s *Store) PublicTags() ([]TagCount, error) {
rows, err := s.db.Query(`
SELECT je.value AS tag, COUNT(*) AS cnt
FROM notes n, json_each(n.tags) je
WHERE ` + PublicNoteFilter + `
GROUP BY je.value
ORDER BY cnt DESC, tag ASC`)
if err != nil {
return nil, err
}
defer rows.Close()
var out []TagCount
for rows.Next() {
var tc TagCount
if err := rows.Scan(&tc.Name, &tc.Count); err != nil {
return nil, err
}
out = append(out, tc)
}
return out, rows.Err()
}
func (s *Store) queryNotes(q string, args ...any) ([]Note, error) {
rows, err := s.db.Query(q, args...)
if err != nil {
return nil, err
}
defer rows.Close()
out := []Note{}
for rows.Next() {
n, err := scanNote(rows.Scan)
if err != nil {
return nil, err
}
out = append(out, *n)
}
return out, rows.Err()
}
func marshalTags(tags []string) (string, error) {
if tags == nil {
tags = []string{}
}
b, err := json.Marshal(tags)
if err != nil {
return "", err
}
return string(b), nil
}
func boolToInt(b bool) int {
if b {
return 1
}
return 0
}
// NormalizeTags 清洗标签:去空白、去空、去重、上限 20 个、单个 ≤64 字符。
func NormalizeTags(tags []string) []string {
seen := map[string]bool{}
out := []string{}
for _, t := range tags {
t = strings.TrimSpace(t)
if t == "" || seen[t] {
continue
}
if len([]rune(t)) > 64 {
t = string([]rune(t)[:64])
}
seen[t] = true
out = append(out, t)
if len(out) >= 20 {
break
}
}
return out
}
+78
View File
@@ -0,0 +1,78 @@
package store
import (
"database/sql"
"errors"
"strconv"
)
func fmtSscanInt(s string, out *int64) (int, error) {
v, err := strconv.ParseInt(s, 10, 64)
if err != nil {
return 0, err
}
*out = v
return 1, nil
}
// Session 会话行(token 明文不入库,只存 SHA-256)。
type Session struct {
TokenHash string
CSRFToken string
CreatedAt int64
ExpiresAt int64
}
// CreateSession 插入会话行。
func (s *Store) CreateSession(tokenHash, csrfToken string, createdAt, expiresAt int64) error {
_, err := s.db.Exec(
`INSERT INTO sessions (token_hash, csrf_token, created_at, expires_at) VALUES (?, ?, ?, ?)`,
tokenHash, csrfToken, createdAt, expiresAt)
return err
}
// RotateSession 重建会话行:新 token、**csrf_token 保持不变**(§7.3-5)、新过期时间。
// 原子替换,避免窗口期内两行并存。
func (s *Store) RotateSession(oldTokenHash, newTokenHash, csrfToken string, createdAt, expiresAt int64) error {
tx, err := s.db.Begin()
if err != nil {
return err
}
defer tx.Rollback()
if _, err := tx.Exec(`DELETE FROM sessions WHERE token_hash=?`, oldTokenHash); err != nil {
return err
}
if _, err := tx.Exec(
`INSERT INTO sessions (token_hash, csrf_token, created_at, expires_at) VALUES (?, ?, ?, ?)`,
newTokenHash, csrfToken, createdAt, expiresAt); err != nil {
return err
}
return tx.Commit()
}
// GetSession 按 token 摘要取会话(含已过期行;过期判定在调用方)。
func (s *Store) GetSession(tokenHash string) (*Session, error) {
row := s.db.QueryRow(
`SELECT token_hash, csrf_token, created_at, expires_at FROM sessions WHERE token_hash=?`, tokenHash)
var sess Session
err := row.Scan(&sess.TokenHash, &sess.CSRFToken, &sess.CreatedAt, &sess.ExpiresAt)
if errors.Is(err, sql.ErrNoRows) {
return nil, ErrNotFound
}
return &sess, err
}
// DeleteSession 删除会话行(登出/轮换)。
func (s *Store) DeleteSession(tokenHash string) error {
_, err := s.db.Exec(`DELETE FROM sessions WHERE token_hash=?`, tokenHash)
return err
}
// DeleteExpiredSessions 清理过期会话,返回删除行数。
func (s *Store) DeleteExpiredSessions(now int64) (int64, error) {
res, err := s.db.Exec(`DELETE FROM sessions WHERE expires_at < ?`, now)
if err != nil {
return 0, err
}
return res.RowsAffected()
}
+70
View File
@@ -0,0 +1,70 @@
package store
import (
"database/sql"
"errors"
"fmt"
)
// Settings 键白名单(§7.1 SettingsDTO 同源;admin_password_hash 永不进入 API 响应)。
const (
KeyAdminPasswordHash = "admin_password_hash"
KeySiteTitle = "site_title"
KeySiteDesc = "site_desc"
KeyPageSize = "page_size"
)
// GetSetting 读取单个设置。
func (s *Store) GetSetting(key string) (string, bool, error) {
var v string
err := s.db.QueryRow(`SELECT value FROM settings WHERE key=?`, key).Scan(&v)
if errors.Is(err, sql.ErrNoRows) {
return "", false, nil
}
return v, err == nil, err
}
// SetSetting 写入单个设置(INSERT OR REPLACE)。
func (s *Store) SetSetting(key, value string) error {
_, err := s.db.Exec(`INSERT INTO settings (key, value) VALUES (?, ?)
ON CONFLICT(key) DO UPDATE SET value=excluded.value`, key, value)
return err
}
// SiteSettings 对外的站点设置白名单视图。
type SiteSettings struct {
SiteTitle string `json:"site_title"`
SiteDesc string `json:"site_desc"`
PageSize int `json:"page_size"`
}
// Defaults,未初始化时兜底。
const (
DefaultSiteTitle = "Pure Note"
DefaultSiteDesc = ""
DefaultPageSize = 10
)
// GetSiteSettings 读取站点设置(白名单三键,带默认值)。
func (s *Store) GetSiteSettings() (*SiteSettings, error) {
ss := &SiteSettings{SiteTitle: DefaultSiteTitle, SiteDesc: DefaultSiteDesc, PageSize: DefaultPageSize}
if v, ok, err := s.GetSetting(KeySiteTitle); err != nil {
return nil, err
} else if ok && v != "" {
ss.SiteTitle = v
}
if v, ok, err := s.GetSetting(KeySiteDesc); err != nil {
return nil, err
} else if ok {
ss.SiteDesc = v
}
if v, ok, err := s.GetSetting(KeyPageSize); err != nil {
return nil, err
} else if ok {
var n int
if _, err := fmt.Sscanf(v, "%d", &n); err == nil && n >= 1 && n <= 100 {
ss.PageSize = n
}
}
return ss, nil
}
+154
View File
@@ -0,0 +1,154 @@
// Package store SQLite 打开、版本化迁移与全部数据访问。
package store
import (
"context"
"database/sql"
"errors"
"fmt"
"net/url"
"path/filepath"
"time"
_ "modernc.org/sqlite"
)
// bgCtx DAO 内部使用的后台 context(个人规模查询均为快速查询)。
var bgCtx = context.Background()
// ErrSchemaNewer 数据库 schema 版本超出代码支持范围(§10.4 启动守卫)。
var ErrSchemaNewer = errors.New("数据库 schema 版本高于本程序支持的最高版本,拒绝启动(可用 --allow-newer 显式放行)")
// migrations 仅追加式:禁止删列/重命名/改类型(§6.2)。
var migrations = []string{
// v1: 初始 schema(§6.1 DDL)
`
CREATE TABLE IF NOT EXISTS notes (
id INTEGER PRIMARY KEY AUTOINCREMENT,
slug TEXT NOT NULL UNIQUE,
title TEXT NOT NULL,
summary TEXT NOT NULL DEFAULT '',
content TEXT NOT NULL DEFAULT '',
status TEXT NOT NULL DEFAULT 'private'
CHECK (status IN ('public','private')),
tags TEXT NOT NULL DEFAULT '[]',
pinned INTEGER NOT NULL DEFAULT 0,
deleted_at INTEGER,
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_notes_public
ON notes (status, pinned, updated_at DESC);
CREATE INDEX IF NOT EXISTS idx_notes_deleted ON notes (deleted_at);
CREATE TABLE IF NOT EXISTS images (
id INTEGER PRIMARY KEY AUTOINCREMENT,
sha256 TEXT NOT NULL UNIQUE,
mime TEXT NOT NULL,
size INTEGER NOT NULL,
data BLOB NOT NULL,
created_at INTEGER NOT NULL
);
CREATE TABLE IF NOT EXISTS image_refs (
image_id INTEGER NOT NULL REFERENCES images(id) ON DELETE CASCADE,
note_id INTEGER NOT NULL REFERENCES notes(id) ON DELETE CASCADE,
PRIMARY KEY (image_id, note_id)
);
CREATE INDEX IF NOT EXISTS idx_image_refs_note ON image_refs (note_id);
CREATE TABLE IF NOT EXISTS sessions (
token_hash TEXT PRIMARY KEY,
csrf_token TEXT NOT NULL,
created_at INTEGER NOT NULL,
expires_at INTEGER NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_sessions_expires ON sessions (expires_at);
CREATE TABLE IF NOT EXISTS settings (
key TEXT PRIMARY KEY,
value TEXT NOT NULL
);
`,
}
// MaxSchemaVersion 代码支持的最高 schema 版本。
var MaxSchemaVersion = len(migrations)
// Store 封装 *sql.DB;单连接串行写(§6.2 并发策略)。
type Store struct {
db *sql.DB
}
// Open 打开数据库并执行迁移。dbPath 为文件绝对/相对路径。
func Open(dbPath string, allowNewer bool) (*Store, error) {
// SQLite URI 不接受相对路径
abs, err := filepath.Abs(dbPath)
if err != nil {
return nil, err
}
dsn := url.URL{
Scheme: "file",
Path: abs,
RawQuery: "_pragma=journal_mode(WAL)" +
"&_pragma=busy_timeout(5000)" +
"&_pragma=foreign_keys(1)" +
"&_pragma=synchronous(NORMAL)",
}
db, err := sql.Open("sqlite", dsn.String())
if err != nil {
return nil, err
}
// 单写者串行化(§6.2)
db.SetMaxOpenConns(1)
s := &Store{db: db}
if err := s.migrate(allowNewer); err != nil {
db.Close()
return nil, err
}
return s, nil
}
func (s *Store) migrate(allowNewer bool) error {
var v int
if err := s.db.QueryRow("PRAGMA user_version").Scan(&v); err != nil {
return err
}
if v > MaxSchemaVersion && !allowNewer {
return fmt.Errorf("%w: 库版本 %d > 支持上限 %d", ErrSchemaNewer, v, MaxSchemaVersion)
}
for i := v; i < MaxSchemaVersion; i++ {
tx, err := s.db.Begin()
if err != nil {
return err
}
if _, err := tx.Exec(migrations[i]); err != nil {
tx.Rollback()
return fmt.Errorf("执行迁移 v%d 失败: %w", i+1, err)
}
// PRAGMA 不能参数化,i 为内部 int 常量,无注入面。
if _, err := tx.Exec(fmt.Sprintf("PRAGMA user_version = %d", i+1)); err != nil {
tx.Rollback()
return err
}
if err := tx.Commit(); err != nil {
return err
}
}
return nil
}
// SchemaVersion 返回当前库版本。
func (s *Store) SchemaVersion() (int, error) {
var v int
err := s.db.QueryRow("PRAGMA user_version").Scan(&v)
return v, err
}
// Close 关闭连接。
func (s *Store) Close() error { return s.db.Close() }
// DB 暴露底层连接(仅备份等维护路径使用)。
func (s *Store) DB() *sql.DB { return s.db }
func nowUnix() int64 { return time.Now().Unix() }
+122
View File
@@ -0,0 +1,122 @@
package store
import (
"errors"
"path/filepath"
"testing"
)
// TestMigrationsFromEmpty 空库 → 最新版本;重复打开幂等(§13 迁移组)。
func TestMigrationsFromEmpty(t *testing.T) {
dir := t.TempDir()
dbPath := filepath.Join(dir, "pure-note.db")
s1, err := Open(dbPath, false)
if err != nil {
t.Fatalf("首次打开失败: %v", err)
}
v1, err := s1.SchemaVersion()
if err != nil {
t.Fatal(err)
}
if v1 != MaxSchemaVersion {
t.Fatalf("期望版本 %d 实际 %d", MaxSchemaVersion, v1)
}
s1.Close()
// 重开:幂等,不再执行迁移
s2, err := Open(dbPath, false)
if err != nil {
t.Fatalf("二次打开失败: %v", err)
}
defer s2.Close()
v2, err := s2.SchemaVersion()
if err != nil {
t.Fatal(err)
}
if v2 != v1 {
t.Fatalf("迁移不幂等: %d → %d", v1, v2)
}
}
// TestMigrationNewerRejected user_version 越界拒绝启动;--allow-newer 放行(§10.4)。
func TestMigrationNewerRejected(t *testing.T) {
dir := t.TempDir()
dbPath := filepath.Join(dir, "pure-note.db")
s1, err := Open(dbPath, false)
if err != nil {
t.Fatal(err)
}
if _, err := s1.DB().Exec("PRAGMA user_version = 99"); err != nil {
t.Fatal(err)
}
s1.Close()
_, err = Open(dbPath, false)
if !errors.Is(err, ErrSchemaNewer) {
t.Fatalf("期望 ErrSchemaNewer,实际 %v", err)
}
s2, err := Open(dbPath, true) // --allow-newer
if err != nil {
t.Fatalf("allow-newer 应放行: %v", err)
}
defer s2.Close()
}
// TestForeignKeysCascade 外键级联生效(image_refs 依赖 foreign_keys=1)。
func TestForeignKeysCascade(t *testing.T) {
s := openTestStore(t)
nid, err := s.CreateNote(&Note{Slug: "casc", Title: "级联", Status: "public",
Tags: []string{}, CreatedAt: 1, UpdatedAt: 1})
if err != nil {
t.Fatal(err)
}
imgID, err := s.UpsertImage("aa", "image/png", 1, []byte{1}, 1)
if err != nil {
t.Fatal(err)
}
if err := s.RebuildImageRefs(nid, "![x](/api/images/"+itoa64(imgID)+")"); err != nil {
t.Fatal(err)
}
if err := s.DeleteNoteForever(nid); err != nil {
t.Fatal(err)
}
// 级联后图片变孤儿
orphans, err := s.ListOrphanImages()
if err != nil {
t.Fatal(err)
}
found := false
for _, o := range orphans {
if o.ID == imgID {
found = true
}
}
if !found {
t.Fatal("物理删除笔记后其引用关系应被级联删除(图片转为孤儿)")
}
}
func openTestStore(t *testing.T) *Store {
t.Helper()
s, err := Open(filepath.Join(t.TempDir(), "test.db"), false)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { s.Close() })
return s
}
func itoa64(v int64) string {
if v == 0 {
return "0"
}
var b [20]byte
i := len(b)
for v > 0 {
i--
b[i] = byte('0' + v%10)
v /= 10
}
return string(b[i:])
}