Files
ONE/backend/internal/store/store.go
T
Sakurasan 1b3eb870da 文件删除前查引用:被文章/项目/站主头像用着就先挡住
- store.FileReferences(key) 反查谁在用这个文件:文章(封面 + content_md +
  content_html)、项目封面、settings.owner_avatar_key。匹配的键是 files.key 本身
  而不是完整 URL——本地 /uploads/{key}、R2 {publicBase}/{key}、缩略图
  /uploads/thumb/{key} 三种形态都以 key 结尾,换存储端后正文里的老链接照样查得到。
  草稿也算:现在没发布,删了将来发出来就是裂的。
- 按需 LIKE 现扫,不维护计数表:写入口有编辑器、外链转存、短文、项目、头像好几处,
  计数一旦漂移就再也信不过;删文件是低频操作,扫全表几十毫秒换一个永远正确的答案。
- 新端点 GET /api/admin/files/{id}/refs;DELETE 默认对在用文件返回 409(消息带引用数
  和前三个位置名),明确带 force=1 才真删。守卫放在动 blob 之前——存储端一删就没法回头。
- 后台删除弹层先把引用清单摊出来(文章《标题》(草稿)、站主头像),仍然要删才带 force。
  引用查询失败就按无引用走:真被引用时后端会 409,不会静默删掉在用文件。

顺手把另一个 agent 工具的本地草稿目录(.zcode/、.zcodeignore)加进 .gitignore。
2026-09-30 02:03:14 +08:00

