Files
artex/db/notification_delivery.go
dela 0335d572de
ci / go (push) Waiting to run
ci / go-db (agent) (push) Waiting to run
ci / go-db (config) (push) Waiting to run
ci / go-db (db) (push) Waiting to run
ci / go-db (evidence) (push) Waiting to run
ci / go-db (llmrec) (push) Waiting to run
ci / go-db (server) (push) Waiting to run
detections / detections (push) Waiting to run
web / web (push) Waiting to run
docs / links (push) Canceled after 0s
First Commit
2026-10-09 08:38:16 +08:00

447 lines
18 KiB
Go
Raw Permalink 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 db
import (
"context"
"database/sql"
"encoding/json"
"fmt"
"strings"
"time"
)
// 本文件是投递任务的领取与状态流转。
//
// 领取用「租约」而非长事务:把行置为 sending 并把 next_attempt_at 推到未来作为
// 租约到期时间,提交事务后再去做网络投递。这样投递期间不持有数据库锁——
// 网络请求可能耗时数秒(客户端超时 15 秒),占着行锁不放会拖垮同库的其它写操作。
//
// 代价是进程若在投递途中崩溃,行会停在 sending。这是**可自愈**的:租约到期后
// next_attempt_at 落入过去,下一轮领取会把同一行重新捞起来(见领取条件里的
// state IN ('pending','sending'))。重试计数在领取时就已 +1,所以崩溃不会造成
// 无限重试——MaxNotifyAttempts 次机会用完后落入 failed 等人工处理。
// MaxNotifyAttempts 是一条投递的最大尝试次数(含首次)。
// 定义在这里而非投递引擎里:它是状态机自身的策略,引擎只是执行者。
const MaxNotifyAttempts = 3
// MaxDigestBatchSize 是单个汇总批次一次最多合并多少条投递。
//
// 存在的理由是资源:一个汇总周期内如果扫出几万个漏洞(完全可能——一次全量扫描
// 就能做到),不设上界的话领取会把全部行读进内存、渲染成一条超长消息,
// 然后被渠道的长度上限截掉大半——既浪费内存,又**静默丢失**被截掉的那些漏洞。
// 设上界后,超出的部分留在库里成为下一个批次,下个周期自然发出去,不会丢。
//
// 取 500 的依据:它是渲染成消息后在企微 4096 字节上限内还"有内容可读"的量级;
// 再大也只是让截断发生在更靠后的位置而已。
const MaxDigestBatchSize = 500
// NotificationDelivery 是一条投递任务,含渲染所需的渠道配置与事件快照。
type NotificationDelivery struct {
ID int64 `json:"id"`
EventID int64 `json:"event_id"`
ChannelID int64 `json:"channel_id"`
State string `json:"state"`
Attempts int `json:"attempts"`
NextAttemptAt time.Time `json:"next_attempt_at"`
LastError string `json:"last_error"`
BatchID *int64 `json:"batch_id,omitempty"`
CreatedAt time.Time `json:"created_at"`
SentAt *time.Time `json:"sent_at,omitempty"`
Snapshot json.RawMessage `json:"snapshot,omitempty"`
// 联合加载的渲染上下文,不进 JSON(由 server 层组装 DTO)。
Channel *NotificationChannel `json:"-"`
// FindingID/EventKind 从事件带出,供历史列表直接跳转漏洞详情。
FindingID int64 `json:"finding_id,string"`
EventKind string `json:"event_kind"`
// ChannelName/ChannelKind 是列表展示用的冗余字段,省掉前端二次查询。
ChannelName string `json:"channel_name"`
ChannelKind string `json:"channel_kind"`
}
const notificationDeliveryCols = `d.id, d.event_id, d.channel_id, d.state, d.attempts, d.next_attempt_at,
d.last_error, d.batch_id, d.created_at, d.sent_at`
// joinedDeliveryQuery 是投递行的统一读取形状:投递 + 事件快照 + 渠道配置。
// 渲染一条消息三者缺一不可,分开查会写出三次往返。
const joinedDeliveryQuery = `SELECT ` + notificationDeliveryCols + `,
e.snapshot, e.kind, e.finding_id,
c.id, c.name, c.kind, c.enabled, c.config, c.mode, c.filter, c.rate_per_min
FROM notification_deliveries d
JOIN notification_events e ON e.id = d.event_id
JOIN notification_channels c ON c.id = d.channel_id`
func scanNotificationDelivery(sc interface{ Scan(...any) error }) (*NotificationDelivery, error) {
var (
dl NotificationDelivery
lastErr sql.NullString
batchID sql.NullInt64
sentAt sql.NullTime
snapshot []byte
eventKind string
channel NotificationChannel
chEnabled bool
)
if err := sc.Scan(&dl.ID, &dl.EventID, &dl.ChannelID, &dl.State, &dl.Attempts, &dl.NextAttemptAt,
&lastErr, &batchID, &dl.CreatedAt, &sentAt,
&snapshot, &eventKind, &dl.FindingID,
&channel.ID, &channel.Name, &channel.Kind, &chEnabled, &channel.Config, &channel.Mode, &channel.Filter, &channel.RatePerMin); err != nil {
return nil, err
}
dl.LastError = lastErr.String
if batchID.Valid {
dl.BatchID = &batchID.Int64
}
if sentAt.Valid {
dl.SentAt = &sentAt.Time
}
dl.Snapshot = json.RawMessage(snapshot)
dl.EventKind = eventKind
dl.ChannelName = channel.Name
dl.ChannelKind = channel.Kind
channel.Enabled = &chEnabled
dl.Channel = &channel
return &dl, nil
}
// claimQuery 描述一次领取:先按 sel 选出候选并加锁,再把它们置为 sending 并
// 延长租约。sel 里的 lease 位置由调用方用 $n 占位并自行传参。
type claimQuery struct {
sql string
args []any
}
// ClaimRealtimeDeliveries 领取某渠道一批到期的实时投递,最多 limit 条。
//
// 刻意按**单个渠道**领取而不是「全局领一批再挑着发」:限流闸在投递引擎里按渠道
// 维护,只有先知道这个渠道这一轮还能发几条、再去领同样多的行,限流才不会消耗
// 重试次数。若反过来先领后弃,被限流挡下的行已经被计过一次 attempts,
// 3 次预算会被纯粹的等待耗光,最后落进 failed。
//
// 条件含「租约已过期的 sending」——那是崩溃自愈的落点。lease 必须显著大于单次
// 投递的最坏耗时(渠道 HTTP 客户端超时 15 秒),否则同一行会被两个 dispatcher
// 同时投递。同时挡掉已停用渠道:停用操作已把存量投递标记为 skipped,
// 这里再拦一道,避免停用与领取并发时的漏网。
func (d *DB) ClaimRealtimeDeliveries(ctx context.Context, channelID int64, limit int, lease time.Duration) ([]*NotificationDelivery, error) {
if limit <= 0 {
return nil, nil
}
return d.claimDeliveries(ctx, lease, claimQuery{
sql: `SELECT dd.id FROM notification_deliveries dd
JOIN notification_channels c ON c.id = dd.channel_id
WHERE dd.channel_id = $1 AND dd.state IN ($2,$3) AND dd.next_attempt_at <= now()
AND c.enabled AND c.mode = $4
ORDER BY dd.next_attempt_at, dd.id
FOR UPDATE OF dd SKIP LOCKED
LIMIT $5`,
args: []any{channelID, NotifyStatePending, NotifyStateSending, NotifyModeRealtime, limit},
}, nil)
}
// DigestBatchDue 报告该渠道是否已攒够一个到期批次:存在待发投递,且**最老的那条**
// 年龄已达到汇总周期。
//
// 判定依据是最老投递的年龄而非墙上时钟:这样刚建好的渠道不会因为对齐到整点而
// 立刻吐出一条只有一条的「汇总」,积压很久的批次也不会再白等一轮。
//
// 与 ClaimDigestBatch 分开是因为语义不同:本函数只回答「该不该发」,
// 而领取要拿走该渠道**全部**待发行(包括尚未满年龄的那些)——否则一个周期
// 会被拆成多条消息,汇总就失去意义了。
func (d *DB) DigestBatchDue(ctx context.Context, channelID int64, minAge time.Duration) (bool, error) {
var due bool
err := d.QueryRowContext(ctx, `SELECT EXISTS (
SELECT 1 FROM notification_deliveries d
JOIN notification_channels c ON c.id = d.channel_id
WHERE d.channel_id = $1 AND d.state IN ($2,$3) AND c.enabled
GROUP BY d.channel_id
HAVING min(d.created_at) <= now() - make_interval(secs => $4)
)`, channelID, NotifyStatePending, NotifyStateSending, int64(minAge.Seconds())).Scan(&due)
return due, err
}
// ClaimDigestBatch 领取某渠道当前到期的待发投递,作为一个汇总批次,
// 单批最多 MaxDigestBatchSize 条。
//
// 同批次的所有投递共享 batch_id,用集合里的最小 id 作批次号(稳定、可读、
// 无需额外序列)。重试时用 COALESCE 保留原批次号,使「这批 N 条是一起发的」
// 在多次重试后依然成立。
//
// 按 id 升序取前 N 条而非随机取:最早产生的投递最先发出去,积压时不会出现
// 「新漏洞先发、老漏洞永远排在后面」的饥饿。
func (d *DB) ClaimDigestBatch(ctx context.Context, channelID int64, limit int, lease time.Duration) ([]*NotificationDelivery, error) {
if limit <= 0 {
return nil, nil
}
// limit 是**内存上界**,调用方传 MaxDigestBatchSize;这里再夹一道,
// 防止调用方传进一个更大的值。
//
// 刻意不接受「限流额度」充当批次大小:限流的单位是消息条数——一个批次只发
// 一条消息、消耗一个令牌,由 server 层的 takeTokens 扣除——与「一批装几条
// 漏洞」是两个不同的量纲。曾经为了让 rate_per_min 对 digest 生效而把每轮
// 请求预算传进来当批次大小,结果 rate=20/min 的渠道每批只装 1 条漏洞,
// digest 退化成带汇总文案的实时推送。要改限流请改 takeTokens 的 want,
// 不要动这里。
if limit > MaxDigestBatchSize {
limit = MaxDigestBatchSize
}
out, err := d.claimDeliveries(ctx, lease, claimQuery{
sql: `SELECT dd.id FROM notification_deliveries dd
JOIN notification_channels c ON c.id = dd.channel_id
WHERE dd.channel_id = $1 AND dd.state IN ($2,$3) AND dd.next_attempt_at <= now() AND c.enabled
ORDER BY dd.id
FOR UPDATE OF dd SKIP LOCKED
LIMIT $4`,
args: []any{channelID, NotifyStatePending, NotifyStateSending, limit},
}, func(tx *sql.Tx, ids []int64) error {
batchID := ids[0]
for _, id := range ids {
if id < batchID {
batchID = id
}
}
ph, idArgs := placeholders(2, ids)
_, err := tx.ExecContext(ctx, `UPDATE notification_deliveries SET batch_id = COALESCE(batch_id, $1)
WHERE id IN (`+ph+`)`, append([]any{batchID}, idArgs...)...)
return err
})
return out, err
}
// claimDeliveries 执行「选取 + 置 sending 延长租约 + 读取完整行」,全在一个事务里。
// postClaim 是可选的附加步骤(汇总批次用它写入 batch_id)。
func (d *DB) claimDeliveries(ctx context.Context, lease time.Duration, cq claimQuery, postClaim func(*sql.Tx, []int64) error) ([]*NotificationDelivery, error) {
tx, err := d.BeginTx(ctx, nil)
if err != nil {
return nil, err
}
defer tx.Rollback() //nolint:errcheck // 提交成功后是 no-op
ids, err := selectForClaim(ctx, tx, cq.sql, cq.args...)
if err != nil {
return nil, err
}
if len(ids) == 0 {
return nil, tx.Commit()
}
// 置 sending 并把 next_attempt_at 推到未来:这个未来时刻即租约到期时间,
// 「租约未到期」与「未到重试时间」因此共用同一个条件表达,不需要新增列。
ph, idArgs := placeholders(3, ids)
if _, err := tx.ExecContext(ctx, `UPDATE notification_deliveries
SET state=$1, attempts=attempts+1, next_attempt_at=now()+make_interval(secs => $2)
WHERE id IN (`+ph+`)`,
append([]any{NotifyStateSending, lease.Seconds()}, idArgs...)...); err != nil {
return nil, err
}
if postClaim != nil {
if err := postClaim(tx, ids); err != nil {
return nil, err
}
}
out, err := loadDeliveriesTx(ctx, tx, ids)
if err != nil {
return nil, err
}
return out, tx.Commit()
}
func selectForClaim(ctx context.Context, tx *sql.Tx, query string, args ...any) ([]int64, error) {
rows, err := tx.QueryContext(ctx, query, args...)
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()
}
func loadDeliveriesTx(ctx context.Context, tx *sql.Tx, ids []int64) ([]*NotificationDelivery, error) {
ph, args := placeholders(1, ids)
rows, err := tx.QueryContext(ctx, joinedDeliveryQuery+` WHERE d.id IN (`+ph+`) ORDER BY d.id`, args...)
if err != nil {
return nil, err
}
defer rows.Close()
out := []*NotificationDelivery{}
for rows.Next() {
dl, err := scanNotificationDelivery(rows)
if err != nil {
return nil, err
}
out = append(out, dl)
}
return out, rows.Err()
}
// MarkDeliveriesSent 把一批投递标记为已送达。
func (d *DB) MarkDeliveriesSent(ctx context.Context, ids []int64) error {
ph, args := placeholders(2, ids)
if len(args) == 0 {
return nil
}
_, err := d.ExecContext(ctx, `UPDATE notification_deliveries
SET state=$1, sent_at=now(), last_error='' WHERE id IN (`+ph+`)`, append([]any{NotifyStateSent}, args...)...)
return err
}
// RescheduleDeliveries 把一批投递退回 pending 并推后重试时间。
//
// 退回 pending 而不是引入新的中间状态,是为了让「还剩几次机会」只由一个地方
// 表达(MaxNotifyAttempts),避免状态机的分支随重试策略膨胀。
func (d *DB) RescheduleDeliveries(ctx context.Context, ids []int64, delay time.Duration, errMsg string) error {
ph, args := placeholders(4, ids)
if len(args) == 0 {
return nil
}
_, err := d.ExecContext(ctx, `UPDATE notification_deliveries
SET state=$1, next_attempt_at=now()+make_interval(secs => $2), last_error=$3
WHERE id IN (`+ph+`)`,
append([]any{NotifyStatePending, delay.Seconds(), truncateNotifyError(errMsg)}, args...)...)
return err
}
// DeferDeliveries 把一批投递退回 pending、立即可再领,并**撤销领取时计的那一次尝试**。
//
// 用途只有一个:汇总消息按渠道长度上限分段发送时,没装进本条的条目要留到下一批。
// 那不是失败,所以不该消耗重试预算——领取时 attempts 已经乐观地 +1 了,
// 这里必须减回去。否则一个 500 条的积压会按每段 20 条切成 25 段,
// 尾部条目在第 3 段就被 MaxNotifyAttempts 判成 failed,而它们从未出过任何错。
//
// GREATEST(...,0) 兜住「有人手工重发把 attempts 清零后又走到这里」的情况,
// 不让计数变成负数。
func (d *DB) DeferDeliveries(ctx context.Context, ids []int64, reason string) error {
ph, args := placeholders(3, ids)
if len(args) == 0 {
return nil
}
_, err := d.ExecContext(ctx, `UPDATE notification_deliveries
SET state=$1, attempts=GREATEST(attempts-1, 0), next_attempt_at=now(), last_error=$2
WHERE id IN (`+ph+`)`,
append([]any{NotifyStatePending, truncateNotifyError(reason)}, args...)...)
return err
}
// FailDeliveries 把一批投递标记为最终失败,等待人工在投递历史里重发。
func (d *DB) FailDeliveries(ctx context.Context, ids []int64, errMsg string) error {
// 占位符从 $3 开始:$1 是 state、$2 是 last_error。
ph, args := placeholders(3, ids)
if len(args) == 0 {
return nil
}
_, err := d.ExecContext(ctx, `UPDATE notification_deliveries SET state=$1, last_error=$2 WHERE id IN (`+ph+`)`,
append([]any{NotifyStateFailed, truncateNotifyError(errMsg)}, args...)...)
return err
}
// RetryNotificationDelivery 手动重发一条投递:重置为 pending、清零重试计数、
// 立即到期。清计数是刻意的——人工点「重发」意味着前几次失败的原因已被处理,
// 再拿旧计数限制它没有道理。
func (d *DB) RetryNotificationDelivery(ctx context.Context, id int64) error {
res, err := d.ExecContext(ctx, `UPDATE notification_deliveries
SET state=$2, attempts=0, next_attempt_at=now(), last_error=''
WHERE id=$1 AND state IN ($3,$4)`, id, NotifyStatePending, NotifyStateFailed, NotifyStateSkipped)
if err != nil {
return err
}
if n, _ := res.RowsAffected(); n == 0 {
return fmt.Errorf("投递 %d 不存在或当前状态不允许重发", id)
}
return nil
}
// NotificationDeliveryFilter 是投递历史的查询条件。
type NotificationDeliveryFilter struct {
ChannelID int64
State string
EventKind string
}
func (f NotificationDeliveryFilter) where() (string, []any) {
var conds []string
var args []any
if f.ChannelID > 0 {
args = append(args, f.ChannelID)
conds = append(conds, fmt.Sprintf("d.channel_id=$%d", len(args)))
}
if f.State != "" {
args = append(args, f.State)
conds = append(conds, fmt.Sprintf("d.state=$%d", len(args)))
}
if f.EventKind != "" {
args = append(args, f.EventKind)
conds = append(conds, fmt.Sprintf("e.kind=$%d", len(args)))
}
if len(conds) == 0 {
return "", nil
}
return " WHERE " + strings.Join(conds, " AND "), args
}
// ListNotificationDeliveries 分页返回投递历史,新的在前。
func (d *DB) ListNotificationDeliveries(ctx context.Context, f NotificationDeliveryFilter, page, pageSize int) ([]*NotificationDelivery, int, error) {
if page < 1 {
page = 1
}
if pageSize <= 0 || pageSize > 200 {
pageSize = 50
}
where, args := f.where()
var total int
if err := d.QueryRowContext(ctx, `SELECT count(*) FROM notification_deliveries d
JOIN notification_events e ON e.id = d.event_id`+where, args...).Scan(&total); err != nil {
return nil, 0, err
}
q := fmt.Sprintf("%s%s ORDER BY d.id DESC LIMIT $%d OFFSET $%d",
joinedDeliveryQuery, where, len(args)+1, len(args)+2)
rows, err := d.QueryContext(ctx, q, append(args, pageSize, (page-1)*pageSize)...)
if err != nil {
return nil, 0, err
}
defer rows.Close()
out := []*NotificationDelivery{}
for rows.Next() {
dl, err := scanNotificationDelivery(rows)
if err != nil {
return nil, 0, err
}
out = append(out, dl)
}
return out, total, rows.Err()
}
// truncateNotifyError 把错误信息截到列可接受的长度。渠道返回的响应体可能很长
// (通用 Webhook 打到自建服务时尤甚),不截断会让历史列表的载荷膨胀。
func truncateNotifyError(msg string) string {
const max = 500
if len(msg) <= max {
return msg
}
// 按字符边界回退,避免留下半个 UTF-8 字符让前端显示成乱码。
cut := max
for cut > 0 && !isUTF8Start(msg[cut]) {
cut--
}
return msg[:cut] + "…"
}
func isUTF8Start(b byte) bool { return b&0xC0 != 0x80 }
// placeholders 生成从 start 开始的 $n 占位串及对应参数,供 IN (...) 使用。
// 例如 start=3, ids=[7,8] → "$3,$4", [7,8]。
func placeholders(start int, ids []int64) (string, []any) {
ph := make([]string, 0, len(ids))
args := make([]any, 0, len(ids))
for i, id := range ids {
ph = append(ph, fmt.Sprintf("$%d", start+i))
args = append(args, id)
}
return strings.Join(ph, ","), args
}