Files
sub2api/backend/internal/repository/audit_log_repo.go
T
李建琦 6d655c9903
Release / update-version (push) Has been cancelled
Release / build-frontend (push) Has been cancelled
Release / release (push) Has been cancelled
Release / sync-version-file (push) Has been cancelled
CI / shell (push) Canceled after 0s
CI / test (push) Canceled after 0s
CI / frontend (push) Canceled after 0s
CI / golangci-lint (push) Canceled after 0s
Security Scan / backend-security (push) Canceled after 0s
Security Scan / frontend-security (push) Canceled after 0s
Sub2API v1.0 - AI API 网关(二开初始版本,基于上游 Wei-Shaw/sub2api)
2026-08-21 18:30:13 +08:00

369 lines
9.6 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 repository
import (
"context"
"database/sql"
"encoding/json"
"fmt"
"strings"
"time"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/lib/pq"
)
// auditLogRepository 审计日志仓储(raw SQLappend-only)。
// 刻意不实现单条删除:审计日志只允许追加、按保留期批量清理、以及带 2FA 的全量清空。
type auditLogRepository struct {
db *sql.DB
}
// NewAuditLogRepository 创建审计日志仓储。
func NewAuditLogRepository(db *sql.DB) service.AuditLogRepository {
return &auditLogRepository{db: db}
}
const auditLogInsertColumns = `created_at, actor_user_id, actor_email, actor_role, auth_method,
credential_masked, action, method, path, request_id, client_ip, user_agent,
request_body, status_code, latency_ms, extra`
func auditLogInsertValues(log *service.AuditLog) []any {
createdAt := log.CreatedAt
if createdAt.IsZero() {
createdAt = time.Now().UTC()
}
extraJSON := "{}"
if len(log.Extra) > 0 {
if encoded, err := json.Marshal(log.Extra); err == nil {
extraJSON = string(encoded)
}
}
return []any{
createdAt.UTC(),
nullInt64Ptr(log.ActorUserID),
truncateString(log.ActorEmail, 255),
truncateString(log.ActorRole, 32),
truncateString(log.AuthMethod, 32),
truncateString(log.CredentialMasked, 160),
truncateString(log.Action, 128),
truncateString(log.Method, 16),
truncateString(log.Path, 512),
truncateString(log.RequestID, 64),
truncateString(log.ClientIP, 64),
truncateString(log.UserAgent, 512),
log.RequestBody,
log.StatusCode,
log.LatencyMs,
extraJSON,
}
}
func (r *auditLogRepository) BatchInsert(ctx context.Context, logs []*service.AuditLog) (int64, error) {
if r == nil || r.db == nil {
return 0, fmt.Errorf("nil audit log repository")
}
if len(logs) == 0 {
return 0, nil
}
tx, err := r.db.BeginTx(ctx, nil)
if err != nil {
return 0, err
}
stmt, err := tx.PrepareContext(ctx, pq.CopyIn(
"audit_logs",
"created_at", "actor_user_id", "actor_email", "actor_role", "auth_method",
"credential_masked", "action", "method", "path", "request_id", "client_ip", "user_agent",
"request_body", "status_code", "latency_ms", "extra",
))
if err != nil {
_ = tx.Rollback()
return 0, err
}
var inserted int64
for _, log := range logs {
if log == nil {
continue
}
if _, err := stmt.ExecContext(ctx, auditLogInsertValues(log)...); err != nil {
_ = stmt.Close()
_ = tx.Rollback()
return inserted, err
}
inserted++
}
if _, err := stmt.ExecContext(ctx); err != nil {
_ = stmt.Close()
_ = tx.Rollback()
return inserted, err
}
if err := stmt.Close(); err != nil {
_ = tx.Rollback()
return inserted, err
}
if err := tx.Commit(); err != nil {
return inserted, err
}
return inserted, nil
}
func (r *auditLogRepository) Insert(ctx context.Context, log *service.AuditLog) error {
if r == nil || r.db == nil {
return fmt.Errorf("nil audit log repository")
}
if log == nil {
return fmt.Errorf("nil audit log")
}
query := `INSERT INTO audit_logs (` + auditLogInsertColumns + `)
VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,$15,$16)`
_, err := r.db.ExecContext(ctx, query, auditLogInsertValues(log)...)
return err
}
func buildAuditLogsWhere(filter *service.AuditLogFilter) (string, []any) {
clauses := make([]string, 0, 10)
args := make([]any, 0, 10)
clauses = append(clauses, "1=1")
if filter.StartTime != nil {
args = append(args, filter.StartTime.UTC())
clauses = append(clauses, "l.created_at >= $"+itoa(len(args)))
}
if filter.EndTime != nil {
args = append(args, filter.EndTime.UTC())
clauses = append(clauses, "l.created_at <= $"+itoa(len(args)))
}
if filter.ActorUserID != nil {
args = append(args, *filter.ActorUserID)
clauses = append(clauses, "l.actor_user_id = $"+itoa(len(args)))
}
if v := strings.TrimSpace(filter.ActorEmail); v != "" {
args = append(args, "%"+escapeLikePattern(v)+"%")
clauses = append(clauses, "l.actor_email ILIKE $"+itoa(len(args)))
}
if v := strings.TrimSpace(filter.AuthMethod); v != "" {
args = append(args, v)
clauses = append(clauses, "l.auth_method = $"+itoa(len(args)))
}
if v := strings.TrimSpace(filter.Action); v != "" {
args = append(args, "%"+escapeLikePattern(v)+"%")
clauses = append(clauses, "l.action ILIKE $"+itoa(len(args)))
}
if v := strings.TrimSpace(filter.Method); v != "" {
args = append(args, strings.ToUpper(v))
clauses = append(clauses, "l.method = $"+itoa(len(args)))
}
if v := strings.TrimSpace(filter.ClientIP); v != "" {
args = append(args, v)
clauses = append(clauses, "l.client_ip = $"+itoa(len(args)))
}
if filter.Success != nil {
if *filter.Success {
clauses = append(clauses, "l.status_code < 400")
} else {
clauses = append(clauses, "l.status_code >= 400")
}
}
if v := strings.TrimSpace(filter.Query); v != "" {
args = append(args, "%"+escapeLikePattern(v)+"%")
idx := itoa(len(args))
clauses = append(clauses, "(l.path ILIKE $"+idx+" OR l.action ILIKE $"+idx+" OR l.actor_email ILIKE $"+idx+")")
}
return "WHERE " + strings.Join(clauses, " AND "), args
}
const auditLogSelectColumns = `
l.id,
l.created_at,
l.actor_user_id,
COALESCE(l.actor_email, ''),
COALESCE(l.actor_role, ''),
COALESCE(l.auth_method, ''),
COALESCE(l.credential_masked, ''),
COALESCE(l.action, ''),
COALESCE(l.method, ''),
COALESCE(l.path, ''),
COALESCE(l.request_id, ''),
COALESCE(l.client_ip, ''),
COALESCE(l.user_agent, ''),
COALESCE(l.request_body, ''),
l.status_code,
l.latency_ms,
COALESCE(l.extra::text, '{}')`
func scanAuditLogRow(scan func(dest ...any) error) (*service.AuditLog, error) {
item := &service.AuditLog{}
var actorUserID sql.NullInt64
var extraRaw string
if err := scan(
&item.ID,
&item.CreatedAt,
&actorUserID,
&item.ActorEmail,
&item.ActorRole,
&item.AuthMethod,
&item.CredentialMasked,
&item.Action,
&item.Method,
&item.Path,
&item.RequestID,
&item.ClientIP,
&item.UserAgent,
&item.RequestBody,
&item.StatusCode,
&item.LatencyMs,
&extraRaw,
); err != nil {
return nil, err
}
if actorUserID.Valid {
v := actorUserID.Int64
item.ActorUserID = &v
}
extraRaw = strings.TrimSpace(extraRaw)
if extraRaw != "" && extraRaw != "null" && extraRaw != "{}" {
extra := make(map[string]any)
if err := json.Unmarshal([]byte(extraRaw), &extra); err == nil {
item.Extra = extra
}
}
return item, nil
}
func (r *auditLogRepository) List(ctx context.Context, filter *service.AuditLogFilter) (*service.AuditLogList, error) {
if r == nil || r.db == nil {
return nil, fmt.Errorf("nil audit log repository")
}
if filter == nil {
filter = &service.AuditLogFilter{}
}
page := filter.Page
if page <= 0 {
page = 1
}
pageSize := filter.PageSize
if pageSize <= 0 {
pageSize = 50
}
if pageSize > 200 {
pageSize = 200
}
where, args := buildAuditLogsWhere(filter)
countSQL := "SELECT COUNT(*) FROM audit_logs l " + where
var total int
if err := r.db.QueryRowContext(ctx, countSQL, args...).Scan(&total); err != nil {
return nil, err
}
offset := (page - 1) * pageSize
argsWithLimit := append(args, pageSize, offset)
query := "SELECT" + auditLogSelectColumns + "\nFROM audit_logs l\n" + where + `
ORDER BY l.created_at DESC, l.id DESC
LIMIT $` + itoa(len(args)+1) + ` OFFSET $` + itoa(len(args)+2)
rows, err := r.db.QueryContext(ctx, query, argsWithLimit...)
if err != nil {
return nil, err
}
defer func() { _ = rows.Close() }()
logs := make([]*service.AuditLog, 0, pageSize)
for rows.Next() {
item, err := scanAuditLogRow(rows.Scan)
if err != nil {
return nil, err
}
// 列表页不返回 body,降低载荷;详情接口返回完整记录。
item.RequestBody = ""
logs = append(logs, item)
}
if err := rows.Err(); err != nil {
return nil, err
}
return &service.AuditLogList{
Logs: logs,
Total: total,
Page: page,
PageSize: pageSize,
}, nil
}
func (r *auditLogRepository) GetByID(ctx context.Context, id int64) (*service.AuditLog, error) {
if r == nil || r.db == nil {
return nil, fmt.Errorf("nil audit log repository")
}
query := "SELECT" + auditLogSelectColumns + "\nFROM audit_logs l WHERE l.id = $1"
row := r.db.QueryRowContext(ctx, query, id)
item, err := scanAuditLogRow(row.Scan)
if err != nil {
if err == sql.ErrNoRows {
return nil, service.ErrAuditLogNotFound
}
return nil, err
}
return item, nil
}
func (r *auditLogRepository) Count(ctx context.Context) (int64, error) {
if r == nil || r.db == nil {
return 0, fmt.Errorf("nil audit log repository")
}
var total int64
if err := r.db.QueryRowContext(ctx, "SELECT COUNT(*) FROM audit_logs").Scan(&total); err != nil {
return 0, err
}
return total, nil
}
func (r *auditLogRepository) TruncateAll(ctx context.Context) error {
if r == nil || r.db == nil {
return fmt.Errorf("nil audit log repository")
}
_, err := r.db.ExecContext(ctx, "TRUNCATE TABLE audit_logs")
return err
}
func (r *auditLogRepository) DeleteBefore(ctx context.Context, cutoff time.Time, batchSize int) (int64, error) {
if r == nil || r.db == nil {
return 0, fmt.Errorf("nil audit log repository")
}
if batchSize <= 0 {
batchSize = 5000
}
res, err := r.db.ExecContext(ctx, `
WITH batch AS (
SELECT id FROM audit_logs WHERE created_at < $1 ORDER BY id LIMIT $2
)
DELETE FROM audit_logs WHERE id IN (SELECT id FROM batch)`, cutoff.UTC(), batchSize)
if err != nil {
return 0, err
}
return res.RowsAffected()
}
func nullInt64Ptr(v *int64) any {
if v == nil || *v <= 0 {
return nil
}
return *v
}
func truncateString(s string, max int) string {
s = strings.TrimSpace(s)
if len(s) <= max {
return s
}
// 按字节截断可能切断多字节字符,按 rune 处理。
runes := []rune(s)
for len(string(runes)) > max && len(runes) > 0 {
runes = runes[:len(runes)-1]
}
return string(runes)
}