2225 lines
67 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// Package store owns the schema and every SQL query. All SQL is written with
// `?` placeholders and re-bound to `$n` when running on PostgreSQL.
package store
import (
"context"
"database/sql"
"encoding/json"
"errors"
"fmt"
"regexp"
"strings"
"time"
"oneblog/internal/db"
"oneblog/internal/model"
"oneblog/internal/render"
)
func renderHTML(md string) string { return render.Markdown(md) }
func readingMinutes(md string) int { return render.ReadingMinutes(md) }
var ErrNotFound = errors.New("not found")
// ErrConflict 表示要建的唯一键已被占(例如某个外部身份已绑到别的账号)。
var ErrConflict = errors.New("conflict")
type Store struct {
db *db.DB
}
func New(d *db.DB) (*Store, error) {
s := &Store{db: d}
if err := s.migrate(); err != nil {
return nil, err
}
return s, nil
}
// Ping 供 /api/health 探活数据库连接。
func (s *Store) Ping(ctx context.Context) error { return s.db.PingContext(ctx) }
func now() string { return time.Now().UTC().Format(time.RFC3339) }
func (s *Store) migrate() error {
ai := s.db.AutoInc()
stmts := []string{
fmt.Sprintf(`CREATE TABLE IF NOT EXISTS posts (
id %s,
kind TEXT NOT NULL DEFAULT 'long',
title TEXT NOT NULL DEFAULT '',
slug TEXT NOT NULL,
summary TEXT NOT NULL DEFAULT '',
cover_url TEXT NOT NULL DEFAULT '',
content_md TEXT NOT NULL DEFAULT '',
content_html TEXT NOT NULL DEFAULT '',
status TEXT NOT NULL DEFAULT 'draft',
published_at TEXT NOT NULL,
created_at TEXT NOT NULL,
updated_at TEXT NOT NULL,
reading_minutes INTEGER NOT NULL DEFAULT 1
)`, ai),
fmt.Sprintf(`CREATE TABLE IF NOT EXISTS tags (
id %s,
name TEXT NOT NULL,
slug TEXT NOT NULL,
color TEXT NOT NULL DEFAULT ''
)`, ai),
`CREATE TABLE IF NOT EXISTS post_tags (
post_id INTEGER NOT NULL,
tag_id INTEGER NOT NULL,
PRIMARY KEY (post_id, tag_id)
)`,
`CREATE TABLE IF NOT EXISTS settings (
key TEXT PRIMARY KEY,
value TEXT NOT NULL DEFAULT ''
)`,
fmt.Sprintf(`CREATE TABLE IF NOT EXISTS projects (
id %s,
title TEXT NOT NULL DEFAULT '',
slug TEXT NOT NULL,
summary TEXT NOT NULL DEFAULT '',
cover_url TEXT NOT NULL DEFAULT '',
url TEXT NOT NULL DEFAULT '',
repo_url TEXT NOT NULL DEFAULT '',
status TEXT NOT NULL DEFAULT 'published',
position INTEGER NOT NULL DEFAULT 0,
created_at TEXT NOT NULL,
updated_at TEXT NOT NULL
)`, ai),
fmt.Sprintf(`CREATE TABLE IF NOT EXISTS files (
id %s,
key TEXT NOT NULL,
name TEXT NOT NULL,
mime TEXT NOT NULL DEFAULT '',
size INTEGER NOT NULL DEFAULT 0,
sha256 TEXT NOT NULL DEFAULT '',
store TEXT NOT NULL DEFAULT 'local',
created_at TEXT NOT NULL
)`, ai),
fmt.Sprintf(`CREATE TABLE IF NOT EXISTS users (
id %s,
provider TEXT NOT NULL,
handle TEXT NOT NULL,
name TEXT NOT NULL DEFAULT '',
avatar_url TEXT NOT NULL DEFAULT '',
url TEXT NOT NULL DEFAULT '',
banned INTEGER NOT NULL DEFAULT 0,
created_at TEXT NOT NULL,
UNIQUE(provider, handle)
)`, ai),
fmt.Sprintf(`CREATE TABLE IF NOT EXISTS comments (
id %s,
post_id INTEGER NOT NULL,
user_id INTEGER NOT NULL,
parent_id INTEGER NOT NULL DEFAULT 0,
root_id INTEGER NOT NULL DEFAULT 0,
body_md TEXT NOT NULL DEFAULT '',
body_html TEXT NOT NULL DEFAULT '',
status TEXT NOT NULL DEFAULT 'visible',
is_deleted INTEGER NOT NULL DEFAULT 0,
created_at TEXT NOT NULL,
edited_at TEXT NOT NULL DEFAULT ''
)`, ai),
// 第三方身份绑定:一个账号每个平台只能绑一条(UNIQUE(user_id,provider)),
// 同一个外部账号也只能属于一个用户(UNIQUE(provider,extern_uid))——
// 后者是防接管的关键:不能靠「先用我的 GitHub 登录、再把你的账号绑上来」占位。
fmt.Sprintf(`CREATE TABLE IF NOT EXISTS user_identities (
id %s,
user_id INTEGER NOT NULL,
provider TEXT NOT NULL,
extern_uid TEXT NOT NULL,
display TEXT NOT NULL DEFAULT '',
created_at TEXT NOT NULL,
UNIQUE(provider, extern_uid),
UNIQUE(user_id, provider)
)`, ai),
// Passkey 一人可多把(笔记本 + 手机),所以不加 UNIQUE(user_id)
fmt.Sprintf(`CREATE TABLE IF NOT EXISTS passkeys (
id %s,
user_id INTEGER NOT NULL,
credential_id TEXT NOT NULL UNIQUE,
public_key TEXT NOT NULL,
sign_count INTEGER NOT NULL DEFAULT 0,
name TEXT NOT NULL DEFAULT '',
created_at TEXT NOT NULL,
last_used_at TEXT NOT NULL DEFAULT ''
)`, ai),
}
for _, q := range stmts {
if _, err := s.db.Exec(s.db.Q(q)); err != nil {
return fmt.Errorf("migrate: %w", err)
}
}
// Lightweight column-add migrations for older databases. SQLite supports
// ADD COLUMN; Postgres doesn't support IF NOT EXISTS on column add, so we
// swallow "already exists" errors either way.
columnAdds := []string{
`ALTER TABLE posts ADD COLUMN cover_url TEXT NOT NULL DEFAULT ''`,
`ALTER TABLE tags ADD COLUMN color TEXT NOT NULL DEFAULT ''`,
// 短文配图(快发盒上传,Twitter 式网格展示),JSON 数组
`ALTER TABLE posts ADD COLUMN images TEXT NOT NULL DEFAULT '[]'`,
// 正文首个外链的预览卡片(og 标题/描述/封面),JSON 对象或空串
`ALTER TABLE posts ADD COLUMN link_card TEXT NOT NULL DEFAULT ''`,
// 账号角色:owner(站主,全库唯一)/ reader(评论区访客)
`ALTER TABLE users ADD COLUMN role TEXT NOT NULL DEFAULT 'reader'`,
}
for _, q := range columnAdds {
if _, err := s.db.Exec(s.db.Q(q)); err != nil && !strings.Contains(err.Error(), "already exists") &&
!strings.Contains(err.Error(), "duplicate column") {
return fmt.Errorf("migrate column: %w", err)
}
}
// Indexes/unique constraints need dialect-specific "IF NOT EXISTS" support.
indexes := []struct{ name, ddl string }{
{"idx_posts_slug", `CREATE UNIQUE INDEX IF NOT EXISTS idx_posts_slug ON posts(slug)`},
{"idx_posts_feed", `CREATE INDEX IF NOT EXISTS idx_posts_feed ON posts(status, published_at DESC)`},
{"idx_tags_slug", `CREATE UNIQUE INDEX IF NOT EXISTS idx_tags_slug ON tags(slug)`},
{"idx_post_tags_tag", `CREATE INDEX IF NOT EXISTS idx_post_tags_tag ON post_tags(tag_id)`},
{"idx_projects_slug", `CREATE UNIQUE INDEX IF NOT EXISTS idx_projects_slug ON projects(slug)`},
{"idx_files_key", `CREATE UNIQUE INDEX IF NOT EXISTS idx_files_key ON files(key)`},
{"idx_files_created", `CREATE INDEX IF NOT EXISTS idx_files_created ON files(created_at DESC)`},
{"idx_comments_post", `CREATE INDEX IF NOT EXISTS idx_comments_post ON comments(post_id, created_at)`},
{"idx_comments_user", `CREATE INDEX IF NOT EXISTS idx_comments_user ON comments(user_id)`},
{"idx_comments_status", `CREATE INDEX IF NOT EXISTS idx_comments_status ON comments(status, created_at DESC)`},
{"idx_identities_user", `CREATE INDEX IF NOT EXISTS idx_identities_user ON user_identities(user_id)`},
{"idx_passkeys_user", `CREATE INDEX IF NOT EXISTS idx_passkeys_user ON passkeys(user_id)`},
}
for _, ix := range indexes {
if _, err := s.db.Exec(s.db.Q(ix.ddl)); err != nil && !strings.Contains(err.Error(), "already exists") {
return fmt.Errorf("migrate index %s: %w", ix.name, err)
}
}
return s.seedSettings()
}
func (s *Store) seedSettings() error {
defs := map[string]string{
"site_title": "ONE · 一个博客",
"site_desc": "长文与短文,同一种节奏。",
"author_name": "ONE",
"author_bio": "写点长的,也写点短的。",
"footer_note": "© ONE · 一个博客",
"icp": "",
"posts_per_page": "10",
"comments_enabled": "1",
"comments_review": "0",
}
for k, v := range defs {
if s.db.Dialect == db.Postgres {
_, err := s.db.Exec(s.db.Q(`INSERT INTO settings(key,value) VALUES (?,?)
ON CONFLICT (key) DO NOTHING`), k, v)
if err != nil {
return err
}
continue
}
if _, err := s.db.Exec(s.db.Q(`INSERT OR IGNORE INTO settings(key,value) VALUES (?,?)`), k, v); err != nil {
return err
}
}
return nil
}
// ---------- settings ----------
func (s *Store) GetSettings() (model.Settings, error) {
rows, err := s.db.Query(s.db.Q(`SELECT key, value FROM settings`))
if err != nil {
return model.Settings{}, err
}
defer rows.Close()
m := map[string]string{}
for rows.Next() {
var k, v string
if err := rows.Scan(&k, &v); err != nil {
return model.Settings{}, err
}
m[k] = v
}
return settingsFromMap(m), rows.Err()
}
func settingsFromMap(m map[string]string) model.Settings {
st := model.Settings{
SiteTitle: m["site_title"],
SiteDesc: m["site_desc"],
AuthorName: m["author_name"],
AuthorBio: m["author_bio"],
FooterNote: m["footer_note"],
ICPLicense: m["icp"],
PostsPerPage: 10,
}
// LightSkinID is the new canonical name; fall back to legacy theme_id
// for clients that haven't been updated, and finally to "paper".
st.LightSkinID = m["light_skin_id"]
if st.LightSkinID == "" {
st.LightSkinID = m["theme_id"]
}
if st.LightSkinID == "" {
st.LightSkinID = "paper"
}
// Mirror the value into ThemeID so legacy API consumers still see it
// in the JSON response.
st.ThemeID = st.LightSkinID
st.UIID = sanitizeUI(m["ui_id"])
st.CustomCSS = decodeCSSMap(m["custom_css"])
st.CustomJS = m["custom_js"]
st.SocialLinks = decodeSocialLinks(m["social_links"])
st.AuthorAvatarKey = m["owner_avatar_key"]
// 开关类:'1' / 'true' 都算开,其余(含空)算关
st.CommentsEnabled = m["comments_enabled"] == "1" || strings.EqualFold(m["comments_enabled"], "true")
st.CommentsReview = m["comments_review"] == "1" || strings.EqualFold(m["comments_review"], "true")
if n := atoi(m["posts_per_page"]); n > 0 {
st.PostsPerPage = n
}
return st
}
// decodeSocialLinks 把 KV 里的 JSON 数组还原成社交链接。容错:空串、坏 JSON、
// 缺 label/url 的条目一律丢弃,返回空切片(前台区块自动隐藏)。
func decodeSocialLinks(s string) []model.SocialLink {
out := []model.SocialLink{}
if strings.TrimSpace(s) == "" {
return out
}
if err := json.Unmarshal([]byte(s), &out); err != nil {
return []model.SocialLink{}
}
clean := make([]model.SocialLink, 0, len(out))
for _, l := range out {
l.Label = strings.TrimSpace(l.Label)
l.URL = strings.TrimSpace(l.URL)
if l.Label != "" && l.URL != "" {
clean = append(clean, l)
}
}
return clean
}
func encodeSocialLinks(ls []model.SocialLink) string {
if len(ls) == 0 {
return ""
}
b, err := json.Marshal(ls)
if err != nil {
return ""
}
return string(b)
}
func atoi(v string) int {
n := 0
for _, r := range v {
if r < '0' || r > '9' {
return 0
}
n = n*10 + int(r-'0')
}
return n
}
// ValidLightSkins is the whitelist of light-side skin ids the API accepts.
// Anything outside this set falls back to "paper" on write.
var ValidLightSkins = map[string]bool{
"paper": true,
"sage": true,
"rose": true,
}
const (
// UIDefault is the UI used when nothing valid is stored.
UIDefault = "classic"
// UIClassic is the original minimalist front-end.
UIClassic = "classic"
// UIVivid is the livelier front-end, which also owns custom-CSS support.
UIVivid = "vivid"
)
// ValidUIs is the whitelist of front-end UI ids. Anything else falls back to
// "classic" on both read and write.
var ValidUIs = map[string]bool{
UIClassic: true,
UIVivid: true,
}
// ValidCSSSections whitelists the page sections an owner may target with
// custom CSS. Keep in sync with frontend/src/ui/sections.js.
var ValidCSSSections = map[string]bool{
"global": true,
"home": true,
"post": true,
"archive": true,
"tags": true,
"projects": true,
"about": true,
}
const (
// maxCustomCSSBytes caps the whole stylesheet set; maxSectionCSSBytes caps
// any single section. Both are generous for hand-written CSS but keep a
// runaway paste from bloating every /api/site response.
maxCustomCSSBytes = 64 << 10
maxSectionCSSBytes = 16 << 10
)
func sanitizeUI(id string) string {
id = strings.TrimSpace(id)
if ValidUIs[id] {
return id
}
return UIDefault
}
// decodeCSSMap turns the stored JSON blob into a section→CSS map. It never
// returns nil, and it drops unknown sections and over-long values even if the
// row was hand-edited, so callers can trust the result.
func decodeCSSMap(raw string) map[string]string {
out := map[string]string{}
if strings.TrimSpace(raw) == "" {
return out
}
var parsed map[string]string
if err := json.Unmarshal([]byte(raw), &parsed); err != nil {
return out
}
total := 0
for section, css := range parsed {
if !ValidCSSSections[section] {
continue
}
css = strings.ReplaceAll(css, "\x00", "")
if strings.TrimSpace(css) == "" || len(css) > maxSectionCSSBytes {
continue
}
if total+len(css) > maxCustomCSSBytes {
continue
}
total += len(css)
out[section] = css
}
return out
}
// encodeCSSMap is the write-side counterpart: unknown sections are dropped and
// oversized values are refused, then the map is serialized. A nil or empty map
// encodes to "{}" so a full-replace PUT clears everything deterministically.
func encodeCSSMap(in map[string]string) string {
clean := map[string]string{}
total := 0
for section, css := range in {
if !ValidCSSSections[section] {
continue
}
css = strings.ReplaceAll(css, "\x00", "")
if strings.TrimSpace(css) == "" || len(css) > maxSectionCSSBytes {
continue
}
if total+len(css) > maxCustomCSSBytes {
continue
}
total += len(css)
clean[section] = css
}
b, err := json.Marshal(clean)
if err != nil {
return "{}"
}
return string(b)
}
func (s *Store) UpdateSettings(st model.Settings) error {
if st.PostsPerPage <= 0 {
st.PostsPerPage = 10
}
if !ValidLightSkins[st.LightSkinID] {
st.LightSkinID = "paper"
}
st.UIID = sanitizeUI(st.UIID)
sets := map[string]string{
"site_title": st.SiteTitle,
"site_desc": st.SiteDesc,
"author_name": st.AuthorName,
"author_bio": st.AuthorBio,
"footer_note": st.FooterNote,
"icp": st.ICPLicense,
"posts_per_page": fmt.Sprint(st.PostsPerPage),
"light_skin_id": st.LightSkinID,
// Mirror to theme_id so any older client still sees something.
"theme_id": st.LightSkinID,
"ui_id": st.UIID,
// Full replace: an omitted/empty map clears every section's CSS.
"custom_css": encodeCSSMap(st.CustomCSS),
// 原样存:站主自己的代码,不做任何转义/清洗。
"custom_js": st.CustomJS,
// 社交 / 源码链接:JSON 数组,空数组存 "[]"。
"social_links": encodeSocialLinks(st.SocialLinks),
// 开关统一存 '1' / '0'。
"comments_enabled": b2s(st.CommentsEnabled),
"comments_review": b2s(st.CommentsReview),
}
for k, v := range sets {
if s.db.Dialect == db.Postgres {
if _, err := s.db.Exec(s.db.Q(`INSERT INTO settings(key,value) VALUES (?,?)
ON CONFLICT (key) DO UPDATE SET value = EXCLUDED.value`), k, v); err != nil {
return err
}
continue
}
if _, err := s.db.Exec(s.db.Q(`INSERT INTO settings(key,value) VALUES (?,?)
ON CONFLICT(key) DO UPDATE SET value = excluded.value`), k, v); err != nil {
return err
}
}
return nil
}
// SetSetting 只写一个 KV,绕开 UpdateSettings 的全量替换。
// 账户页改头像用它的自己的键,免得「站点设置」保存时被顺带清掉。
func (s *Store) SetSetting(key, value string) error {
if s.db.Dialect == db.Postgres {
_, err := s.db.Exec(s.db.Q(`INSERT INTO settings(key,value) VALUES (?,?)
ON CONFLICT (key) DO UPDATE SET value = EXCLUDED.value`), key, value)
return err
}
_, err := s.db.Exec(s.db.Q(`INSERT INTO settings(key,value) VALUES (?,?)
ON CONFLICT(key) DO UPDATE SET value = excluded.value`), key, value)
return err
}
// ---------- posts ----------
type ListOptions struct {
Kind string
Tag string
Query string
Status string // "" = published only (public), "any" = all (admin)
Page int
Size int
OrderBy string
}
// OrderBy 来自查询参数,必须白名单校验后才能拼进 SQL。
var allowedOrder = map[string]bool{
"published_at desc": true,
"published_at asc": true,
"updated_at desc": true,
"updated_at asc": true,
"created_at desc": true,
"reading_minutes desc": true,
"reading_minutes asc": true,
}
func sanitizeOrder(o string) string {
key := strings.ToLower(strings.Join(strings.Fields(o), " "))
if allowedOrder[key] {
return key
}
return "published_at desc"
}
const postCols = `id, kind, title, slug, summary, cover_url, content_md, content_html, status,
published_at, created_at, updated_at, reading_minutes, LENGTH(content_md), images, link_card`
// listCols 用于列表/时间线:不传 content_md(前端不用),
// 长文 content_html 只截 600 字符供无摘要时提取纯文本,短文保留全文渲染。
const listCols = `id, kind, title, slug, summary, cover_url,
'' AS content_md,
CASE WHEN kind = 'short' THEN content_html ELSE substr(content_html, 1, 600) END AS content_html,
status, published_at, created_at, updated_at, reading_minutes, LENGTH(content_md), images, link_card`
// images 列的 JSON 编解码(列存 '[]',Go 侧 []string;坏数据静默为空)。
func decodeImages(s string) []string {
out := []string{}
if strings.TrimSpace(s) == "" {
return out
}
if err := json.Unmarshal([]byte(s), &out); err != nil {
return []string{}
}
return out
}
func encodeImages(imgs []string) string {
if len(imgs) == 0 {
return "[]"
}
b, err := json.Marshal(imgs)
if err != nil {
return "[]"
}
return string(b)
}
// link_card 列的 JSON 编解码:空串/坏数据一律当作「没有卡片」。
func decodeLinkCard(s string) *model.LinkCard {
s = strings.TrimSpace(s)
if s == "" {
return nil
}
var c model.LinkCard
if err := json.Unmarshal([]byte(s), &c); err != nil || c.Empty() {
return nil
}
return &c
}
func encodeLinkCard(c *model.LinkCard) string {
if c.Empty() {
return ""
}
b, err := json.Marshal(c)
if err != nil {
return ""
}
return string(b)
}
func scanPost(rows interface{ Scan(...any) error }) (model.Post, error) {
var p model.Post
var imgs, card string
err := rows.Scan(&p.ID, &p.Kind, &p.Title, &p.Slug, &p.Summary, &p.CoverURL,
&p.ContentMd, &p.ContentHTML, &p.Status, &p.PublishedAt, &p.CreatedAt, &p.UpdatedAt,
&p.ReadingMinutes, &p.ContentLen, &imgs, &card)
p.Tags = []string{}
p.Images = decodeImages(imgs)
p.LinkCard = decodeLinkCard(card)
return p, err
}
func (s *Store) List(o ListOptions) (model.Page, error) {
if o.Page < 1 {
o.Page = 1
}
if o.Size < 1 || o.Size > 100 {
o.Size = 10
}
where := []string{}
args := []any{}
if o.Status == "any" {
// admin: no status filter
} else if o.Status != "" {
where = append(where, "status = ?")
args = append(args, o.Status)
} else {
where = append(where, "status = 'published'")
}
if o.Kind != "" {
where = append(where, "kind = ?")
args = append(args, o.Kind)
}
if o.Tag != "" {
where = append(where, `id IN (SELECT pt.post_id FROM post_tags pt JOIN tags t ON t.id = pt.tag_id
WHERE t.slug = ? OR t.name = ?)`)
args = append(args, o.Tag, o.Tag)
}
if o.Query != "" {
like := "%" + strings.ToLower(o.Query) + "%"
where = append(where, `(lower(title) LIKE ? OR lower(summary) LIKE ? OR lower(content_md) LIKE ?)`)
args = append(args, like, like, like)
}
w := ""
if len(where) > 0 {
w = "WHERE " + strings.Join(where, " AND ")
}
order := sanitizeOrder(o.OrderBy)
var total int
if err := s.db.QueryRow(s.db.Q(`SELECT COUNT(*) FROM posts `+w), args...).Scan(&total); err != nil {
return model.Page{}, err
}
q := s.db.Q(fmt.Sprintf(`SELECT %s FROM posts %s ORDER BY %s LIMIT ? OFFSET ?`, listCols, w, order))
rows, err := s.db.Query(q, append(args, o.Size, (o.Page-1)*o.Size)...)
if err != nil {
return model.Page{}, err
}
defer rows.Close()
items := []model.Post{}
for rows.Next() {
p, err := scanPost(rows)
if err != nil {
return model.Page{}, err
}
items = append(items, p)
}
if err := rows.Err(); err != nil {
return model.Page{}, err
}
if err := s.attachTags(items); err != nil {
return model.Page{}, err
}
return model.Page{Items: items, Total: total, Page: o.Page, Size: o.Size}, nil
}
func (s *Store) attachTags(posts []model.Post) error {
if len(posts) == 0 {
return nil
}
ids := make([]any, 0, len(posts))
idx := map[int64]int{}
for i, p := range posts {
ids = append(ids, p.ID)
idx[p.ID] = i
}
ph := strings.TrimSuffix(strings.Repeat("?,", len(ids)), ",")
q := s.db.Q(fmt.Sprintf(`SELECT pt.post_id, t.name FROM post_tags pt
JOIN tags t ON t.id = pt.tag_id WHERE pt.post_id IN (%s) ORDER BY t.name`, ph))
rows, err := s.db.Query(q, ids...)
if err != nil {
return err
}
defer rows.Close()
for rows.Next() {
var pid int64
var name string
if err := rows.Scan(&pid, &name); err != nil {
return err
}
if i, ok := idx[pid]; ok {
posts[i].Tags = append(posts[i].Tags, name)
}
}
return rows.Err()
}
func (s *Store) Get(id int64) (model.Post, error) {
var p model.Post
row := s.db.QueryRow(s.db.Q(`SELECT `+postCols+` FROM posts WHERE id = ?`), id)
if err := scanPostInto(row, &p); err != nil {
return p, err
}
items := []model.Post{p}
if err := s.attachTags(items); err != nil {
return p, err
}
return items[0], nil
}
// Neighbors 返回某篇已发布文章在时间线上的前后邻居(详情页侧栏用)。
// 返回顺序与时间线一致(published_at DESC):最旧在最前、最新在最后,
// 目标文章夹在中间,前后各取最多 2 篇。
func (s *Store) Neighbors(slug string) ([]model.Post, error) {
rows, err := s.db.Query(s.db.Q(`SELECT `+listCols+` FROM posts
WHERE status = ? ORDER BY published_at DESC`), model.StatusPublished)
if err != nil {
return nil, err
}
defer rows.Close()
all := []model.Post{}
for rows.Next() {
p, err := scanPost(rows)
if err != nil {
return nil, err
}
all = append(all, p)
}
if err := rows.Err(); err != nil {
return nil, err
}
idx := -1
for i, p := range all {
if p.Slug == slug {
idx = i
break
}
}
if idx < 0 {
return []model.Post{}, nil
}
lo, hi := idx-2, idx+2
if lo < 0 {
lo = 0
}
if hi > len(all)-1 {
hi = len(all) - 1
}
out := all[lo : hi+1]
if err := s.attachTags(out); err != nil {
return nil, err
}
return out, nil
}
func (s *Store) GetBySlug(slug string) (model.Post, error) {
var p model.Post
row := s.db.QueryRow(s.db.Q(`SELECT `+postCols+` FROM posts WHERE slug = ?`), slug)
if err := scanPostInto(row, &p); err != nil {
return p, err
}
items := []model.Post{p}
if err := s.attachTags(items); err != nil {
return p, err
}
return items[0], nil
}
func scanPostInto(row *sql.Row, p *model.Post) error {
var imgs, card string
err := row.Scan(&p.ID, &p.Kind, &p.Title, &p.Slug, &p.Summary, &p.CoverURL,
&p.ContentMd, &p.ContentHTML, &p.Status, &p.PublishedAt, &p.CreatedAt, &p.UpdatedAt,
&p.ReadingMinutes, &p.ContentLen, &imgs, &card)
if err == sql.ErrNoRows {
return ErrNotFound
}
if err != nil {
return err
}
p.Tags = []string{}
p.Images = decodeImages(imgs)
p.LinkCard = decodeLinkCard(card)
return nil
}
// ---------- slug ----------
var slugSep = regexp.MustCompile(`[^\p{L}\p{N}]+`)
func Slugify(s string) string {
s = strings.TrimSpace(strings.ToLower(s))
s = slugSep.ReplaceAllString(s, "-")
s = strings.Trim(s, "-")
return s
}
func (s *Store) uniqueSlug(base string, excludeID int64) string {
base = Slugify(base)
if base == "" {
base = "post"
}
candidate := base
for i := 2; ; i++ {
var id int64
err := s.db.QueryRow(s.db.Q(`SELECT id FROM posts WHERE slug = ? AND id <> ?`), candidate, excludeID).Scan(&id)
if err == sql.ErrNoRows {
return candidate
}
if err != nil {
return fmt.Sprintf("%s-%d", base, time.Now().Unix())
}
candidate = fmt.Sprintf("%s-%d", base, i)
}
}
func (s *Store) uniqueProjectSlug(base string, excludeID int64) string {
base = Slugify(base)
if base == "" {
base = "project"
}
candidate := base
for i := 2; ; i++ {
var id int64
err := s.db.QueryRow(s.db.Q(`SELECT id FROM projects WHERE slug = ? AND id <> ?`), candidate, excludeID).Scan(&id)
if err == sql.ErrNoRows {
return candidate
}
if err != nil {
return fmt.Sprintf("%s-%d", base, time.Now().Unix())
}
candidate = fmt.Sprintf("%s-%d", base, i)
}
}
// ---------- projects ----------
func scanProject(row interface{ Scan(...any) error }) (model.Project, error) {
var p model.Project
err := row.Scan(&p.ID, &p.Title, &p.Slug, &p.Summary, &p.CoverURL, &p.URL,
&p.RepoURL, &p.Status, &p.Position, &p.CreatedAt, &p.UpdatedAt)
return p, err
}
// ListProjects returns projects with an optional status filter. Pass "" to
// get every project (admin), or model.StatusPublished / model.StatusDraft to
// narrow. Results are ordered by explicit Position then recency.
func (s *Store) ListProjects(status string) ([]model.Project, error) {
where := ""
args := []any{}
if status == model.StatusPublished || status == model.StatusDraft {
where = " WHERE status = ?"
args = append(args, status)
}
q := s.db.Q(`SELECT id,title,slug,summary,cover_url,url,repo_url,status,position,created_at,updated_at
FROM projects` + where + ` ORDER BY position ASC, created_at DESC`)
rows, err := s.db.Query(q, args...)
if err != nil {
return nil, err
}
defer rows.Close()
out := []model.Project{}
for rows.Next() {
p, err := scanProject(rows)
if err != nil {
return nil, err
}
out = append(out, p)
}
return out, rows.Err()
}
func (s *Store) GetProject(id int64) (model.Project, error) {
var p model.Project
err := s.db.QueryRow(s.db.Q(`SELECT id,title,slug,summary,cover_url,url,repo_url,status,position,created_at,updated_at
FROM projects WHERE id = ?`), id).Scan(&p.ID, &p.Title, &p.Slug, &p.Summary,
&p.CoverURL, &p.URL, &p.RepoURL, &p.Status, &p.Position, &p.CreatedAt, &p.UpdatedAt)
if err == sql.ErrNoRows {
return p, ErrNotFound
}
return p, err
}
func (s *Store) CreateProject(in model.ProjectInput) (model.Project, error) {
p := model.Project{
Title: strings.TrimSpace(in.Title),
Slug: strings.TrimSpace(in.Slug),
Summary: strings.TrimSpace(in.Summary),
CoverURL: strings.TrimSpace(in.CoverURL),
URL: strings.TrimSpace(in.URL),
RepoURL: strings.TrimSpace(in.RepoURL),
Status: NormalizeStatus(in.Status),
Position: in.Position,
}
if p.Status == "" {
p.Status = model.StatusPublished
}
if p.Slug == "" {
p.Slug = Slugify(p.Title)
}
if p.Slug == "" {
p.Slug = "project-" + time.Now().UTC().Format("20060102-150405")
}
p.Slug = s.uniqueProjectSlug(p.Slug, 0)
p.CreatedAt = now()
p.UpdatedAt = p.CreatedAt
q := s.db.Q(`INSERT INTO projects (title,slug,summary,cover_url,url,repo_url,status,position,created_at,updated_at)
VALUES (?,?,?,?,?,?,?,?,?,?)`)
var id int64
if s.db.Dialect == db.Postgres {
err := s.db.QueryRow(q, p.Title, p.Slug, p.Summary, p.CoverURL, p.URL, p.RepoURL,
p.Status, p.Position, p.CreatedAt, p.UpdatedAt).Scan(&id)
if err != nil {
return p, err
}
} else {
res, err := s.db.Exec(q, p.Title, p.Slug, p.Summary, p.CoverURL, p.URL, p.RepoURL,
p.Status, p.Position, p.CreatedAt, p.UpdatedAt)
if err != nil {
return p, err
}
id, err = res.LastInsertId()
if err != nil {
return p, err
}
}
return s.GetProject(id)
}
func (s *Store) UpdateProject(id int64, in model.ProjectInput) (model.Project, error) {
cur, err := s.GetProject(id)
if err != nil {
return cur, err
}
p := cur
if in.Title != "" {
p.Title = strings.TrimSpace(in.Title)
}
if in.Summary != "" {
p.Summary = strings.TrimSpace(in.Summary)
}
if in.CoverURL != "" {
p.CoverURL = strings.TrimSpace(in.CoverURL)
}
if in.URL != "" {
p.URL = strings.TrimSpace(in.URL)
}
if in.RepoURL != "" {
p.RepoURL = strings.TrimSpace(in.RepoURL)
}
if in.Status != "" {
p.Status = NormalizeStatus(in.Status)
}
if in.Slug != "" && in.Slug != cur.Slug {
p.Slug = s.uniqueProjectSlug(in.Slug, id)
}
p.Position = in.Position
p.UpdatedAt = now()
if _, err := s.db.Exec(s.db.Q(`UPDATE projects SET title=?,slug=?,summary=?,cover_url=?,url=?,repo_url=?,
status=?,position=?,updated_at=? WHERE id=?`),
p.Title, p.Slug, p.Summary, p.CoverURL, p.URL, p.RepoURL, p.Status, p.Position, p.UpdatedAt, id); err != nil {
return p, err
}
return s.GetProject(id)
}
func (s *Store) DeleteProject(id int64) error {
res, err := s.db.Exec(s.db.Q(`DELETE FROM projects WHERE id = ?`), id)
if err != nil {
return err
}
if n, err := res.RowsAffected(); err == nil && n == 0 {
return ErrNotFound
}
return nil
}
// ---------- write ----------
func (s *Store) Create(in model.PostInput) (model.Post, error) {
p := model.Post{
Kind: in.Kind,
Title: strings.TrimSpace(in.Title),
Slug: in.Slug,
Status: in.Status,
CoverURL: strings.TrimSpace(in.CoverURL),
Tags: []string{},
}
if p.Kind == "" {
p.Kind = model.KindLong
}
if p.Status == "" {
p.Status = model.StatusDraft
}
if p.Slug == "" {
p.Slug = Slugify(p.Title)
}
if p.Slug == "" {
// Titles without any latin characters (typical for short notes) get a
// date-based slug instead of colliding on "post", "post-2", ...
p.Slug = "s-" + time.Now().UTC().Format("20060102-150405")
}
p.Slug = s.uniqueSlug(p.Slug, 0)
// Short posts have no visible title, but archive and tag listings still
// need something to index them by.
if p.Kind == model.KindShort && p.Title == "" {
p.Title = render.TitleFromMarkdown(in.ContentMd)
}
p.Summary = strings.TrimSpace(in.Summary)
p.ContentMd = in.ContentMd
p.ContentHTML = renderHTML(in.ContentMd)
p.PublishedAt = in.PublishedAt
p.CreatedAt = now()
p.UpdatedAt = p.CreatedAt
if p.PublishedAt == "" {
p.PublishedAt = p.CreatedAt
}
if in.ReadingMinutes != nil && *in.ReadingMinutes > 0 {
p.ReadingMinutes = *in.ReadingMinutes
} else {
p.ReadingMinutes = readingMinutes(in.ContentMd)
}
var id int64
q := s.db.Q(`INSERT INTO posts (kind,title,slug,summary,cover_url,content_md,content_html,status,
published_at,created_at,updated_at,reading_minutes,images,link_card)
VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?)`)
if s.db.Dialect == db.Postgres {
err := s.db.QueryRow(q, p.Kind, p.Title, p.Slug, p.Summary, p.CoverURL, p.ContentMd, p.ContentHTML,
p.Status, p.PublishedAt, p.CreatedAt, p.UpdatedAt, p.ReadingMinutes, encodeImages(in.Images),
encodeLinkCard(in.LinkCard)).Scan(&id)
if err != nil {
return p, err
}
} else {
res, err := s.db.Exec(q, p.Kind, p.Title, p.Slug, p.Summary, p.CoverURL, p.ContentMd, p.ContentHTML,
p.Status, p.PublishedAt, p.CreatedAt, p.UpdatedAt, p.ReadingMinutes, encodeImages(in.Images),
encodeLinkCard(in.LinkCard))
if err != nil {
return p, err
}
id, err = res.LastInsertId()
if err != nil {
return p, err
}
}
p.ID = id
// 标签是长文的组织方式;短文(类推微博)一律不打标签
tags := in.Tags
if p.Kind == model.KindShort {
tags = nil
}
if err := s.setTags(id, tags); err != nil {
return p, err
}
return s.Get(id)
}
func (s *Store) Update(id int64, in model.PostInput) (model.Post, error) {
cur, err := s.Get(id)
if err != nil {
return cur, err
}
p := cur
if in.Kind != "" {
p.Kind = in.Kind
}
if in.Title != "" || in.Kind == model.KindShort {
p.Title = strings.TrimSpace(in.Title)
}
if in.Summary != "" {
p.Summary = strings.TrimSpace(in.Summary)
}
if in.ContentMd != "" {
p.ContentMd = in.ContentMd
p.ContentHTML = renderHTML(in.ContentMd)
}
if in.Status != "" {
p.Status = in.Status
}
if in.Slug != "" && in.Slug != cur.Slug {
p.Slug = s.uniqueSlug(in.Slug, id)
}
if in.PublishedAt != "" {
p.PublishedAt = in.PublishedAt
}
// CoverURL is taken verbatim from input. The frontend always sends the
// current value (empty string included), so any update round-trips with
// whatever the user last saved.
p.CoverURL = strings.TrimSpace(in.CoverURL)
// Images 同理:nil = 不修改(快发盒只发新帖,后台编辑器全量回传)
if in.Images != nil {
p.Images = in.Images
}
// LinkCard 由 admin 层按正文重算后传入;空卡片(含只有 URL)= 清空
if in.LinkCard != nil {
p.LinkCard = in.LinkCard
}
p.UpdatedAt = now()
if in.ReadingMinutes != nil && *in.ReadingMinutes > 0 {
p.ReadingMinutes = *in.ReadingMinutes
} else if in.ContentMd != "" {
p.ReadingMinutes = readingMinutes(in.ContentMd)
}
if _, err := s.db.Exec(s.db.Q(`UPDATE posts SET kind=?,title=?,slug=?,summary=?,cover_url=?,content_md=?,
content_html=?,status=?,published_at=?,updated_at=?,reading_minutes=?,images=?,link_card=? WHERE id=?`),
p.Kind, p.Title, p.Slug, p.Summary, p.CoverURL, p.ContentMd, p.ContentHTML, p.Status,
p.PublishedAt, p.UpdatedAt, p.ReadingMinutes, encodeImages(p.Images), encodeLinkCard(p.LinkCard), id); err != nil {
return p, err
}
// 短文一律无标签:切换类型或直接保存时都把旧标签清掉
if in.Tags != nil || p.Kind == model.KindShort {
tags := in.Tags
if p.Kind == model.KindShort {
tags = nil
}
if err := s.setTags(id, tags); err != nil {
return p, err
}
}
return s.Get(id)
}
// BulkUpdateStatus flips the status of many posts in a single transaction.
func (s *Store) BulkUpdateStatus(ids []int64, status string) (int, error) {
if len(ids) == 0 {
return 0, nil
}
status = NormalizeStatus(status)
placeholders := strings.TrimSuffix(strings.Repeat("?,", len(ids)), ",")
args := []any{status, now()}
for _, id := range ids {
args = append(args, id)
}
q := s.db.Q(`UPDATE posts SET status=?, updated_at=? WHERE id IN (` + placeholders + `)`)
res, err := s.db.Exec(q, args...)
if err != nil {
return 0, err
}
n, _ := res.RowsAffected()
return int(n), nil
}
// NormalizeStatus coerces an arbitrary input string into a known post status
// value, defaulting to draft when nothing matches.
func NormalizeStatus(s string) string {
switch strings.ToLower(strings.TrimSpace(s)) {
case model.StatusDraft, model.StatusPublished:
return strings.ToLower(strings.TrimSpace(s))
default:
return model.StatusDraft
}
}
func (s *Store) Delete(id int64) error {
return s.tx(func(tx *sql.Tx) error {
if _, err := tx.Exec(s.db.Q(`DELETE FROM post_tags WHERE post_id = ?`), id); err != nil {
return err
}
res, err := tx.Exec(s.db.Q(`DELETE FROM posts WHERE id = ?`), id)
if err != nil {
return err
}
if n, err := res.RowsAffected(); err == nil && n == 0 {
return ErrNotFound
}
return nil
})
}
// tx 把多语句写包进一个事务:中途失败整体回滚,不留半截状态。
func (s *Store) tx(fn func(tx *sql.Tx) error) error {
tx, err := s.db.Begin()
if err != nil {
return err
}
if err := fn(tx); err != nil {
_ = tx.Rollback()
return err
}
return tx.Commit()
}
// ---------- tags ----------
func (s *Store) setTags(postID int64, names []string) error {
return s.tx(func(tx *sql.Tx) error {
if _, err := tx.Exec(s.db.Q(`DELETE FROM post_tags WHERE post_id = ?`), postID); err != nil {
return err
}
seen := map[string]bool{}
for _, raw := range names {
name := strings.TrimSpace(raw)
if name == "" || seen[name] {
continue
}
seen[name] = true
tagID, err := s.upsertTagOn(tx, name)
if err != nil {
return err
}
if _, err := tx.Exec(s.db.Q(`INSERT INTO post_tags(post_id, tag_id) VALUES (?,?)`), postID, tagID); err != nil {
return err
}
}
return nil
})
}
// execer 让 upsertTag 在普通连接和事务里都能跑。
type execer interface {
Exec(query string, args ...any) (sql.Result, error)
QueryRow(query string, args ...any) *sql.Row
}
func (s *Store) upsertTag(name string) (int64, error) {
return s.upsertTagOn(s.db, name)
}
func (s *Store) upsertTagOn(x execer, name string) (int64, error) {
slug := Slugify(name)
if slug == "" {
slug = "tag"
}
var id int64
err := x.QueryRow(s.db.Q(`SELECT id FROM tags WHERE slug = ?`), slug).Scan(&id)
if err == nil {
return id, nil
}
if err != sql.ErrNoRows {
return 0, err
}
if s.db.Dialect == db.Postgres {
err = x.QueryRow(s.db.Q(`INSERT INTO tags(name,slug) VALUES (?,?) RETURNING id`), name, slug).Scan(&id)
return id, err
}
res, err := x.Exec(s.db.Q(`INSERT INTO tags(name,slug) VALUES (?,?)`), name, slug)
if err != nil {
return 0, err
}
return res.LastInsertId()
}
func (s *Store) ListTags() ([]model.Tag, error) {
q := s.db.Q(`SELECT t.id, t.name, t.slug, t.color, COUNT(pt.post_id) AS c
FROM tags t LEFT JOIN post_tags pt ON pt.tag_id = t.id
GROUP BY t.id, t.name, t.slug, t.color ORDER BY c DESC, t.name`)
rows, err := s.db.Query(q)
if err != nil {
return nil, err
}
defer rows.Close()
out := []model.Tag{}
for rows.Next() {
var t model.Tag
if err := rows.Scan(&t.ID, &t.Name, &t.Slug, &t.Color, &t.Count); err != nil {
return nil, err
}
out = append(out, t)
}
return out, rows.Err()
}
func (s *Store) CreateTag(name string) (model.Tag, error) {
name = strings.TrimSpace(name)
if name == "" {
return model.Tag{}, errors.New("tag name required")
}
id, err := s.upsertTag(name)
if err != nil {
return model.Tag{}, err
}
return model.Tag{ID: id, Name: name, Slug: Slugify(name)}, nil
}
// CreateTagFull is used by the admin API when color is supplied.
func (s *Store) CreateTagFull(name, color string) (model.Tag, error) {
name = strings.TrimSpace(name)
if name == "" {
return model.Tag{}, errors.New("tag name required")
}
slug := Slugify(name)
if slug == "" {
return model.Tag{}, errors.New("tag slug required")
}
var id int64
err := s.db.QueryRow(s.db.Q(`SELECT id FROM tags WHERE slug = ?`), slug).Scan(&id)
if err != nil && err != sql.ErrNoRows {
return model.Tag{}, err
}
if err == nil {
// already exists · update color and return
if _, err := s.db.Exec(s.db.Q(`UPDATE tags SET color=? WHERE id=?`), color, id); err != nil {
return model.Tag{}, err
}
return s.getTag(id)
}
color = strings.TrimSpace(color)
if s.db.Dialect == db.Postgres {
err := s.db.QueryRow(s.db.Q(`INSERT INTO tags(name,slug,color) VALUES (?,?,?) RETURNING id`),
name, slug, color).Scan(&id)
if err != nil {
return model.Tag{}, err
}
} else {
res, err := s.db.Exec(s.db.Q(`INSERT INTO tags(name,slug,color) VALUES (?,?,?)`), name, slug, color)
if err != nil {
return model.Tag{}, err
}
id, err = res.LastInsertId()
if err != nil {
return model.Tag{}, err
}
}
return s.getTag(id)
}
func (s *Store) getTag(id int64) (model.Tag, error) {
var t model.Tag
err := s.db.QueryRow(s.db.Q(`SELECT t.id, t.name, t.slug, t.color, COUNT(pt.post_id)
FROM tags t LEFT JOIN post_tags pt ON pt.tag_id = t.id
WHERE t.id = ? GROUP BY t.id`), id).Scan(&t.ID, &t.Name, &t.Slug, &t.Color, &t.Count)
if err == sql.ErrNoRows {
return t, ErrNotFound
}
return t, err
}
func (s *Store) RenameTag(id int64, name string) (model.Tag, error) {
name = strings.TrimSpace(name)
slug := Slugify(name)
if name == "" {
return model.Tag{}, errors.New("tag name required")
}
if _, err := s.db.Exec(s.db.Q(`UPDATE tags SET name=?, slug=? WHERE id=?`), name, slug, id); err != nil {
return model.Tag{}, err
}
return s.getTag(id)
}
// UpdateTag is the combined rename + recolor used by the admin UI.
func (s *Store) UpdateTag(id int64, name, color string) (model.Tag, error) {
name = strings.TrimSpace(name)
slug := Slugify(name)
if name == "" {
return model.Tag{}, errors.New("tag name required")
}
color = strings.TrimSpace(color)
if _, err := s.db.Exec(s.db.Q(`UPDATE tags SET name=?, slug=?, color=? WHERE id=?`), name, slug, color, id); err != nil {
return model.Tag{}, err
}
return s.getTag(id)
}
// MergeTags moves every post from fromID to toID, then deletes fromID. It is a
// no-op if the ids are equal or either doesn't exist.
func (s *Store) MergeTags(fromID, toID int64) (model.Tag, error) {
if fromID == toID {
return model.Tag{}, errors.New("source and target are the same tag")
}
if _, err := s.getTag(fromID); err != nil {
return model.Tag{}, err
}
if _, err := s.getTag(toID); err != nil {
return model.Tag{}, err
}
// De-duplicate before the move so we don't end up with two rows pointing
// at the same post.
if err := s.tx(func(tx *sql.Tx) error {
if _, err := tx.Exec(s.db.Q(`DELETE FROM post_tags
WHERE post_id IN (SELECT post_id FROM post_tags WHERE tag_id = ?)
AND tag_id = ?`), toID, toID); err != nil {
return err
}
if _, err := tx.Exec(s.db.Q(`UPDATE post_tags SET tag_id=? WHERE tag_id=?`), toID, fromID); err != nil {
return err
}
_, err := tx.Exec(s.db.Q(`DELETE FROM tags WHERE id=?`), fromID)
return err
}); err != nil {
return model.Tag{}, err
}
return s.getTag(toID)
}
func (s *Store) DeleteTag(id int64) error {
return s.tx(func(tx *sql.Tx) error {
if _, err := tx.Exec(s.db.Q(`DELETE FROM post_tags WHERE tag_id = ?`), id); err != nil {
return err
}
_, err := tx.Exec(s.db.Q(`DELETE FROM tags WHERE id = ?`), id)
return err
})
}
// ---------- archive ----------
func (s *Store) Archive() ([]model.ArchiveYear, error) {
var items []model.Post
if err := s.scanPostsInto(s.db.Q(`SELECT `+listCols+` FROM posts
WHERE status = ? ORDER BY published_at DESC`), []any{model.StatusPublished}, &items); err != nil {
return nil, err
}
years := []model.ArchiveYear{}
yearIdx := map[string]int{}
monthIdx := map[string]int{}
for _, p := range items {
y, m := splitDate(p.PublishedAt)
if y == "" {
continue
}
yi, ok := yearIdx[y]
if !ok {
years = append(years, model.ArchiveYear{Year: y, Months: []model.ArchiveMonth{}})
yi = len(years) - 1
yearIdx[y] = yi
}
key := y + "-" + m
mi, ok := monthIdx[key]
if !ok {
years[yi].Months = append(years[yi].Months, model.ArchiveMonth{Month: m, Posts: []model.Post{}})
mi = len(years[yi].Months) - 1
monthIdx[key] = mi
}
years[yi].Count++
years[yi].Months[mi].Posts = append(years[yi].Months[mi].Posts, p)
}
return years, nil
}
func splitDate(rfc3339 string) (string, string) {
if len(rfc3339) < 7 {
return "", ""
}
return rfc3339[0:4], rfc3339[5:7]
}
// ---------- dashboard ----------
// Dashboard produces the snapshot rendered on the admin home page.
func (s *Store) Dashboard() (model.Dashboard, error) {
d := model.Dashboard{
RecentPosts: []model.Post{},
RecentDrafts: []model.Post{},
TopTags: []model.Tag{},
PublishedByMonth: []model.MonthBucket{},
}
// status / kind counts --------------------------------------------------
rows, err := s.db.Query(s.db.Q(`SELECT status, kind, COUNT(*) FROM posts GROUP BY status, kind`))
if err != nil {
return d, err
}
for rows.Next() {
var status, kind string
var n int
if err := rows.Scan(&status, &kind, &n); err != nil {
rows.Close()
return d, err
}
d.TotalPosts += n
switch status {
case model.StatusPublished:
d.PublishedPosts += n
case model.StatusDraft:
d.DraftPosts += n
}
switch kind {
case model.KindShort:
d.ShortPosts += n
case model.KindLong:
d.LongPosts += n
}
}
rows.Close()
// total tag count -------------------------------------------------------
if err := s.db.QueryRow(s.db.Q(`SELECT COUNT(*) FROM tags`)).Scan(&d.TotalTags); err != nil {
return d, err
}
// total word count: sum of content_md length ----------------------------
if err := s.db.QueryRow(s.db.Q(`SELECT COALESCE(SUM(LENGTH(content_md)),0) FROM posts`)).Scan(&d.TotalWords); err != nil {
return d, err
}
// recent published posts (5) --------------------------------------------
if err := s.scanPostsInto(s.db.Q(`SELECT `+listCols+` FROM posts
WHERE status=? ORDER BY published_at DESC LIMIT 5`), []any{model.StatusPublished}, &d.RecentPosts); err != nil {
return d, err
}
// recent drafts (5) -----------------------------------------------------
if err := s.scanPostsInto(s.db.Q(`SELECT `+listCols+` FROM posts
WHERE status=? ORDER BY updated_at DESC LIMIT 5`), []any{model.StatusDraft}, &d.RecentDrafts); err != nil {
return d, err
}
// top tags (10) ---------------------------------------------------------
if err := s.db.QueryRow(s.db.Q(`SELECT COUNT(*) FROM tags`)).Scan(&d.TotalTags); err != nil {
return d, err
}
tagRows, err := s.db.Query(s.db.Q(`SELECT t.id, t.name, t.slug, t.color, COUNT(pt.post_id) AS c
FROM tags t LEFT JOIN post_tags pt ON pt.tag_id = t.id
GROUP BY t.id, t.name, t.slug, t.color ORDER BY c DESC, t.name LIMIT 10`))
if err != nil {
return d, err
}
for tagRows.Next() {
var t model.Tag
if err := tagRows.Scan(&t.ID, &t.Name, &t.Slug, &t.Color, &t.Count); err != nil {
tagRows.Close()
return d, err
}
d.TopTags = append(d.TopTags, t)
}
tagRows.Close()
// published_by_month · last 12 months -----------------------------------
monthRows, err := s.db.Query(s.db.Q(`SELECT substr(published_at,1,7) AS m, COUNT(*) AS c
FROM posts WHERE status=? AND length(published_at) >= 7
GROUP BY m ORDER BY m DESC LIMIT 12`), model.StatusPublished)
if err != nil {
return d, err
}
defer monthRows.Close()
for monthRows.Next() {
var b model.MonthBucket
if err := monthRows.Scan(&b.Month, &b.Count); err != nil {
return d, err
}
d.PublishedByMonth = append(d.PublishedByMonth, b)
}
// reverse so we can render left-to-right on the chart
for i, j := 0, len(d.PublishedByMonth)-1; i < j; i, j = i+1, j-1 {
d.PublishedByMonth[i], d.PublishedByMonth[j] = d.PublishedByMonth[j], d.PublishedByMonth[i]
}
return d, nil
}
// scanPostsInto is a small helper that runs a query and scans rows into the
// given slice, attaching tags at the end.
func (s *Store) scanPostsInto(query string, args []any, dst *[]model.Post) error {
rows, err := s.db.Query(query, args...)
if err != nil {
return err
}
defer rows.Close()
out := []model.Post{}
for rows.Next() {
p, err := scanPost(rows)
if err != nil {
return err
}
out = append(out, p)
}
if err := rows.Err(); err != nil {
return err
}
if err := s.attachTags(out); err != nil {
return err
}
*dst = out
return nil
}
// ---------- files ----------
func scanFile(sc interface{ Scan(...any) error }) (model.File, error) {
var f model.File
err := sc.Scan(&f.ID, &f.Key, &f.Name, &f.Mime, &f.Size, &f.SHA256, &f.Store, &f.CreatedAt)
return f, err
}
func (s *Store) CreateFile(f model.File) (model.File, error) {
f.CreatedAt = now()
q := s.db.Q(`INSERT INTO files (key,name,mime,size,sha256,store,created_at) VALUES (?,?,?,?,?,?,?)`)
var id int64
if s.db.Dialect == db.Postgres {
err := s.db.QueryRow(q+` RETURNING id`, f.Key, f.Name, f.Mime, f.Size, f.SHA256, f.Store, f.CreatedAt).Scan(&id)
if err != nil {
return model.File{}, err
}
} else {
res, err := s.db.Exec(q, f.Key, f.Name, f.Mime, f.Size, f.SHA256, f.Store, f.CreatedAt)
if err != nil {
return model.File{}, err
}
if id, err = res.LastInsertId(); err != nil {
return model.File{}, err
}
}
f.ID = id
return f, nil
}
// ListFiles 按创建时间倒序分页,q 模糊匹配原始文件名 / key。
func (s *Store) ListFiles(page, size int, q string) (model.FilePage, error) {
if page < 1 {
page = 1
}
if size < 1 || size > 100 {
size = 20
}
where := ""
var args []any
if query := strings.TrimSpace(q); query != "" {
where = ` WHERE name LIKE ? OR key LIKE ?`
like := "%" + query + "%"
args = append(args, like, like)
}
var total int
if err := s.db.QueryRow(s.db.Q(`SELECT COUNT(*) FROM files`+where), args...).Scan(&total); err != nil {
return model.FilePage{}, err
}
args = append(args, size, (page-1)*size)
rows, err := s.db.Query(s.db.Q(`SELECT id,key,name,mime,size,sha256,store,created_at FROM files`+
where+` ORDER BY created_at DESC, id DESC LIMIT ? OFFSET ?`), args...)
if err != nil {
return model.FilePage{}, err
}
defer rows.Close()
out := []model.File{}
for rows.Next() {
f, err := scanFile(rows)
if err != nil {
return model.FilePage{}, err
}
out = append(out, f)
}
return model.FilePage{Items: out, Total: total, Page: page, Size: size}, rows.Err()
}
func (s *Store) GetFile(id int64) (model.File, error) {
f, err := scanFile(s.db.QueryRow(s.db.Q(`SELECT id,key,name,mime,size,sha256,store,created_at FROM files WHERE id = ?`), id))
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return model.File{}, ErrNotFound
}
return model.File{}, err
}
return f, nil
}
func (s *Store) GetFileByKey(key string) (model.File, error) {
f, err := scanFile(s.db.QueryRow(s.db.Q(`SELECT id,key,name,mime,size,sha256,store,created_at FROM files WHERE key = ?`), key))
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return model.File{}, ErrNotFound
}
return model.File{}, err
}
return f, nil
}
// FileReferences 反查谁在用这个文件:文章(封面或正文)、项目封面、站主头像。
// 匹配的键是 files.key 本身而不是完整 URL——本地存成 /uploads/{key}、R2 存成
// {publicBase}/{key}、缩略图是 /uploads/thumb/{key},三种形态都以 key 结尾,
// 换存储端后正文里的老链接也照样能查到。
// 用 LIKE 现扫而不维护计数表:写入口有编辑器、外链转存、短文、项目、头像好几处,
// 计数一旦漂移就再也信不过;删文件是低频操作,扫全表几十毫秒换一个永远正确的答案。
func (s *Store) FileReferences(key string) ([]model.FileRef, error) {
if key == "" {
return []model.FileRef{}, nil
}
like := "%" + key + "%"
out := []model.FileRef{}
rows, err := s.db.Query(s.db.Q(`SELECT id,title,slug,status FROM posts
WHERE cover_url LIKE ? OR content_md LIKE ? OR content_html LIKE ?
ORDER BY published_at DESC LIMIT 100`), like, like, like)
if err != nil {
return nil, err
}
for rows.Next() {
var r model.FileRef
r.Kind = "post"
if err := rows.Scan(&r.ID, &r.Title, &r.Slug, &r.Status); err != nil {
rows.Close()
return nil, err
}
out = append(out, r)
}
if err := rows.Err(); err != nil {
rows.Close()
return nil, err
}
rows.Close()
proj, err := s.db.Query(s.db.Q(`SELECT id,title,slug,status FROM projects WHERE cover_url LIKE ? LIMIT 100`), like)
if err != nil {
return nil, err
}
for proj.Next() {
var r model.FileRef
r.Kind = "project"
if err := proj.Scan(&r.ID, &r.Title, &r.Slug, &r.Status); err != nil {
proj.Close()
return nil, err
}
out = append(out, r)
}
if err := proj.Err(); err != nil {
proj.Close()
return nil, err
}
proj.Close()
// 头像存的就是 key(不是 URL),等值比较
if st, err := s.GetSettings(); err != nil {
return nil, err
} else if st.AuthorAvatarKey == key {
out = append(out, model.FileRef{Kind: "avatar"})
}
return out, nil
}
// DeleteFile 删索引行并返回被删的行(调用方负责先删对象存储里的本体)。
func (s *Store) DeleteFile(id int64) (model.File, error) {
f, err := s.GetFile(id)
if err != nil {
return model.File{}, err
}
res, err := s.db.Exec(s.db.Q(`DELETE FROM files WHERE id = ?`), id)
if err != nil {
return model.File{}, err
}
if n, _ := res.RowsAffected(); n == 0 {
return model.File{}, ErrNotFound
}
return f, nil
}
// b2s 布尔转 KV 开关值
func b2s(b bool) string {
if b {
return "1"
}
return "0"
}
// ---------- readers(评论区的登录用户) ----------
// UpsertReader 按 (provider, handle) 找人:找到就更新资料,找不到就建档。
// banned 是状态位,不随资料更新覆盖。
func (s *Store) UpsertReader(r model.Reader) (model.Reader, error) {
r.CreatedAt = now()
q := s.db.Q(`INSERT INTO users (provider,handle,name,avatar_url,url,banned,created_at)
VALUES (?,?,?,?,?,0,?)
ON CONFLICT (provider,handle) DO UPDATE SET
name=excluded.name, avatar_url=excluded.avatar_url, url=excluded.url`)
// SQLite 的 ON CONFLICT 语法 Postgres 也认(现代版);老库退化走下面分支
if s.db.Dialect == db.Postgres {
if _, err := s.db.Exec(s.db.Q(`INSERT INTO users (provider,handle,name,avatar_url,url,banned,created_at)
VALUES (?,?,?,?,?,0,?) ON CONFLICT (provider,handle) DO UPDATE SET
name=excluded.name, avatar_url=excluded.avatar_url, url=excluded.url`),
r.Provider, r.Handle, r.Name, r.AvatarURL, r.URL, r.CreatedAt); err != nil {
return model.Reader{}, err
}
return s.GetReaderByProviderHandle(r.Provider, r.Handle)
}
if _, err := s.db.Exec(q, r.Provider, r.Handle, r.Name, r.AvatarURL, r.URL, r.CreatedAt); err != nil {
return model.Reader{}, err
}
return s.GetReaderByProviderHandle(r.Provider, r.Handle)
}
// userCols 是 users 表的读取列清单。列在多处 SELECT 复用,抽出来免得加一列
// 就要同步改一遍(漏一处就是扫错位置)。
const userCols = `id,provider,handle,name,avatar_url,url,banned,role,created_at`
// userColsU 是 JOIN 查询里带 u. 前缀的同一份列清单。和 userCols 成对改。
const userColsU = `u.id,u.provider,u.handle,u.name,u.avatar_url,u.url,u.banned,u.role,u.created_at`
func (s *Store) GetReader(id int64) (model.Reader, error) {
r, err := scanReader(s.db.QueryRow(s.db.Q(`SELECT `+userCols+` FROM users WHERE id = ?`), id))
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return model.Reader{}, ErrNotFound
}
return model.Reader{}, err
}
return r, nil
}
func (s *Store) GetReaderByProviderHandle(provider, handle string) (model.Reader, error) {
r, err := scanReader(s.db.QueryRow(s.db.Q(`SELECT `+userCols+` FROM users WHERE provider = ? AND handle = ?`), provider, handle))
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return model.Reader{}, ErrNotFound
}
return model.Reader{}, err
}
return r, nil
}
func (s *Store) SetReaderBanned(id int64, banned bool) error {
res, err := s.db.Exec(s.db.Q(`UPDATE users SET banned = ? WHERE id = ?`), b2i(banned), id)
if err != nil {
return err
}
if n, _ := res.RowsAffected(); n == 0 {
return ErrNotFound
}
return nil
}
// ListReaders 后台的用户列表:带评论数,站主最前、禁言中的排前面
func (s *Store) ListReaders() ([]model.Reader, error) {
rows, err := s.db.Query(s.db.Q(`SELECT u.id,u.provider,u.handle,u.name,u.avatar_url,u.url,u.banned,u.role,u.created_at,
(SELECT COUNT(*) FROM comments c WHERE c.user_id = u.id AND c.is_deleted = 0) AS cnt
FROM users u ORDER BY CASE WHEN u.role = 'owner' THEN 0 ELSE 1 END, u.banned DESC, u.created_at DESC`))
if err != nil {
return nil, err
}
defer rows.Close()
out := []model.Reader{}
for rows.Next() {
var r model.Reader
var cnt int64
if err := rows.Scan(&r.ID, &r.Provider, &r.Handle, &r.Name, &r.AvatarURL, &r.URL, &r.Banned, &r.Role, &r.CreatedAt, &cnt); err != nil {
return nil, err
}
r.CommentCount = cnt
out = append(out, r)
}
return out, rows.Err()
}
// UpdateProfile 改站主行的显示名。只碰 name —— provider/handle/role 是身份锚,
// 不随资料编辑变动。
func (s *Store) UpdateProfile(userID int64, name string) (model.Reader, error) {
if _, err := s.db.Exec(s.db.Q(`UPDATE users SET name = ? WHERE id = ?`), name, userID); err != nil {
return model.Reader{}, err
}
return s.GetReader(userID)
}
// ---------- 站主账号 / 身份绑定 / passkey ----------
// EnsureOwner 取(或创建)站主账号行:provider=admin、role=owner。
// handle 跟着 ONE_ADMIN_USER 走——改了环境变量后这里会同步,
// 但 role=owner 只此一行,是「哪些身份登录算管理员」的锚点。
func (s *Store) EnsureOwner(handle string) (model.Reader, error) {
cur, err := s.GetOwner()
if err == nil {
if cur.Handle != handle {
if _, uerr := s.db.Exec(s.db.Q(`UPDATE users SET handle = ? WHERE id = ?`), handle, cur.ID); uerr != nil {
return model.Reader{}, uerr
}
cur.Handle = handle
}
return cur, nil
}
if !errors.Is(err, ErrNotFound) {
return model.Reader{}, err
}
if _, err := s.db.Exec(s.db.Q(`INSERT INTO users (provider,handle,name,avatar_url,url,banned,role,created_at)
VALUES (?,?,?,?,?,0,'owner',?)`), "admin", handle, "", "", "", now()); err != nil {
return model.Reader{}, err
}
return s.GetOwner()
}
// GetOwner 取站主行。role='owner' 全库唯一,按 provider 兜底兼容老数据
// (老库里站主行只有 provider='admin',没有 role)。
func (s *Store) GetOwner() (model.Reader, error) {
r, err := scanReader(s.db.QueryRow(s.db.Q(
`SELECT ` + userCols + ` FROM users WHERE role = 'owner' OR provider = 'admin' ORDER BY id LIMIT 1`)))
if errors.Is(err, sql.ErrNoRows) {
return model.Reader{}, ErrNotFound
}
return r, err
}
// BindIdentity 把 (provider, extern_uid) 绑到某个用户上。
// 外部账号已被别人占用时返回 ErrConflict —— 调用方必须原样拒绝,
// 不能「后来者覆盖」,否则任何人都能抢先把别人的 GitHub 账号登记成自己的。
func (s *Store) BindIdentity(userID int64, provider, externUID, display string) error {
owner, err := s.GetUserByIdentity(provider, externUID)
if err == nil {
if owner.ID == userID {
return nil // 重复绑定同一个,幂等放过
}
return ErrConflict
}
if !errors.Is(err, ErrNotFound) {
return err
}
_, err = s.db.Exec(s.db.Q(`INSERT INTO user_identities (user_id,provider,extern_uid,display,created_at)
VALUES (?,?,?,?,?)`), userID, provider, externUID, display, now())
return err
}
// UnbindIdentity 解绑某平台的绑定。站主始终还有环境变量密码这条退路,
// 所以这里不需要「不能解绑唯一登录方式」的护栏。
func (s *Store) UnbindIdentity(userID int64, provider string) error {
res, err := s.db.Exec(s.db.Q(`DELETE FROM user_identities WHERE user_id = ? AND provider = ?`), userID, provider)
if err != nil {
return err
}
if n, _ := res.RowsAffected(); n == 0 {
return ErrNotFound
}
return nil
}
func (s *Store) ListIdentities(userID int64) ([]model.UserIdentity, error) {
rows, err := s.db.Query(s.db.Q(`SELECT id,user_id,provider,extern_uid,display,created_at
FROM user_identities WHERE user_id = ? ORDER BY created_at`), userID)
if err != nil {
return nil, err
}
defer rows.Close()
out := []model.UserIdentity{}
for rows.Next() {
var it model.UserIdentity
if err := rows.Scan(&it.ID, &it.UserID, &it.Provider, &it.ExternUID, &it.Display, &it.CreatedAt); err != nil {
return nil, err
}
out = append(out, it)
}
return out, rows.Err()
}
// GetUserByIdentity 按外部身份找用户:登录时先问它,命中即知道该发哪种会话。
func (s *Store) GetUserByIdentity(provider, externUID string) (model.Reader, error) {
r, err := scanReader(s.db.QueryRow(s.db.Q(
`SELECT `+userColsU+`
FROM user_identities i JOIN users u ON u.id = i.user_id
WHERE i.provider = ? AND i.extern_uid = ?`), provider, externUID))
if errors.Is(err, sql.ErrNoRows) {
return model.Reader{}, ErrNotFound
}
return r, err
}
// ---------- passkey ----------
func (s *Store) AddPasskey(p model.Passkey) (model.Passkey, error) {
p.CreatedAt = now()
res, err := s.db.Exec(s.db.Q(`INSERT INTO passkeys (user_id,credential_id,public_key,sign_count,name,created_at,last_used_at)
VALUES (?,?,?,?,?,?,?)`), p.UserID, p.CredentialID, p.PublicKey, int64(p.SignCount), p.Name, p.CreatedAt, "")
if err != nil {
return model.Passkey{}, err
}
p.ID, _ = res.LastInsertId()
return p, nil
}
// ListPasskeys 不返回 public_key:管理页只列名字与时间,凭据公钥
// 没必要顺着列表接口到处走。
func (s *Store) ListPasskeys(userID int64) ([]model.Passkey, error) {
rows, err := s.db.Query(s.db.Q(`SELECT id,user_id,credential_id,'',sign_count,name,created_at,last_used_at
FROM passkeys WHERE user_id = ? ORDER BY created_at`), userID)
if err != nil {
return nil, err
}
defer rows.Close()
out := []model.Passkey{}
for rows.Next() {
var p model.Passkey
var sc int64
if err := rows.Scan(&p.ID, &p.UserID, &p.CredentialID, &p.PublicKey, &sc, &p.Name, &p.CreatedAt, &p.LastUsedAt); err != nil {
return nil, err
}
p.SignCount = uint32(sc)
out = append(out, p)
}
return out, rows.Err()
}
func (s *Store) GetPasskeyByCredentialID(credentialID string) (model.Passkey, error) {
var p model.Passkey
var sc int64
err := s.db.QueryRow(s.db.Q(`SELECT id,user_id,credential_id,public_key,sign_count,name,created_at,last_used_at
FROM passkeys WHERE credential_id = ?`), credentialID).
Scan(&p.ID, &p.UserID, &p.CredentialID, &p.PublicKey, &sc, &p.Name, &p.CreatedAt, &p.LastUsedAt)
if errors.Is(err, sql.ErrNoRows) {
return model.Passkey{}, ErrNotFound
}
p.SignCount = uint32(sc)
return p, err
}
// TouchPasskey 回写签名计数与使用时间。计数只增不减:
// 新计数比库里记的小,说明凭据被克隆到多个 authenticator 上用过。
func (s *Store) TouchPasskey(id int64, signCount uint32) error {
_, err := s.db.Exec(s.db.Q(`UPDATE passkeys SET sign_count = ?, last_used_at = ? WHERE id = ?`),
int64(signCount), now(), id)
return err
}
// DeletePasskey 带 user_id 条件删:免得拿别人的 id 越权删凭据。
func (s *Store) DeletePasskey(id, userID int64) error {
res, err := s.db.Exec(s.db.Q(`DELETE FROM passkeys WHERE id = ? AND user_id = ?`), id, userID)
if err != nil {
return err
}
if n, _ := res.RowsAffected(); n == 0 {
return ErrNotFound
}
return nil
}
func b2i(b bool) int64 {
if b {
return 1
}
return 0
}
func scanReader(sc interface{ Scan(...any) error }) (model.Reader, error) {
var r model.Reader
var banned int64
err := sc.Scan(&r.ID, &r.Provider, &r.Handle, &r.Name, &r.AvatarURL, &r.URL, &banned, &r.Role, &r.CreatedAt)
r.Banned = banned == 1
return r, err
}
// ---------- comments ----------
const commentCols = `c.id, c.post_id, c.parent_id, c.root_id, c.user_id, c.body_md, c.body_html,
c.status, c.is_deleted, c.created_at, c.edited_at,
u.id, u.provider, u.handle, u.name, u.avatar_url, u.url, u.banned, u.created_at`
func scanComment(sc interface{ Scan(...any) error }) (model.Comment, error) {
var c model.Comment
var u model.Reader
var banned int64
err := sc.Scan(&c.ID, &c.PostID, &c.ParentID, &c.RootID, &c.UserID, &c.BodyMd, &c.BodyHTML,
&c.Status, &c.IsDeleted, &c.CreatedAt, &c.EditedAt,
&u.ID, &u.Provider, &u.Handle, &u.Name, &u.AvatarURL, &u.URL, &banned, &u.CreatedAt)
if err != nil {
return c, err
}
c.User = &model.Reader{ID: u.ID, Provider: u.Provider, Handle: u.Handle, Name: u.Name,
AvatarURL: u.AvatarURL, URL: u.URL, Banned: banned == 1, CreatedAt: u.CreatedAt}
c.Replies = []model.Comment{}
return c, nil
}
// CreateComment 新建评论;parent/root 归属与审核状态由调用方决定
func (s *Store) CreateComment(c model.Comment) (model.Comment, error) {
c.CreatedAt = now()
q := s.db.Q(`INSERT INTO comments (post_id,user_id,parent_id,root_id,body_md,body_html,status,created_at)
VALUES (?,?,?,?,?,?,?,?)`)
var id int64
if s.db.Dialect == db.Postgres {
err := s.db.QueryRow(q+` RETURNING id`, c.PostID, c.UserID, c.ParentID, c.RootID,
c.BodyMd, c.BodyHTML, c.Status, c.CreatedAt).Scan(&id)
if err != nil {
return model.Comment{}, err
}
} else {
res, err := s.db.Exec(q, c.PostID, c.UserID, c.ParentID, c.RootID, c.BodyMd, c.BodyHTML, c.Status, c.CreatedAt)
if err != nil {
return model.Comment{}, err
}
if id, err = res.LastInsertId(); err != nil {
return model.Comment{}, err
}
}
return s.GetComment(id)
}
// GetComment 单条(含用户)
func (s *Store) GetComment(id int64) (model.Comment, error) {
c, err := scanComment(s.db.QueryRow(s.db.Q(`SELECT `+commentCols+` FROM comments c
JOIN users u ON u.id = c.user_id WHERE c.id = ?`), id))
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return model.Comment{}, ErrNotFound
}
return model.Comment{}, err
}
return c, nil
}
// ListCommentsByPost 一篇文章的公开评论树:顶层可见 + 访客自己的待审;
// 每条顶层内嵌全部可见回复。deleted 的行保留(墓碑:正文清空,楼层不塌)。
func (s *Store) ListCommentsByPost(postID, viewerID int64, newestFirst bool) ([]model.Comment, error) {
order := `ASC`
if newestFirst {
order = `DESC`
}
rows, err := s.db.Query(s.db.Q(`SELECT `+commentCols+` FROM comments c
JOIN users u ON u.id = c.user_id
WHERE c.post_id = ? AND c.parent_id = 0
AND (c.status = 'visible' OR (c.status = 'pending' AND c.user_id = ?))
ORDER BY c.created_at `+order), postID, viewerID)
if err != nil {
return nil, err
}
defer rows.Close()
roots := []model.Comment{}
idx := map[int64]int{}
for rows.Next() {
c, err := scanComment(rows)
if err != nil {
return nil, err
}
c.Replies = []model.Comment{}
idx[c.ID] = len(roots)
roots = append(roots, c)
}
if err := rows.Err(); err != nil {
return nil, err
}
// 回复:可见的(+访客自己的待审),按时间正序挂在各自根下
rows2, err := s.db.Query(s.db.Q(`SELECT `+commentCols+` FROM comments c
JOIN users u ON u.id = c.user_id
WHERE c.post_id = ? AND c.parent_id <> 0
AND (c.status = 'visible' OR (c.status = 'pending' AND c.user_id = ?))
ORDER BY c.created_at ASC`), postID, viewerID)
if err != nil {
return nil, err
}
defer rows2.Close()
for rows2.Next() {
c, err := scanComment(rows2)
if err != nil {
return nil, err
}
if at, ok := idx[c.RootID]; ok {
roots[at].Replies = append(roots[at].Replies, c)
}
}
for i := range roots {
roots[i].ReplyCount = len(roots[i].Replies)
}
return roots, rows2.Err()
}
// ListCommentsAdmin 后台的评论列表(平铺,含用户与文章标题),status 过滤
func (s *Store) ListCommentsAdmin(status string, page, size int) (model.CommentPage, error) {
if page < 1 {
page = 1
}
if size < 1 || size > 100 {
size = 20
}
where := ""
var args []any
switch status {
case "pending", "visible":
where = ` WHERE c.status = '` + status + `' AND c.is_deleted = 0`
default:
where = ` WHERE c.is_deleted = 0`
}
var total int
if err := s.db.QueryRow(s.db.Q(`SELECT COUNT(*) FROM comments c`+where), args...).Scan(&total); err != nil {
return model.CommentPage{}, err
}
args = append(args, size, (page-1)*size)
rows, err := s.db.Query(s.db.Q(`SELECT `+commentCols+`, COALESCE(p.title, '') FROM comments c
JOIN users u ON u.id = c.user_id
JOIN posts p ON p.id = c.post_id`+where+`
ORDER BY c.created_at DESC LIMIT ? OFFSET ?`), args...)
if err != nil {
return model.CommentPage{}, err
}
defer rows.Close()
out := []model.Comment{}
for rows.Next() {
c, title, err := scanCommentAdmin(rows)
if err != nil {
return model.CommentPage{}, err
}
c.PostTitle = title
out = append(out, c)
}
return model.CommentPage{Items: out, Total: total, Page: page, Size: size}, rows.Err()
}
// CommentPage 后台评论管理的分页容器
type modelCommentPageAlias = struct{}
func scanCommentAdmin(sc interface{ Scan(...any) error }) (model.Comment, string, error) {
var c model.Comment
var u model.Reader
var title string
var banned int64
err := sc.Scan(&c.ID, &c.PostID, &c.ParentID, &c.RootID, &c.UserID, &c.BodyMd, &c.BodyHTML,
&c.Status, &c.IsDeleted, &c.CreatedAt, &c.EditedAt,
&u.ID, &u.Provider, &u.Handle, &u.Name, &u.AvatarURL, &u.URL, &banned, &u.CreatedAt, &title)
c.User = &model.Reader{ID: u.ID, Provider: u.Provider, Handle: u.Handle, Name: u.Name,
AvatarURL: u.AvatarURL, URL: u.URL, Banned: banned == 1, CreatedAt: u.CreatedAt}
return c, title, err
}
// SetCommentStatus 审核通过 / 退回待审
func (s *Store) SetCommentStatus(id int64, status string) error {
res, err := s.db.Exec(s.db.Q(`UPDATE comments SET status = ? WHERE id = ?`), status, id)
if err != nil {
return err
}
if n, _ := res.RowsAffected(); n == 0 {
return ErrNotFound
}
return nil
}
// DeleteComment 软删:留壳(「该评论已删除」),正文清空
func (s *Store) DeleteComment(id int64) error {
res, err := s.db.Exec(s.db.Q(`UPDATE comments SET is_deleted = 1, body_md = '', body_html = '' WHERE id = ?`), id)
if err != nil {
return err
}
if n, _ := res.RowsAffected(); n == 0 {
return ErrNotFound
}
return nil
}
// UpdateCommentBody 编辑后的正文回写
func (s *Store) UpdateCommentBody(id int64, md, html string) error {
res, err := s.db.Exec(s.db.Q(`UPDATE comments SET body_md = ?, body_html = ?, edited_at = ? WHERE id = ?`),
md, html, now(), id)
if err != nil {
return err
}
if n, _ := res.RowsAffected(); n == 0 {
return ErrNotFound
}
return nil
}