This commit is contained in:
tao.chen
2026-08-27 12:41:28 +08:00
commit c8e4996c34
7 changed files with 2010 additions and 0 deletions
+594
View File
@@ -0,0 +1,594 @@
package main
import (
_ "embed"
"encoding/json"
"flag"
"fmt"
"log"
"maps"
"net/http"
"os"
"path/filepath"
"strings"
"sync"
"time"
"github.com/gin-gonic/gin"
"github.com/robfig/cron/v3"
)
// 嵌入式静态页面 (单文件部署, 把 HTML 编译进二进制, 无需运行时外部文件)
//
//go:embed static/index.html
var indexHTML []byte
// -----------------------------------------------------------------------------
// 配置 (可被命令行参数覆盖)
// -----------------------------------------------------------------------------
var (
uploadDir = flag.String("dir", "./data/uploads", "上传文件保存目录")
listen = flag.String("listen", ":8080", "HTTP 监听地址")
fileTTL = flag.Duration("ttl", 24*time.Hour, "文件过期 TTL")
scanEvery = flag.Duration("scan", 1*time.Hour, "清理扫描间隔")
auditDir = flag.String("audit-dir", "./data/audit", "审计日志目录 (每天一个 YYYY-MM-DD.log 文件, 空字符串禁用)")
quota = flag.Int64("quota", 10<<30, "总目录配额 (字节, 默认 10GB, 0=不限)")
)
// -----------------------------------------------------------------------------
// 路由处理器
// -----------------------------------------------------------------------------
const (
fieldName = "file"
)
type fileResponse struct {
Filename string `json:"filename"`
URL string `json:"url"`
Size int64 `json:"size"`
UploadedAt time.Time `json:"uploaded_at"`
ExpiresAt time.Time `json:"expires_at"`
}
// -----------------------------------------------------------------------------
// 审计日志 (JSON Lines, 一行一条事件, 落盘到 *auditDir/YYYY-MM-DD.log, 跨日自动轮转)
// -----------------------------------------------------------------------------
var (
auditMu sync.Mutex
auditFile *os.File
auditCurDay string // 当前打开文件对应的日期 (YYYY-MM-DD)
)
// initAudit 创建目录, 打开今天的日志文件
func initAudit() error {
if *auditDir == "" {
log.Println("[audit] disabled (audit dir is empty)")
return nil
}
if err := os.MkdirAll(*auditDir, 0o755); err != nil {
return fmt.Errorf("create audit dir: %w", err)
}
if err := openAuditFor(time.Now()); err != nil {
return err
}
log.Printf("[audit] writing to %s/%s.log", *auditDir, auditCurDay)
return nil
}
// openAuditFor 为指定时间打开对应日期的日志文件. 调用方需持有 auditMu.
func openAuditFor(t time.Time) error {
day := t.Format("2006-01-02")
if day == auditCurDay && auditFile != nil {
return nil
}
if auditFile != nil {
_ = auditFile.Close()
}
path := filepath.Join(*auditDir, day+".log")
f, err := os.OpenFile(path, os.O_CREATE|os.O_APPEND|os.O_WRONLY, 0o644)
if err != nil {
return fmt.Errorf("open %s: %w", path, err)
}
auditFile = f
auditCurDay = day
return nil
}
// audit 写一条审计事件. fields 里的键会合并到事件 JSON 中.
// 写入失败仅记录到 stdout, 不会影响主流程.
func audit(c *gin.Context, action string, fields map[string]any) {
if *auditDir == "" {
return
}
entry := map[string]any{
"ts": time.Now().Format(time.RFC3339),
"action": action,
"ip": c.ClientIP(),
"method": c.Request.Method,
"path": c.Request.URL.Path,
"ua": c.Request.UserAgent(),
"status": c.Writer.Status(),
"latency": time.Since(c.GetTime("t0")).String(),
}
maps.Copy(entry, fields)
b, err := json.Marshal(entry)
if err != nil {
log.Printf("[audit] marshal err: %v", err)
return
}
b = append(b, '\n')
auditMu.Lock()
defer auditMu.Unlock()
if err := openAuditFor(time.Now()); err != nil {
log.Printf("[audit] open err: %v", err)
return
}
if _, err := auditFile.Write(b); err != nil {
log.Printf("[audit] write err: %v", err)
}
}
// -----------------------------------------------------------------------------
// 配额 (总目录字节数限制, 默认 10GB, 0=不限)
// -----------------------------------------------------------------------------
var (
quotaMu sync.Mutex
quotaUsed int64 // 当前目录已用字节数
)
// initQuota 启动时扫描目录, 初始化配额计数
func initQuota() {
if *quota <= 0 {
log.Printf("[quota] disabled")
return
}
rescanQuota()
log.Printf("[quota] %s / %s", formatSize(quotaUsed), formatSize(*quota))
}
// rescanQuota 重新扫描目录并重置计数. 调用频率低 (启动 + 清理器每次跑完), 用来修正漂移.
func rescanQuota() {
quotaMu.Lock()
defer quotaMu.Unlock()
var total int64
_ = filepath.Walk(*uploadDir, func(_ string, info os.FileInfo, err error) error {
if err != nil || info.IsDir() {
return nil
}
total += info.Size()
return nil
})
quotaUsed = total
}
// quotaSnapshot 原子返回 used/cap
func quotaSnapshot() (used, cap_ int64) {
quotaMu.Lock()
defer quotaMu.Unlock()
return quotaUsed, *quota
}
// quotaRemaining 返回剩余可上传字节数. quota<=0 时返回 math.MaxInt64.
func quotaRemaining() int64 {
quotaMu.Lock()
defer quotaMu.Unlock()
if *quota <= 0 {
return 1 << 62
}
return *quota - quotaUsed
}
// quotaReserve 原子地预留 n 字节配额: 超限返回 false, 否则扣减并返回 true.
func quotaReserve(n int64) bool {
quotaMu.Lock()
defer quotaMu.Unlock()
if *quota > 0 && quotaUsed+n > *quota {
return false
}
quotaUsed += n
return true
}
func quotaSub(n int64) {
if n <= 0 {
return
}
quotaMu.Lock()
quotaUsed -= n
if quotaUsed < 0 { // 漂移防护
quotaUsed = 0
}
quotaMu.Unlock()
}
// formatSize 字节 → 人类可读
func formatSize(b int64) string {
const (
KB = 1 << 10
MB = 1 << 20
GB = 1 << 30
TB = 1 << 40
)
switch {
case b >= TB:
return fmt.Sprintf("%.2f TB", float64(b)/TB)
case b >= GB:
return fmt.Sprintf("%.2f GB", float64(b)/GB)
case b >= MB:
return fmt.Sprintf("%.2f MB", float64(b)/MB)
case b >= KB:
return fmt.Sprintf("%.2f KB", float64(b)/KB)
default:
return fmt.Sprintf("%d B", b)
}
}
func uploadHandler(c *gin.Context) {
f, err := c.FormFile(fieldName)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "missing form field 'file'"})
audit(c, "upload", gin.H{"file": "", "result": "missing_form_field"})
return
}
// 早 reject: 防止 NUL 字节 + 超长文件名打爆 Stat/Save
if strings.ContainsRune(f.Filename, 0) {
c.JSON(http.StatusBadRequest, gin.H{"error": "filename contains null byte"})
audit(c, "upload", gin.H{"file": f.Filename, "result": "null_byte"})
return
}
// Linux/macOS 路径分量 NAME_MAX = 255 字节
if len(f.Filename) > 200 {
c.JSON(http.StatusBadRequest, gin.H{"error": "filename too long (max 200 bytes)"})
audit(c, "upload", gin.H{"file": f.Filename, "result": "filename_too_long"})
return
}
if err := os.MkdirAll(*uploadDir, 0o755); err != nil {
log.Printf("[upload] mkdir failed: %v", err)
c.JSON(http.StatusInternalServerError, gin.H{"error": "internal error"})
audit(c, "upload", gin.H{"file": f.Filename, "result": "mkdir_failed"})
return
}
// 配额检查 (单文件能塞下 + 不会超总配额)
if *quota > 0 && f.Size > quotaRemaining() {
used, cap_ := quotaSnapshot()
c.JSON(http.StatusInsufficientStorage, gin.H{
"error": fmt.Sprintf("quota exceeded: %s used / %s cap, file size %s",
formatSize(used), formatSize(cap_), formatSize(f.Size)),
})
audit(c, "upload", gin.H{
"file": f.Filename,
"size": f.Size,
"result": "quota_exceeded",
})
return
}
name := uniqueFilename(*uploadDir, f.Filename)
dst := filepath.Join(*uploadDir, name)
if err := c.SaveUploadedFile(f, dst); err != nil {
log.Printf("[upload] save failed: %v", err)
c.JSON(http.StatusInternalServerError, gin.H{"error": "internal error"})
audit(c, "upload", gin.H{"file": f.Filename, "result": "save_failed"})
return
}
// 实际写入的字节数可能与 f.Size 略有差异 (SaveUploadedFile 内部处理), 用 stat 拿真实值
var written int64 = f.Size
if info, statErr := os.Stat(dst); statErr == nil {
written = info.Size()
}
// 原子预留配额, 兜住并发下 pre-save 检查放行的多个请求
if !quotaReserve(written) {
_ = os.Remove(dst)
used, cap_ := quotaSnapshot()
c.JSON(http.StatusInsufficientStorage, gin.H{
"error": fmt.Sprintf("quota exceeded: %s used / %s cap, file size %s",
formatSize(used), formatSize(cap_), formatSize(written)),
})
audit(c, "upload", gin.H{
"file": name,
"size": written,
"result": "quota_exceeded",
})
return
}
now := time.Now()
_ = os.Chtimes(dst, now, now)
c.JSON(http.StatusOK, fileResponse{
Filename: name,
URL: "/download/" + name,
Size: f.Size,
UploadedAt: now,
ExpiresAt: now.Add(*fileTTL),
})
audit(c, "upload", gin.H{
"file": name, // 落盘后的名字 (可能加了 (1) 后缀)
"orig": f.Filename,
"size": f.Size,
"result": "ok",
})
}
// uniqueFilename 把 original 清洗后, 在 dir 中找一个不存在的名字.
// 规则: 原名 → 原名 (1) → 原名 (2) → ...
// 例: report.pdf -> report.pdf
//
// report.pdf (1) -> report.pdf (1)
// .gitignore -> .gitignore
// .gitignore (1) -> .gitignore (1)
func uniqueFilename(dir, original string) string {
// 1. 剥离路径组件, 防穿越
base := filepath.Base(original)
if base == "" || base == "." || base == ".." {
base = "file"
}
// 2. 拆分 stem 和 ext, 处理 .gitignore 这类隐藏文件
ext := filepath.Ext(base)
stem := strings.TrimSuffix(base, ext)
if stem == "" {
// 整个名字都是扩展名 (如 .gitignore), 把整个当 stem
ext = ""
stem = base
}
// 3. 依次尝试 base, base (1), base (2), ...
// 用 O_CREATE|O_EXCL 原子占位, 解决 TOCTOU 竞态 (20 路并发同名上传必须各自拿到不同名字)
// 占位的 0 字节 placeholder 会在 SaveUploadedFile 时被覆盖
for i := range 10000 {
var candidate string
if i == 0 {
candidate = base
} else {
candidate = fmt.Sprintf("%s (%d)%s", stem, i, ext)
}
p := filepath.Join(dir, candidate)
f, err := os.OpenFile(p, os.O_CREATE|os.O_EXCL|os.O_WRONLY, 0o644)
if err == nil {
f.Close()
return candidate
}
if !os.IsExist(err) {
// 非 "已存在" 错误 (权限/磁盘满等), 兜底用时间戳
log.Printf("[uniqueFilename] open err=%v, fallback to ts", err)
return fmt.Sprintf("%s (%d)%s", stem, time.Now().UnixNano(), ext)
}
}
return fmt.Sprintf("%s (%d)%s", stem, time.Now().UnixNano(), ext)
}
func downloadHandler(c *gin.Context) {
raw := c.Param("filename")
name := filepath.Base(raw)
if name == "" || name == "." || name == ".." ||
strings.ContainsAny(name, "/\\\x00") {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid filename"})
audit(c, "download", gin.H{"file": raw, "result": "invalid_filename"})
return
}
p := filepath.Join(*uploadDir, name)
info, err := os.Stat(p)
if os.IsNotExist(err) {
c.JSON(http.StatusNotFound, gin.H{"error": "file not found or expired"})
audit(c, "download", gin.H{"file": name, "result": "not_found"})
return
}
if err != nil {
log.Printf("[download] stat failed: %v", err)
c.JSON(http.StatusInternalServerError, gin.H{"error": "internal error"})
audit(c, "download", gin.H{"file": name, "result": "stat_failed"})
return
}
if info.IsDir() {
c.JSON(http.StatusBadRequest, gin.H{"error": "not a file"})
audit(c, "download", gin.H{"file": name, "result": "is_dir"})
return
}
// 强制下载 + 防 XSS (浏览器不会按扩展名/Content-Type 渲染)
c.Header("Content-Disposition", `attachment; filename="`+name+`"`)
c.Header("X-Content-Type-Options", "nosniff")
// 下载即续期
now := time.Now()
_ = os.Chtimes(p, now, now)
c.File(p)
audit(c, "download", gin.H{
"file": name,
"size": info.Size(),
"result": "ok",
})
}
type fileInfo struct {
Filename string `json:"filename"`
URL string `json:"url"`
Size int64 `json:"size"`
ModTime time.Time `json:"mod_time"`
ExpiresAt time.Time `json:"expires_at"`
}
func listHandler(c *gin.Context) {
entries, err := os.ReadDir(*uploadDir)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
files := make([]fileInfo, 0, len(entries))
for _, e := range entries {
if e.IsDir() {
continue
}
info, err := e.Info()
if err != nil {
continue
}
files = append(files, fileInfo{
Filename: e.Name(),
URL: "/download/" + e.Name(),
Size: info.Size(),
ModTime: info.ModTime(),
ExpiresAt: info.ModTime().Add(*fileTTL),
})
}
c.JSON(http.StatusOK, gin.H{
"count": len(files),
"now": time.Now(),
"ttl": fileTTL.String(),
"quota": func() gin.H {
used, cap_ := quotaSnapshot()
return gin.H{
"used": used,
"cap": cap_,
"used_str": formatSize(used),
"cap_str": formatSize(cap_),
}
}(),
"files": files,
})
}
func deleteHandler(c *gin.Context) {
raw := c.Param("filename")
name := filepath.Base(raw)
if name == "" || name == "." || name == ".." ||
strings.ContainsAny(name, "/\\\x00") {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid filename"})
audit(c, "delete", gin.H{"file": raw, "result": "invalid_filename"})
return
}
p := filepath.Join(*uploadDir, name)
info, err := os.Stat(p)
if os.IsNotExist(err) {
c.JSON(http.StatusNotFound, gin.H{"error": "file not found or already deleted"})
audit(c, "delete", gin.H{"file": name, "result": "not_found"})
return
}
if err != nil {
log.Printf("[delete] stat failed: %v", err)
c.JSON(http.StatusInternalServerError, gin.H{"error": "internal error"})
audit(c, "delete", gin.H{"file": name, "result": "stat_failed"})
return
}
if info.IsDir() {
c.JSON(http.StatusBadRequest, gin.H{"error": "not a file"})
audit(c, "delete", gin.H{"file": name, "result": "is_dir"})
return
}
if err := os.Remove(p); err != nil {
log.Printf("[delete] remove failed: %v", err)
c.JSON(http.StatusInternalServerError, gin.H{"error": "internal error"})
audit(c, "delete", gin.H{"file": name, "result": "remove_failed"})
return
}
quotaSub(info.Size())
c.JSON(http.StatusOK, gin.H{"deleted": name})
audit(c, "delete", gin.H{"file": name, "size": info.Size(), "result": "ok"})
}
// -----------------------------------------------------------------------------
// TTL 清理器
// -----------------------------------------------------------------------------
var cleanerOnce sync.Once
func startCleaner() {
cleanerOnce.Do(func() {
if err := os.MkdirAll(*uploadDir, 0o755); err != nil {
log.Fatalf("create upload dir failed: %v", err)
}
c := cron.New()
_, err := c.AddFunc("@every "+scanEvery.String(), func() {
now := time.Now()
removed := 0
var removedBytes int64
_ = filepath.Walk(*uploadDir, func(path string, info os.FileInfo, err error) error {
if err != nil || info.IsDir() {
return nil
}
if now.Sub(info.ModTime()) > *fileTTL {
if rmErr := os.Remove(path); rmErr != nil {
log.Printf("[cleaner] remove failed: %s err=%v", path, rmErr)
} else {
log.Printf("[cleaner] removed: %s (age=%v, size=%s)", path, now.Sub(info.ModTime()), formatSize(info.Size()))
removed++
removedBytes += info.Size()
}
}
return nil
})
quotaSub(removedBytes)
// 漂移修正: 每小时一次 rescan 校正任何来源不明的字节数偏差
rescanQuota()
log.Printf("[cleaner] scan done, removed=%d (%s), quota used now=%s",
removed, formatSize(removedBytes), formatSize(quotaUsedSnapshot()))
})
if err != nil {
log.Fatalf("add cron job failed: %v", err)
}
c.Start()
log.Printf("[cleaner] started, schedule=@every %s dir=%s ttl=%s",
scanEvery, *uploadDir, *fileTTL)
})
}
// quotaUsedSnapshot 不持锁读, 仅用于 cleaner 自身的日志 (允许轻微漂移)
func quotaUsedSnapshot() int64 {
quotaMu.Lock()
defer quotaMu.Unlock()
return quotaUsed
}
// -----------------------------------------------------------------------------
// main
// -----------------------------------------------------------------------------
func main() {
flag.Parse()
startCleaner()
if err := initAudit(); err != nil {
log.Fatalf("init audit: %v", err)
}
initQuota()
gin.SetMode(gin.ReleaseMode)
r := gin.New()
r.Use(gin.Logger(), gin.Recovery(), auditT0())
// 嵌入式首页
r.GET("/", func(c *gin.Context) {
c.Data(http.StatusOK, "text/html; charset=utf-8", indexHTML)
})
// 业务路由
r.POST("/upload", uploadHandler)
r.GET("/download/:filename", downloadHandler)
r.DELETE("/files/:filename", deleteHandler)
r.GET("/files", listHandler)
log.Printf("server listening on %s, uploads -> %s, ttl=%s", *listen, *uploadDir, *fileTTL)
if err := r.Run(*listen); err != nil {
log.Fatal(err)
}
}
// auditT0 记录请求进入时间, 供 audit() 计算 latency
func auditT0() gin.HandlerFunc {
return func(c *gin.Context) {
c.Set("t0", time.Now())
c.Next()
}
}