FileCodeBox Go 重写版 v2.5.6(安全审计修复版)

Go 1.27.1 (Gin+GORM) + Vue 3 文件快传服务:

- 安全审计全部修复(docs/security-audit-2026-09-05.md):
  bcrypt 密码哈希与自动升级、presign 直传服务端大小/内容校验、
  全局请求体上限、依赖升级(govulncheck 0 命中)、janitor 后台清理、
  管理端审计动作落库、/admin CORS 收紧、通知内容白名单净化、
  会话默认 7 天、限流缓存故障降级、robots.txt 端点等
- 前端:取件链接复制修复(不再重复拼接提取码)、markdown 净化器加固
- Redis 支持库号(FCB_REDIS_DB / redis://…/db URL)
- 文档:docs/api/* 与 openapi.yaml 同步最新行为(robots.txt、
  提码 5 位起、chunk 32MiB 上限、admin 审计动作等)

验证:gofmt/go vet/go test 全绿;二进制端到端冒烟通过
This commit is contained in:
2026-09-05 04:22:41 +08:00
commit 9686fe887a
173 changed files with 32455 additions and 0 deletions
+176
View File
@@ -0,0 +1,176 @@
# 文件快传 Go 后端(server/
Go 1.27.1 + Gin + GORM 重写的文件快传后端。数据库**默认 SQLite**modernc.org/sqlite
纯 Go 驱动,CGO_ENABLED=0 交叉编译友好,零外部依赖),配置 `FCB_DB_DRIVER=postgres`
切换 Postgres(需求 ⑧);Redis 为**可选**增强(未配置 `FCB_REDIS_ADDR` 时自动降级为进程内存缓存)。
## 目录结构
```
server/
├── cmd/server/main.go # 入口:装配 配置→DB→缓存→设置→审计→中间件→路由
└── internal/
├── config/ # 配置:FCB_* 环境变量基线 + DB settings KV 运行时覆盖
│ └── schema.go # v2 新配置键 schema(键名/类型/默认值/边界)
├── model/ # GORM 模型 + 双方言 AutoMigrate
├── cache/ # 缓存统一接口:redis.go / memory.go 双实现
├── database/ # 双方言连接与迁移(sqlite 默认 / postgres 可选)
├── settings/ # settings KV 读写、密码哈希(sha256$salt$hash)、键 schema re-export
├── audit/ # 审计服务(Sink 抽象、失败重试队列)+ UA 设备解析
├── middleware/ # JWT 认证、IP 限流、审计中间件、CORS、IP 解析
├── storage/ # 存储引擎契约(interface.go 为 go-storage 的实现契约)
└── response/ # 统一响应 {"code":200,"msg":"","data":...}
```
## 环境变量
| 变量 | 必需 | 默认 | 说明 |
|---|---|---|---|
| `FCB_DB_DRIVER` | ❌ | `sqlite` | 数据库驱动:`sqlite` \| `postgres`(需求 ⑧) |
| `FCB_DB_DSN` | 视驱动 | `./data/filecodebox.db` | postgres:连接串(**必需**),如 `postgres://user:pass@host:5432/filecodebox?sslmode=disable`;sqlite:数据库文件路径(可空,自动创建 `data/` 目录) |
| `FCB_REDIS_ADDR` | ❌ | 空 | 为空时缓存降级为内存实现 |
| `FCB_LISTEN` | ❌ | `:8466` | 监听地址 |
| `FCB_STORAGE_ENGINE` | ❌ | `local` | `local` \| `s3` \| `webdav` |
| `FCB_TRUSTED_PROXIES` | ❌ | 空 | 可信代理 CIDR(逗号分隔),用于解析真实客户端 IP |
### 数据库模式(需求 ⑧)
- **SQLite(默认)**`FCB_DB_DRIVER=sqlite`(或缺省)。零 DSN 零依赖启动,数据库文件
默认 `./data/filecodebox.db``FCB_DB_DSN` 可覆盖路径;父目录自动创建)。
连接参数:`busy_timeout=10s` + `WAL` 日志模式 + `foreign_keys=1`(通过 DSN pragma 注入)。
- **Postgres(可选)**`FCB_DB_DRIVER=postgres` 且必须提供 `FCB_DB_DSN`,否则启动报错。
连接池沿用 v1 参数(32/8、1h 轮换)。
- 双方言共用 GORM 抽象:AutoMigrate、settings KV、全部业务查询方言无关;
唯一原生 DDLmigrates 台账表)在 `database.createMigratesTable` 内部分支处理。
## v2 新增配置键(internal/config/schema.go 为单一事实来源)
| 键 | 类型 | 默认 | 说明 |
|---|---|---|---|
| `background_url` | string | `""` | 需求 ①:背景图 URL/上传地址(空=主题默认;legacy `background` 键兜底) |
| `footer_text` | string | `""` | 需求 ②:页脚自定义内容(≤2000 字符) |
| `footer_beian` | string | `""` | 需求 ②:备案号(≤128 字符) |
| `notify_enabled` | int | `1` | 需求 ③:通知开关(1 开 / 0 关) |
| `notify_title` / `notify_content` | string | 见 defaults | 需求 ③:通知标题/内容(沿用参考实现语义) |
| `max_save_seconds` | int64 | `0` | 需求 ④:最长保存秒数上限(0=仅默认 7 天兜底,≤365 天) |
| `max_save_count` | int | `0` | 需求 ④:单次分享最大可取次数上限(0=不限制,≤100000) |
| `max_file_size` | int64 | `0` | 需求 ⑩:存储策略-单文件上限字节(0=回落 `uploadSize`,≤10GiB |
| `uploadSize` / `allowed_file_types` / `storageLimit` / `openUpload` | 既有 | - | 存储策略既有键(语义不变) |
| `uploadCount` / `uploadMinute` | 既有 | `10` / `1` | 上传频率限制(对齐参考 `ip_limit["upload"]` |
管理端与文档(t2/t4)以 `config.KVSchema()``settings.KVSchema()` re-export)为元数据源;
schema 同步测试保证 `KVSchema()``defaults()` 逐键一致。
## 数据模型
- `file_codes`:文件/文本分享(对齐参考 `apps/base/models.py::FileCodes`
- `upload_chunks`:分片上传记录
- `key_values`:运行时配置 KV`settings` / `sys_start` 键)
- `presign_upload_sessions`:预签名直传会话
- `storage_reservations`:上传容量预留
- `audit_logs`:上传/下载审计(时间/IP/UA/设备/动作/结果/字节数/耗时),需求 ③
- `migrates`:迁移台账表(双方言 DDL 分支,见 `database.createMigratesTable`
## 存储引擎契约(internal/storage/interface.go
`go-storage` 按此契约实现 local/s3/webdav 三引擎(**接口签名已冻结**):
```go
type Storage interface {
SaveFile(ctx, r io.Reader, savePath) (int64, error) // 流式保存
DeleteFile(ctx, savePath) error
Open(ctx, savePath, rng *Range) (*Download, error) // Range 下载
Stat(ctx, savePath) (*FileMeta, error)
SaveChunk(ctx, uploadID, chunkIndex, r io.Reader, savePath) (int64, error)
MergeChunks(ctx, uploadID, total, verifyHash, savePath) (int64, string, error)
CleanChunks(ctx, uploadID, savePath) error
FileExists(ctx, savePath) (bool, error)
PresignGetURL(ctx, savePath, expires) (string, error) // 不支持→ErrNotSupported
PresignPutURL(ctx, savePath, expires) (string, error)
HealthCheck(ctx) error
}
```
错误映射约定:`ErrNotFound`→404、`ErrInvalidPath`→400、`ErrUnavailable`→503、`ErrNotSupported`→501、`ErrRangeNotSatisfiable`→416。
引擎通过 `storage.RegisterEngine("local"|"s3"|"webdav", factory)` 注册,
`storage.NewEngine(ctx, name)` 构造。分片路径约定:`<父目录>/chunks/<uploadID>/<index>.part`
### 引擎构造选项(t3 API 层接入)
构造引擎前必须先注入选项(否则 local 退到系统临时目录、s3/webdav 因缺配置失败):
```go
storage.SetEngineOptions(storage.EngineOptions{
Local: storage.LocalOptions{Root: cfg.GetString("local_storage_path")},
S3: storage.S3Options{
AccessKeyID: cfg.GetString("s3_access_key_id"), SecretAccessKey: cfg.GetString("s3_secret_access_key"),
Bucket: cfg.GetString("s3_bucket_name"), Endpoint: cfg.GetString("s3_endpoint_url"),
Region: cfg.GetString("s3_region_name"), AddressingStyle: cfg.GetString("s3_addressing_style"),
},
WebDAV: storage.WebDAVOptions{
BaseURL: cfg.GetString("webdav_url"), Username: cfg.GetString("webdav_username"),
Password: cfg.GetString("webdav_password"), RootPath: cfg.GetString("webdav_root_path"),
},
})
st, err := storage.NewEngine(ctx, cfg.Engine()) // 构造后调 st.HealthCheck(ctx) 完成启动自检
```
引擎要点:local 原子写(临时文件+fsync+rename)、防穿越+符号链接逃逸校验;
s3 原生 multipart 流式合并、预签名直链、SDK 内置 5xx 重试;webdav 连接池复用、
Basic/Digest 自动协商、Range 透传、按需逐级 MKCOL(带缓存)、5xx/429 指数退避重试、
下载 io.Pipe 流式不落盘。三引擎均支持分片上传/合并(SHA256 校验)与 Range 下载。
## 审计埋点用法(API 层)
```go
r.Use(middleware.Audit(auditSvc, nil)) // nil=默认按路由前缀分类 upload/download
// handler 内填充业务字段并显式落库(推荐):
middleware.AuditSet(c, func(e *audit.Entry) { e.FileCode = code; e.FileName = name })
middleware.AuditRecordRequest(c, auditSvc, model.AuditResultSuccess, "")
// 或不显式落库:中间件按 HTTP 状态兜底(4xx/5xx→failed401/403/423/429/428→denied
```
下载响应字节数由中间件自动统计;上传字节数由 handler 填 `TransferredBytes`
## 限流语义(对齐参考实现)
- `error`(取件错误)/`login`(登录失败):**仅在失败时计数**,handler 调用 `limiter.Add(c, kind)`
- `upload`**成功上传才计数**(先 `Check` 放行,成功后 `Add`
- `metadata`:每次访问即计数,可用 `RequireRateLimit` 中间件
- 超限返回 HTTP 423;规则来自 settings KVerrorCount/errorMinute 等),可运行时调整
## 本地开发
```bash
# 默认 SQLite 模式:零依赖,数据库落 ./data/filecodebox.db
go run ./cmd/server # 启动于 :8466
curl localhost:8466/api/v1/health
# Postgres 模式(可选)
export FCB_DB_DRIVER=postgres
export FCB_DB_DSN='postgres://postgres:postgres@localhost:5432/filecodebox?sslmode=disable'
go run ./cmd/server
# 双方言单测:sqlite 始终执行;postgres 需真实实例(FCB_TEST_PG_DSN 指向测试库)
FCB_TEST_PG_DSN='postgres://postgres:postgres@localhost:5432/fcb_test?sslmode=disable' go test ./internal/database/ ./internal/settings/
go test ./... # 全量单测(未设置 FCB_TEST_PG_DSN 时自动跳过 PG 用例)
```
系统未初始化时除 `/setup``/api/v1/health` 外一律返回 428,初始化路由由 API 层任务接入。
## 开发注意
- **GOCACHE/GOMODCACHE 必须用 `export` 设置**(沙箱环境):本地与 CI 沙箱通常禁止写默认 Go 缓存
目录(`~/Library/Caches/go-build`)。注意一个隐蔽的坑——用**内联前缀变量**的方式
`GOCACHE=... GOMODCACHE=... go build ./... && go vet ./... && go test ./...`)时,
只有第一条命令继承这些变量,`&&` 后续命令(vet/test)会丢失前缀、落回默认缓存路径,
报错 `operation not permitted`(指向 `~/Library/Caches/...`)却极易误判为代码问题。
正确写法:
```bash
cd server
export GOCACHE=$(pwd)/../.gocache GOMODCACHE=$(pwd)/../.gomodcache GOSUMDB=off
go build ./... && go vet ./... && go test ./... # 后续命令也能继承 export 的变量
```
+301
View File
@@ -0,0 +1,301 @@
// 文件快递柜 Go 重写版服务入口:
// 装配 配置 → Postgres → 缓存 → 设置 → 审计 → 存储引擎 → 中间件 → API 路由。
package main
import (
"context"
"log"
"net/http"
"os"
"os/signal"
"strconv"
"strings"
"syscall"
"time"
"github.com/gin-gonic/gin"
"filecodebox/internal/api"
"filecodebox/internal/audit"
"filecodebox/internal/cache"
"filecodebox/internal/config"
"filecodebox/internal/database"
"filecodebox/internal/janitor"
"filecodebox/internal/middleware"
"filecodebox/internal/settings"
"filecodebox/internal/storage"
)
// APP_VERSION 版本号,对齐参考仓库 VERSION。
const APP_VERSION = "2.5.6"
func main() {
log.SetFlags(log.LstdFlags | log.Lshortfile)
ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
defer stop()
// 1. 配置(环境变量基线;引擎相关 env 种子在此注入,DB KV 仍可覆盖)
cfg, err := config.New()
if err != nil {
log.Fatalf("[boot] 配置加载失败: %v", err)
}
applyEnvEngineSeeds(cfg)
// 2. 数据库(需求 ⑧:默认 SQLite 零依赖;配置 FCB_DB_DRIVER=postgres+DSN 后走 Postgres
db, err := database.Open(ctx, database.Options{Driver: cfg.Env.DBDriver, DSN: cfg.Env.DBDSN})
if err != nil {
log.Fatalf("[boot] 数据库初始化失败: %v", err)
}
defer func() { _ = database.Close(db) }()
if err := database.Migrate(ctx, db); err != nil {
log.Fatalf("[boot] 数据库迁移失败: %v", err)
}
if cfg.Env.DBDriver == config.DBDriverPostgres {
log.Printf("[boot] Postgres 数据库已启用")
} else {
log.Printf("[boot] SQLite 数据库: %s", cfg.SQLitePath())
}
// 3. 缓存:Redis 可选,未配置则内存降级(需求 ⑧:默认即内存,Redis 为增强项)
cacheImpl, err := cache.New(ctx, cache.RedisOptions{Addr: cfg.Env.RedisAddr, DB: cfg.Env.RedisDB})
if err != nil {
log.Printf("[boot] Redis 连接失败,降级为内存缓存: %v", err)
cacheImpl, err = cache.NewMemory(), nil
if err != nil {
log.Fatalf("[boot] 内存缓存初始化失败: %v", err)
}
}
defer func() { _ = cacheImpl.Close() }()
if cfg.Env.RedisAddr != "" {
if cfg.Env.RedisDB != 0 {
log.Printf("[boot] Redis 缓存已启用: %s (db=%d)", cfg.Env.RedisAddr, cfg.Env.RedisDB)
} else {
log.Printf("[boot] Redis 缓存已启用: %s", cfg.Env.RedisAddr)
}
} else {
log.Println("[boot] 未配置 FCB_REDIS_ADDR,使用内存缓存")
}
// 4. 设置管理器(env 基线 + DB settings KV 运行时覆盖)
mgr, err := settings.NewManager(ctx, db, cfg)
if err != nil {
log.Fatalf("[boot] 设置加载失败: %v", err)
}
mgr.SystemStart(ctx)
autoInitIfNeeded(mgr)
// 5. 审计服务 + 重试循环(需求 ③)
auditSvc := audit.NewService(audit.NewDBSink(db))
auditStop := make(chan struct{})
defer close(auditStop)
auditSvc.StartRetryLoop(auditStop)
// 6. 限流器(规则来自 settings KV,可被管理端动态调整)
limiter := middleware.NewRateLimiter(cacheImpl, map[string]middleware.LimitRule{
middleware.LimitError: {Count: cfg.GetInt("errorCount"), Window: minutes(cfg.GetInt("errorMinute"))},
middleware.LimitUpload: {Count: cfg.GetInt("uploadCount"), Window: minutes(cfg.GetInt("uploadMinute"))},
middleware.LimitLogin: {Count: cfg.GetInt("loginCount"), Window: minutes(cfg.GetInt("loginMinute"))},
middleware.LimitMeta: {Count: cfg.GetInt("errorCount"), Window: minutes(cfg.GetInt("errorMinute"))},
})
// 7. 存储引擎:注入配置 → 构造 → 健康预检(需求 ④;v3 包装为可热切换 Manager)
storage.SetEngineOptions(buildEngineOptions(cfg))
bootStore, err := storage.NewEngine(ctx, cfg.Engine())
if err != nil {
log.Fatalf("[boot] 存储引擎 %s 初始化失败: %v", cfg.Engine(), err)
}
if err := bootStore.HealthCheck(ctx); err != nil {
// 预检失败仅告警不阻断启动(存储可能在运行中恢复)
log.Printf("[boot] 警告: 存储引擎 %s 健康检查未通过: %v", cfg.Engine(), err)
} else {
log.Printf("[boot] 存储引擎 %s 健康检查通过", cfg.Engine())
}
// v3Manager 包装——保存/读取委托当前引擎;管理端可热切换
//(构建闭包在每次切换前用最新 KV 刷新 EngineOptions,参数改动即时生效)
store := storage.NewManager(cfg.Engine(), bootStore, func(name string) (storage.Storage, error) {
storage.SetEngineOptions(buildEngineOptions(cfg))
return storage.NewEngine(ctx, name)
})
// 8. 路由装配
gin.SetMode(gin.ReleaseMode)
r := gin.New()
r.Use(gin.Logger(), gin.Recovery())
r.Use(middleware.ClientIP(middleware.ParseTrustedProxies(cfg.Env.TrustedProxies)))
// L6:管理端 CORS 收紧(允许站点对外域名 site_domain 跨域调用,其余拒绝)
r.Use(middleware.Cors(cfg.SiteDomain()))
// 未初始化守卫:除 /setup 与 /api/v1/health 外一律 428
r.Use(middleware.GuardNotInitialized(mgr.IsInitialized))
// M3:全局请求体大小上限(管理端 1MiB;上传类 = 单文件上限 + 表单开销)
r.Use(middleware.BodyLimit(bodyLimitFn(cfg)))
// 审计中间件:按 DefaultClassifier 路由模式自动分类 upload/download(需求 ③)
r.Use(middleware.Audit(auditSvc, nil))
api.Register(r, &api.Deps{
DB: db,
Cfg: cfg,
Mgr: mgr,
AuditSvc: auditSvc,
Limiter: limiter,
Store: store,
Version: APP_VERSION,
})
// 9. HTTP 服务
// M5:后台清理循环(过期预留/超时会话/直传残留对象),启动后 10 分钟首跑
janitor.Start(ctx, db, store, 10*time.Minute)
srv := &http.Server{
Addr: cfg.Env.Listen,
Handler: r,
ReadHeaderTimeout: 15 * time.Second,
}
go func() {
log.Printf("[boot] 文件快递柜 %s 启动于 %s(存储引擎 %s)", APP_VERSION, cfg.Env.Listen, cfg.Engine())
if err := srv.ListenAndServe(); err != nil && err != http.ErrServerClosed {
log.Fatalf("[boot] HTTP 服务异常退出: %v", err)
}
}()
<-ctx.Done()
log.Println("[boot] 收到退出信号,开始优雅关闭...")
shutdownCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
if err := srv.Shutdown(shutdownCtx); err != nil {
log.Printf("[boot] 优雅关闭超时: %v", err)
}
log.Println("[boot] 服务已退出")
}
// autoInitIfNeeded 安全审计 L1:系统未初始化且配置了 FCB_ADMIN_PASSWORD 时,
// 启动即自动完成管理员初始化,消除「部署到公网后被抢先访问 /setup 接管」的窗口。
// 密码不足 8 位时拒绝并保持 /setup 可用(不静默采用弱口令)。
func autoInitIfNeeded(mgr *settings.Manager) {
if mgr.IsInitialized() {
return
}
pwd := strings.TrimSpace(os.Getenv("FCB_ADMIN_PASSWORD"))
if pwd == "" {
return
}
if len(pwd) < 8 {
log.Printf("[boot] 警告: FCB_ADMIN_PASSWORD 少于 8 位,已忽略;/setup 初始化向导保持可用")
return
}
ctx := context.Background()
if err := mgr.UpdateKV(ctx, map[string]any{
"admin_token": settings.HashPassword(pwd),
"jwt_secret": settings.GenerateJWTSecret(),
}); err != nil {
log.Printf("[boot] 警告: FCB_ADMIN_PASSWORD 自动初始化写入失败: %v/setup 仍可用)", err)
return
}
if err := mgr.Reload(ctx); err != nil {
log.Printf("[boot] 警告: 自动初始化后配置重载失败: %v", err)
return
}
log.Printf("[boot] 已通过 FCB_ADMIN_PASSWORD 自动完成管理员初始化")
}
// envKVSeeds 容器编排常用的 FCB_* 环境变量 → settings KV 键映射。
// config 包只解析进程必需的 5 个变量;引擎相关配置通过本桥接注入,
// 使 docker-compose 无需进库即可完成引擎配置。优先级:defaults < 环境变量种子 < DB settings KV。
var envKVSeeds = []struct {
envKey string
kvKey string
}{
{"FCB_LOCAL_STORAGE_PATH", "local_storage_path"},
{"FCB_STORAGE_PATH", "storage_path"},
{"FCB_S3_ACCESS_KEY_ID", "s3_access_key_id"},
{"FCB_S3_SECRET_ACCESS_KEY", "s3_secret_access_key"},
{"FCB_AWS_SESSION_TOKEN", "aws_session_token"},
{"FCB_S3_BUCKET_NAME", "s3_bucket_name"},
{"FCB_S3_ENDPOINT_URL", "s3_endpoint_url"},
{"FCB_S3_REGION_NAME", "s3_region_name"},
{"FCB_S3_ADDRESSING_STYLE", "s3_addressing_style"},
{"FCB_WEBDAV_URL", "webdav_url"},
{"FCB_WEBDAV_USERNAME", "webdav_username"},
{"FCB_WEBDAV_PASSWORD", "webdav_password"},
{"FCB_WEBDAV_ROOT_PATH", "webdav_root_path"},
}
// applyEnvEngineSeeds 把已设置的环境变量作为 KV 种子写入配置(仅种子,不落库)。
func applyEnvEngineSeeds(cfg *config.Config) {
seeds := map[string]any{}
for _, m := range envKVSeeds {
if v := strings.TrimSpace(os.Getenv(m.envKey)); v != "" {
seeds[m.kvKey] = v
}
}
// L8:管理员会话有效期支持环境变量(秒;1~365 整天校验在 AdminSessionExpireSeconds
if v := strings.TrimSpace(os.Getenv("FCB_ADMIN_SESSION_EXPIRE")); v != "" {
if n, err := strconv.Atoi(v); err == nil {
seeds["adminSessionExpire"] = n
}
}
if len(seeds) > 0 {
cfg.ApplyKV(seeds)
log.Printf("[boot] 已从环境变量注入 %d 项配置种子", len(seeds))
}
}
// buildEngineOptions 从配置构造引擎选项(对齐 go-storage 的注入约定:
// 必须在 NewEngine 之前调用;字段来源为 config KV,键名对齐参考实现)。
func buildEngineOptions(cfg *config.Config) storage.EngineOptions {
return storage.EngineOptions{
Local: storage.LocalOptions{
Root: cfg.GetString("local_storage_path"),
},
S3: storage.S3Options{
AccessKeyID: cfg.GetString("s3_access_key_id"),
SecretAccessKey: cfg.GetString("s3_secret_access_key"),
SessionToken: cfg.GetString("aws_session_token"),
Bucket: cfg.GetString("s3_bucket_name"),
Endpoint: cfg.GetString("s3_endpoint_url"),
Region: cfg.GetString("s3_region_name"),
AddressingStyle: cfg.GetString("s3_addressing_style"),
},
WebDAV: storage.WebDAVOptions{
BaseURL: cfg.GetString("webdav_url"),
Username: cfg.GetString("webdav_username"),
Password: cfg.GetString("webdav_password"),
RootPath: cfg.GetString("webdav_root_path"),
},
}
}
// minutes 分钟数转 Duration。
func minutes(n int) time.Duration {
if n <= 0 {
n = 1
}
return time.Duration(n) * time.Minute
}
// bodyLimitFn 构造按路径分类的请求体上限函数(M3):
// - /admin/*1MiB(管理接口均为小 JSON/表单);
// - /share/text|metadata|select、/setup1MiB(文本内容本身限 222KB);
// - 其余(上传类):max_file_size0=回落 uploadSize,再 0=64MiB+ 2MiB 表单开销。
func bodyLimitFn(cfg *config.Config) func(c *gin.Context) int64 {
const adminLimit = int64(1) << 20
const textLimit = int64(1) << 20
const formOverhead = int64(2) << 20
const fallbackUpload = int64(64) << 20
return func(c *gin.Context) int64 {
p := c.Request.URL.Path
switch {
case p == "/setup" || strings.HasPrefix(p, "/admin/"):
return adminLimit
case p == "/share/text" || p == "/share/metadata" || p == "/share/select":
return textLimit
}
limit := cfg.MaxFileSize()
if limit <= 0 {
limit = cfg.UploadSize()
}
if limit <= 0 {
limit = fallbackUpload
}
// 分片/直传单请求体 ≤ 单片大小;share/file 的 multipart 有边界开销
return limit + formOverhead
}
}
+79
View File
@@ -0,0 +1,79 @@
module filecodebox
go 1.27.1
require (
github.com/aws/aws-sdk-go-v2 v1.45.1
github.com/aws/aws-sdk-go-v2/config v1.33.2
github.com/aws/aws-sdk-go-v2/credentials v1.20.2
github.com/aws/aws-sdk-go-v2/feature/s3/manager v1.23.2
github.com/aws/aws-sdk-go-v2/service/s3 v1.110.0
github.com/aws/smithy-go v1.28.1
github.com/gin-gonic/gin v1.11.0
github.com/glebarez/sqlite v1.11.0
github.com/golang-jwt/jwt/v5 v5.3.0
github.com/redis/go-redis/v9 v9.17.0
gorm.io/driver/postgres v1.6.2
gorm.io/gorm v1.31.2
)
require (
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.20 // indirect
github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.19.1 // indirect
github.com/aws/aws-sdk-go-v2/internal/configsources v1.5.1 // indirect
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.8.1 // indirect
github.com/aws/aws-sdk-go-v2/internal/v4a v1.5.1 // indirect
github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.19 // indirect
github.com/aws/aws-sdk-go-v2/service/internal/checksum v1.11.1 // indirect
github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.14.1 // indirect
github.com/aws/aws-sdk-go-v2/service/internal/s3shared v1.20.1 // indirect
github.com/aws/aws-sdk-go-v2/service/signin v1.8.0 // indirect
github.com/aws/aws-sdk-go-v2/service/sso v1.36.0 // indirect
github.com/aws/aws-sdk-go-v2/service/ssooidc v1.41.0 // indirect
github.com/aws/aws-sdk-go-v2/service/sts v1.48.0 // indirect
github.com/bytedance/sonic v1.14.0 // indirect
github.com/bytedance/sonic/loader v0.3.0 // indirect
github.com/cespare/xxhash/v2 v2.3.0 // indirect
github.com/cloudwego/base64x v0.1.6 // indirect
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f // indirect
github.com/dustin/go-humanize v1.0.1 // indirect
github.com/gabriel-vasile/mimetype v1.4.8 // indirect
github.com/gin-contrib/sse v1.1.0 // indirect
github.com/glebarez/go-sqlite v1.21.2 // indirect
github.com/go-playground/locales v0.14.1 // indirect
github.com/go-playground/universal-translator v0.18.1 // indirect
github.com/go-playground/validator/v10 v10.27.0 // indirect
github.com/goccy/go-json v0.10.2 // indirect
github.com/goccy/go-yaml v1.18.0 // indirect
github.com/google/uuid v1.6.0 // indirect
github.com/jackc/pgpassfile v1.0.0 // indirect
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect
github.com/jackc/pgx/v5 v5.10.0 // indirect
github.com/jackc/puddle/v2 v2.2.2 // indirect
github.com/jinzhu/inflection v1.0.0 // indirect
github.com/jinzhu/now v1.1.5 // indirect
github.com/json-iterator/go v1.1.12 // indirect
github.com/klauspost/cpuid/v2 v2.3.0 // indirect
github.com/leodido/go-urn v1.4.0 // indirect
github.com/mattn/go-isatty v0.0.20 // indirect
github.com/modern-go/concurrent v0.0.0-20180228061459-e0a39a4cb421 // indirect
github.com/modern-go/reflect2 v1.0.2 // indirect
github.com/ncruces/go-strftime v0.1.9 // indirect
github.com/pelletier/go-toml/v2 v2.2.4 // indirect
github.com/quic-go/qpack v0.6.0 // indirect
github.com/quic-go/quic-go v0.59.1 // indirect
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect
github.com/twitchyliquid64/golang-asm v0.15.1 // indirect
github.com/ugorji/go/codec v1.3.0 // indirect
golang.org/x/arch v0.20.0 // indirect
golang.org/x/crypto v0.56.0 // indirect
golang.org/x/net v0.58.0 // indirect
golang.org/x/sync v0.22.0 // indirect
golang.org/x/sys v0.47.0 // indirect
golang.org/x/text v0.41.0 // indirect
google.golang.org/protobuf v1.36.9 // indirect
modernc.org/libc v1.55.3 // indirect
modernc.org/mathutil v1.6.0 // indirect
modernc.org/memory v1.8.0 // indirect
modernc.org/sqlite v1.34.5 // indirect
)
+197
View File
@@ -0,0 +1,197 @@
github.com/aws/aws-sdk-go-v2 v1.45.1 h1:iIoG3NaLhV6UZpPXyPXlDj2I9oS8tV/nMcMnITCC6Ks=
github.com/aws/aws-sdk-go-v2 v1.45.1/go.mod h1:bttEH6JqnUL8LepvDVfdrds/fZ5bCIxzpe3abyUrhDU=
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.20 h1:GPRlPwz40I2B2VrBEASOA3Bi77NyeqejNLkifosX0rs=
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.20/go.mod h1:g7PNzKcsOKWb4fkSRBA7BZVAS6Y8IcxzN+nRohhQ1Q8=
github.com/aws/aws-sdk-go-v2/config v1.33.2 h1:Pj4+nF2kc4Z+1BJysVPnX9d5dMN7IYFXR4UJaWK2IpA=
github.com/aws/aws-sdk-go-v2/config v1.33.2/go.mod h1:Igw+HTwbR2tsTU/ydifAS9EHAFJ2s/FCgkwQWFnAdE4=
github.com/aws/aws-sdk-go-v2/credentials v1.20.2 h1:VQjZODPNfdikCX2ZZrltw4zNLkcwjyUFDUl2vT9yTwg=
github.com/aws/aws-sdk-go-v2/credentials v1.20.2/go.mod h1:OmeHCn28vZylsBvalLDf7t8fuJ2rHYQprJs+7WuxniI=
github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.19.1 h1:YIEBqcqRnpi4Pfv0YHImtgi6czGCwKHANC7SwmUAVD0=
github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.19.1/go.mod h1:imEf0oufgAo8KAkCHhrOdqGEC0YWx1PPBQH82shSxGw=
github.com/aws/aws-sdk-go-v2/feature/s3/manager v1.23.2 h1:yNAPkIRXwrXV3x4NMXi2oAveMy5WUaiBAY6X42K+vUs=
github.com/aws/aws-sdk-go-v2/feature/s3/manager v1.23.2/go.mod h1:+/m7PPNzeC3wq8n5kgw39kAj7pIE3fkAKHrgCyVnMO0=
github.com/aws/aws-sdk-go-v2/internal/configsources v1.5.1 h1:pc138gM1CW+XPc60rEwUlwwuwWFQK16CI1T7v1F9Oec=
github.com/aws/aws-sdk-go-v2/internal/configsources v1.5.1/go.mod h1:1+koxpPIbfBdfzP6vojm5/zTpTQ/micYwlxIiNB3TxI=
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.8.1 h1:K0JsbZQj+1h208Ro1zHeA4l7bMp0NvRffHQ91q8Ol1s=
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.8.1/go.mod h1:W3/vL6EtCIatICGy9ab29QhMuae+cOKPWcMxv02CO+Q=
github.com/aws/aws-sdk-go-v2/internal/v4a v1.5.1 h1:yhw5KD1phVyP9vijxOUzDfEtJx+bt+L63k+VfuiYFAA=
github.com/aws/aws-sdk-go-v2/internal/v4a v1.5.1/go.mod h1:ZW2e0d7DYlRxlS9hEiMXE47gTdX5KRN4byUiNbUpG+Q=
github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.19 h1:bAdDl/HkGCcGPoe25ToSHEw23VIxt6CT5fLcg111BKg=
github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.19/go.mod h1:KaUzbLxv4CeSxh6ZCl9B4m7CuFenS8kUEaDs+f/DQr4=
github.com/aws/aws-sdk-go-v2/service/internal/checksum v1.11.1 h1:s67hBfG5t9rn1NCvDuB4E3QIep3UFhHPtaIqFDjV3N8=
github.com/aws/aws-sdk-go-v2/service/internal/checksum v1.11.1/go.mod h1:FpvjBMXtSNMLPmDJsWwcY5cRnqJlpS2y1R6n4pvzs4k=
github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.14.1 h1:RmmWQPREQdk9U+PfqeHW3MqZaBaNK7TpV9W3RY+b+7g=
github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.14.1/go.mod h1:0A3W4F+68ZnNk5XcNL/e9HFMwnP8RlEicFfy6eOEDyw=
github.com/aws/aws-sdk-go-v2/service/internal/s3shared v1.20.1 h1:ZMbtPZZQRca+3+XYQne9PBvRiYpHZlNJJOZfE9WNfT0=
github.com/aws/aws-sdk-go-v2/service/internal/s3shared v1.20.1/go.mod h1:YAGWQdCYlVCoqrzvfv3RLxO6zKwti7gsAULOGWPLYv4=
github.com/aws/aws-sdk-go-v2/service/s3 v1.110.0 h1:He8vaTTqAAJrux/KdpjFXNWueLJZyKqE49QEXoqAu4I=
github.com/aws/aws-sdk-go-v2/service/s3 v1.110.0/go.mod h1:CUr46sCpGAg/rHaclRyhJX0LJAmH73uWSJPPSaMUrSk=
github.com/aws/aws-sdk-go-v2/service/signin v1.8.0 h1:bSvKIoLuRGFqGwASgeCQncCJDi9YKKBDEmCEZzOX1uU=
github.com/aws/aws-sdk-go-v2/service/signin v1.8.0/go.mod h1:9IqUlsJDbUPcg6cgx3WEzXdjrbWzLDQrak0aaSqlTcI=
github.com/aws/aws-sdk-go-v2/service/sso v1.36.0 h1:iivsh357VnfIc18IFWSuoyQEluf8frfWf4cL2Y0JUQw=
github.com/aws/aws-sdk-go-v2/service/sso v1.36.0/go.mod h1:tWuiVBUtPBr8/rgRiYS8Uf85sHcAN+G7XS3D3CEoUh8=
github.com/aws/aws-sdk-go-v2/service/ssooidc v1.41.0 h1:wVxM3QzSKIK8tSN6OGgezp9OK91lCLH2zhmRInN9rFM=
github.com/aws/aws-sdk-go-v2/service/ssooidc v1.41.0/go.mod h1:naFe83jSMuYkH+QjQPX8n1MLhBkeCFM5Lsnh5m5wz3c=
github.com/aws/aws-sdk-go-v2/service/sts v1.48.0 h1:RzZVCzYM19vhJCT5s6vO2wN8ie770Li/TmbAZ9B6N7E=
github.com/aws/aws-sdk-go-v2/service/sts v1.48.0/go.mod h1:mKo/CzaCz8qytGW70NG4vIIGAx1HXTlb5lHNkC5k3lk=
github.com/aws/smithy-go v1.28.1 h1:R/nXH00c8qcfCzQVELtRw+eLQWtzv+VAIEFJ1/xxXlQ=
github.com/aws/smithy-go v1.28.1/go.mod h1:YE2RhdIuDbA5E5bTdciG9KrW3+TiEONeUWCqxX9i1Fc=
github.com/bsm/ginkgo/v2 v2.12.0 h1:Ny8MWAHyOepLGlLKYmXG4IEkioBysk6GpaRTLC8zwWs=
github.com/bsm/ginkgo/v2 v2.12.0/go.mod h1:SwYbGRRDovPVboqFv0tPTcG1sN61LM1Z4ARdbAV9g4c=
github.com/bsm/gomega v1.27.10 h1:yeMWxP2pV2fG3FgAODIY8EiRE3dy0aeFYt4l7wh6yKA=
github.com/bsm/gomega v1.27.10/go.mod h1:JyEr/xRbxbtgWNi8tIEVPUYZ5Dzef52k01W3YH0H+O0=
github.com/bytedance/sonic v1.14.0 h1:/OfKt8HFw0kh2rj8N0F6C/qPGRESq0BbaNZgcNXXzQQ=
github.com/bytedance/sonic v1.14.0/go.mod h1:WoEbx8WTcFJfzCe0hbmyTGrfjt8PzNEBdxlNUO24NhA=
github.com/bytedance/sonic/loader v0.3.0 h1:dskwH8edlzNMctoruo8FPTJDF3vLtDT0sXZwvZJyqeA=
github.com/bytedance/sonic/loader v0.3.0/go.mod h1:N8A3vUdtUebEY2/VQC0MyhYeKUFosQU6FxH2JmUe6VI=
github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
github.com/cloudwego/base64x v0.1.6 h1:t11wG9AECkCDk5fMSoxmufanudBtJ+/HemLstXDLI2M=
github.com/cloudwego/base64x v0.1.6/go.mod h1:OFcloc187FXDaYHvrNIjxSe8ncn0OOM8gEHfghB2IPU=
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f h1:lO4WD4F/rVNCu3HqELle0jiPLLBs70cWOduZpkS1E78=
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f/go.mod h1:cuUVRXasLTGF7a8hSLbxyZXjz+1KgoB3wDUb6vlszIc=
github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY=
github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto=
github.com/gabriel-vasile/mimetype v1.4.8 h1:FfZ3gj38NjllZIeJAmMhr+qKL8Wu+nOoI3GqacKw1NM=
github.com/gabriel-vasile/mimetype v1.4.8/go.mod h1:ByKUIKGjh1ODkGM1asKUbQZOLGrPjydw3hYPU2YU9t8=
github.com/gin-contrib/sse v1.1.0 h1:n0w2GMuUpWDVp7qSpvze6fAu9iRxJY4Hmj6AmBOU05w=
github.com/gin-contrib/sse v1.1.0/go.mod h1:hxRZ5gVpWMT7Z0B0gSNYqqsSCNIJMjzvm6fqCz9vjwM=
github.com/gin-gonic/gin v1.11.0 h1:OW/6PLjyusp2PPXtyxKHU0RbX6I/l28FTdDlae5ueWk=
github.com/gin-gonic/gin v1.11.0/go.mod h1:+iq/FyxlGzII0KHiBGjuNn4UNENUlKbGlNmc+W50Dls=
github.com/glebarez/go-sqlite v1.21.2 h1:3a6LFC4sKahUunAmynQKLZceZCOzUthkRkEAl9gAXWo=
github.com/glebarez/go-sqlite v1.21.2/go.mod h1:sfxdZyhQjTM2Wry3gVYWaW072Ri1WMdWJi0k6+3382k=
github.com/glebarez/sqlite v1.11.0 h1:wSG0irqzP6VurnMEpFGer5Li19RpIRi2qvQz++w0GMw=
github.com/glebarez/sqlite v1.11.0/go.mod h1:h8/o8j5wiAsqSPoWELDUdJXhjAhsVliSn7bWZjOhrgQ=
github.com/go-playground/assert/v2 v2.2.0 h1:JvknZsQTYeFEAhQwI4qEt9cyV5ONwRHC+lYKSsYSR8s=
github.com/go-playground/assert/v2 v2.2.0/go.mod h1:VDjEfimB/XKnb+ZQfWdccd7VUvScMdVu0Titje2rxJ4=
github.com/go-playground/locales v0.14.1 h1:EWaQ/wswjilfKLTECiXz7Rh+3BjFhfDFKv/oXslEjJA=
github.com/go-playground/locales v0.14.1/go.mod h1:hxrqLVvrK65+Rwrd5Fc6F2O76J/NuW9t0sjnWqG1slY=
github.com/go-playground/universal-translator v0.18.1 h1:Bcnm0ZwsGyWbCzImXv+pAJnYK9S473LQFuzCbDbfSFY=
github.com/go-playground/universal-translator v0.18.1/go.mod h1:xekY+UJKNuX9WP91TpwSH2VMlDf28Uj24BCp08ZFTUY=
github.com/go-playground/validator/v10 v10.27.0 h1:w8+XrWVMhGkxOaaowyKH35gFydVHOvC0/uWoy2Fzwn4=
github.com/go-playground/validator/v10 v10.27.0/go.mod h1:I5QpIEbmr8On7W0TktmJAumgzX4CA1XNl4ZmDuVHKKo=
github.com/goccy/go-json v0.10.2 h1:CrxCmQqYDkv1z7lO7Wbh2HN93uovUHgrECaO5ZrCXAU=
github.com/goccy/go-json v0.10.2/go.mod h1:6MelG93GURQebXPDq3khkgXZkazVtN9CRI+MGFi0w8I=
github.com/goccy/go-yaml v1.18.0 h1:8W7wMFS12Pcas7KU+VVkaiCng+kG8QiFeFwzFb+rwuw=
github.com/goccy/go-yaml v1.18.0/go.mod h1:XBurs7gK8ATbW4ZPGKgcbrY1Br56PdM69F7LkFRi1kA=
github.com/golang-jwt/jwt/v5 v5.3.0 h1:pv4AsKCKKZuqlgs5sUmn4x8UlGa0kEVt/puTpKx9vvo=
github.com/golang-jwt/jwt/v5 v5.3.0/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE=
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg=
github.com/google/pprof v0.0.0-20240409012703-83162a5b38cd h1:gbpYu9NMq8jhDVbvlGkMFWCjLFlqqEZjEmObmhUy6Vo=
github.com/google/pprof v0.0.0-20240409012703-83162a5b38cd/go.mod h1:kf6iHlnVGwgKolg33glAes7Yg/8iWP8ukqeldJSO7jw=
github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsIM=
github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg=
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo=
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM=
github.com/jackc/pgx/v5 v5.10.0 h1:VhSvgU2jSli8o3AqIEOTJr7rZwAEUVo4E4XhR94Zfr0=
github.com/jackc/pgx/v5 v5.10.0/go.mod h1:mal1tBGAFfLHvZzaYh77YS/eC6IX9OWbRV1QIIM0Jn4=
github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo=
github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4=
github.com/jinzhu/inflection v1.0.0 h1:K317FqzuhWc8YvSVlFMCCUb36O/S9MCKRDI7QkRKD/E=
github.com/jinzhu/inflection v1.0.0/go.mod h1:h+uFLlag+Qp1Va5pdKtLDYj+kHp5pxUVkryuEj+Srlc=
github.com/jinzhu/now v1.1.5 h1:/o9tlHleP7gOFmsnYNz3RGnqzefHA47wQpKrrdTIwXQ=
github.com/jinzhu/now v1.1.5/go.mod h1:d3SSVoowX0Lcu0IBviAWJpolVfI5UJVZZ7cO71lE/z8=
github.com/json-iterator/go v1.1.12 h1:PV8peI4a0ysnczrg+LtxykD8LfKY9ML6u2jnxaEnrnM=
github.com/json-iterator/go v1.1.12/go.mod h1:e30LSqwooZae/UwlEbR2852Gd8hjQvJoHmT4TnhNGBo=
github.com/klauspost/cpuid/v2 v2.3.0 h1:S4CRMLnYUhGeDFDqkGriYKdfoFlDnMtqTiI/sFzhA9Y=
github.com/klauspost/cpuid/v2 v2.3.0/go.mod h1:hqwkgyIinND0mEev00jJYCxPNVRVXFQeu1XKlok6oO0=
github.com/leodido/go-urn v1.4.0 h1:WT9HwE9SGECu3lg4d/dIA+jxlljEa1/ffXKmRjqdmIQ=
github.com/leodido/go-urn v1.4.0/go.mod h1:bvxc+MVxLKB4z00jd1z+Dvzr47oO32F/QSNjSBOlFxI=
github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY=
github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y=
github.com/mattn/go-sqlite3 v1.14.22 h1:2gZY6PC6kBnID23Tichd1K+Z0oS6nE/XwU+Vz/5o4kU=
github.com/mattn/go-sqlite3 v1.14.22/go.mod h1:Uh1q+B4BYcTPb+yiD3kU8Ct7aC0hY9fxUwlHK0RXw+Y=
github.com/modern-go/concurrent v0.0.0-20180228061459-e0a39a4cb421 h1:ZqeYNhU3OHLH3mGKHDcjJRFFRrJa6eAM5H+CtDdOsPc=
github.com/modern-go/concurrent v0.0.0-20180228061459-e0a39a4cb421/go.mod h1:6dJC0mAP4ikYIbvyc7fijjWJddQyLn8Ig3JB5CqoB9Q=
github.com/modern-go/reflect2 v1.0.2 h1:xBagoLtFs94CBntxluKeaWgTMpvLxC4ur3nMaC9Gz0M=
github.com/modern-go/reflect2 v1.0.2/go.mod h1:yWuevngMOJpCy52FWWMvUC8ws7m/LJsjYzDa0/r8luk=
github.com/ncruces/go-strftime v0.1.9 h1:bY0MQC28UADQmHmaF5dgpLmImcShSi2kHU9XLdhx/f4=
github.com/ncruces/go-strftime v0.1.9/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls=
github.com/pelletier/go-toml/v2 v2.2.4 h1:mye9XuhQ6gvn5h28+VilKrrPoQVanw5PMw/TB0t5Ec4=
github.com/pelletier/go-toml/v2 v2.2.4/go.mod h1:2gIqNv+qfxSVS7cM2xJQKtLSTLUE9V8t9Stt+h56mCY=
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
github.com/quic-go/qpack v0.6.0 h1:g7W+BMYynC1LbYLSqRt8PBg5Tgwxn214ZZR34VIOjz8=
github.com/quic-go/qpack v0.6.0/go.mod h1:lUpLKChi8njB4ty2bFLX2x4gzDqXwUpaO1DP9qMDZII=
github.com/quic-go/quic-go v0.59.1 h1:0Gmua0HW1Tv7ANR7hUYwRyD0MG5OJfgvYSZasGZzBic=
github.com/quic-go/quic-go v0.59.1/go.mod h1:upnsH4Ju1YkqpLXC305eW3yDZ4NfnNbmQRCMWS58IKU=
github.com/redis/go-redis/v9 v9.17.0 h1:K6E+ZlYN95KSMmZeEQPbU/c++wfmEvfFB17yEAq/VhM=
github.com/redis/go-redis/v9 v9.17.0/go.mod h1:u410H11HMLoB+TP67dz8rL9s6QW2j76l0//kSOd3370=
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE=
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo=
github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME=
github.com/stretchr/objx v0.4.0/go.mod h1:YvHI0jy2hoMjB+UWwv71VJQ9isScKT/TqJzVSSt89Yw=
github.com/stretchr/objx v0.5.0/go.mod h1:Yh+to48EsGEfYuaHDzXPcE3xhTkx73EhmCGUpEOglKo=
github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI=
github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
github.com/stretchr/testify v1.7.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO+kdMU+MU=
github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4=
github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
github.com/twitchyliquid64/golang-asm v0.15.1 h1:SU5vSMR7hnwNxj24w34ZyCi/FmDZTkS4MhqMhdFk5YI=
github.com/twitchyliquid64/golang-asm v0.15.1/go.mod h1:a1lVb/DtPvCB8fslRZhAngC2+aY1QWCk3Cedj/Gdt08=
github.com/ugorji/go/codec v1.3.0 h1:Qd2W2sQawAfG8XSvzwhBeoGq71zXOC/Q1E9y/wUcsUA=
github.com/ugorji/go/codec v1.3.0/go.mod h1:pRBVtBSKl77K30Bv8R2P+cLSGaTtex6fsA2Wjqmfxj4=
go.uber.org/mock v0.5.2 h1:LbtPTcP8A5k9WPXj54PPPbjcI4Y6lhyOZXn+VS7wNko=
go.uber.org/mock v0.5.2/go.mod h1:wLlUxC2vVTPTaE3UD51E0BGOAElKrILxhVSDYQLld5o=
golang.org/x/arch v0.20.0 h1:dx1zTU0MAE98U+TQ8BLl7XsJbgze2WnNKF/8tGp/Q6c=
golang.org/x/arch v0.20.0/go.mod h1:bdwinDaKcfZUGpH09BB7ZmOfhalA8lQdzl62l8gGWsk=
golang.org/x/crypto v0.56.0 h1:GUh5Ii4J5jtcseSMiRqr1jXCNHoxjeV9Fmekc2oLy6Y=
golang.org/x/crypto v0.56.0/go.mod h1:OMW5y6CY9l38uPLmxU6l6pwcXp1obtLo3e6gT7gQR2I=
golang.org/x/mod v0.38.0 h1:MECBjubtXD7yj4HrhIUcywNaGeNVUdfVnxmPajOk4yk=
golang.org/x/mod v0.38.0/go.mod h1:V6Xz0pq8TQ3dGqVQ1FVHuelZpAL0uNhSkk9ogYP3c40=
golang.org/x/net v0.58.0 h1:ynWG7rqYi4ccpTEuPZ2QGWHktVEM9DMCj9yzDE0Q7To=
golang.org/x/net v0.58.0/go.mod h1:YwCddHnFlT7eLQqVprV19OnhLGtc5xOKgE0RyqgfWAU=
golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek=
golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8=
golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M=
golang.org/x/tools v0.48.0 h1:3+hClM1aLL5mjMKm5ovokw9epgRXPuu2tILgismM6RE=
golang.org/x/tools v0.48.0/go.mod h1:08xX0orndb/F7jJxGDicx061tyd5pcMto75YMAXr6lk=
google.golang.org/protobuf v1.36.9 h1:w2gp2mA27hUeUzj9Ex9FBjsBm40zfaDtEWow293U7Iw=
google.golang.org/protobuf v1.36.9/go.mod h1:fuxRtAxBytpl4zzqUh6/eyUujkJdNiuEkXntxiD/uRU=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
gorm.io/driver/postgres v1.6.2 h1:BvXQ/cNUg63q5TFNg672DmDcowZSFrNLkkA3Xe6GXq4=
gorm.io/driver/postgres v1.6.2/go.mod h1:0c4fQA44XhOklXDkgtuKqysHCycTa5i9e3EIpDGCwXk=
gorm.io/driver/sqlite v1.6.0 h1:WHRRrIiulaPiPFmDcod6prc4l2VGVWHz80KspNsxSfQ=
gorm.io/driver/sqlite v1.6.0/go.mod h1:AO9V1qIQddBESngQUKWL9yoH93HIeA1X6V633rBwyT8=
gorm.io/gorm v1.31.2 h1:3o8FXNo9v9S858gil+3LlZA1LkCOzgb4g5BL64FgaCo=
gorm.io/gorm v1.31.2/go.mod h1:XyQVbO2k6YkOis7C2437jSit3SsDK72s7n7rsSHd+Gs=
modernc.org/cc/v4 v4.21.4 h1:3Be/Rdo1fpr8GrQ7IVw9OHtplU4gWbb+wNgeoBMmGLQ=
modernc.org/cc/v4 v4.21.4/go.mod h1:HM7VJTZbUCR3rV8EYBi9wxnJ0ZBRiGE5OeGXNA0IsLQ=
modernc.org/ccgo/v4 v4.19.2 h1:lwQZgvboKD0jBwdaeVCTouxhxAyN6iawF3STraAal8Y=
modernc.org/ccgo/v4 v4.19.2/go.mod h1:ysS3mxiMV38XGRTTcgo0DQTeTmAO4oCmJl1nX9VFI3s=
modernc.org/fileutil v1.3.0 h1:gQ5SIzK3H9kdfai/5x41oQiKValumqNTDXMvKo62HvE=
modernc.org/fileutil v1.3.0/go.mod h1:XatxS8fZi3pS8/hKG2GH/ArUogfxjpEKs3Ku3aK4JyQ=
modernc.org/gc/v2 v2.4.1 h1:9cNzOqPyMJBvrUipmynX0ZohMhcxPtMccYgGOJdOiBw=
modernc.org/gc/v2 v2.4.1/go.mod h1:wzN5dK1AzVGoH6XOzc3YZ+ey/jPgYHLuVckd62P0GYU=
modernc.org/libc v1.55.3 h1:AzcW1mhlPNrRtjS5sS+eW2ISCgSOLLNyFzRh/V3Qj/U=
modernc.org/libc v1.55.3/go.mod h1:qFXepLhz+JjFThQ4kzwzOjA/y/artDeg+pcYnY+Q83w=
modernc.org/mathutil v1.6.0 h1:fRe9+AmYlaej+64JsEEhoWuAYBkOtQiMEU7n/XgfYi4=
modernc.org/mathutil v1.6.0/go.mod h1:Ui5Q9q1TR2gFm0AQRqQUaBWFLAhQpCwNcuhBOSedWPo=
modernc.org/memory v1.8.0 h1:IqGTL6eFMaDZZhEWwcREgeMXYwmW83LYW8cROZYkg+E=
modernc.org/memory v1.8.0/go.mod h1:XPZ936zp5OMKGWPqbD3JShgd/ZoQ7899TUuQqxY+peU=
modernc.org/opt v0.1.3 h1:3XOZf2yznlhC+ibLltsDGzABUGVx8J6pnFMS3E4dcq4=
modernc.org/opt v0.1.3/go.mod h1:WdSiB5evDcignE70guQKxYUl14mgWtbClRi5wmkkTX0=
modernc.org/sortutil v1.2.0 h1:jQiD3PfS2REGJNzNCMMaLSp/wdMNieTbKX920Cqdgqc=
modernc.org/sortutil v1.2.0/go.mod h1:TKU2s7kJMf1AE84OoiGppNHJwvB753OYfNl2WRb++Ss=
modernc.org/sqlite v1.34.5 h1:Bb6SR13/fjp15jt70CL4f18JIN7p7dnMExd+UFnF15g=
modernc.org/sqlite v1.34.5/go.mod h1:YLuNmX9NKs8wRNK2ko1LW1NGYcc9FkBO69JOt1AR9JE=
modernc.org/strutil v1.2.0 h1:agBi9dp1I+eOnxXeiZawM8F4LawKv4NzGWSaLfyeNZA=
modernc.org/strutil v1.2.0/go.mod h1:/mdcBmfOibveCTBxUl5B5l6W+TTH1FXPLHZE6bTosX0=
modernc.org/token v1.1.0 h1:Xl7Ap9dKaEs5kLoOQeQmPWevfnk/DM5qcLcYlA8ys6Y=
modernc.org/token v1.1.0/go.mod h1:UGzOrNV1mAFSEB63lOFHIpNRUVMvYTc6yu1SMY/XTDM=
File diff suppressed because it is too large Load Diff
+667
View File
@@ -0,0 +1,667 @@
package api
import (
"bytes"
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"fmt"
"io"
"mime/multipart"
"net/http"
"strconv"
"time"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"filecodebox/internal/middleware"
"filecodebox/internal/model"
"filecodebox/internal/response"
"filecodebox/internal/storage"
)
// chunkExpireTTL 分片会话保留时长(M5:预留窗口由 24h 缩短为 2h;
// 会话本身保留 24h 支持断点续传,见 janitor 的清理周期)。
const chunkExpireTTL = 2 * time.Hour
// maxChunkSizeBytes 单分片大小上限 32MBM3:限制 io.ReadAll 内存占用)。
const maxChunkSizeBytes = 32 * 1024 * 1024
// ============ POST /chunk/upload/init 初始化分片会话 ============
// requireChunkEnabled L4enableChunk 开关后端强制(此前仅前端隐藏入口,
// 开关关闭后 /chunk/* 接口仍可直接调用)。
func (d *Deps) requireChunkEnabled(c *gin.Context) bool {
if d.Cfg.EnableChunk() {
return true
}
auditRecordFailed(c, d.AuditSvc, "分片上传未启用")
response.Fail(c, http.StatusForbidden, "分片上传未启用")
return false
}
// chunkInitRequest init 请求体(JSON 或表单)。
type chunkInitRequest struct {
FileName string `json:"file_name" form:"file_name"`
ChunkSize int64 `json:"chunk_size" form:"chunk_size"`
FileSize int64 `json:"file_size" form:"file_size"`
FileHash string `json:"file_hash" form:"file_hash"`
}
// chunkInit 创建分片上传会话(对齐参考 init_chunk_upload):
// 支持断点续传(相同 hash/大小/文件名的未完成会话直接续传)。
func (d *Deps) chunkInit(c *gin.Context) {
if !d.requireChunkEnabled(c) {
return
}
if !d.requireShareLogin(c) {
return
}
if !requireUploadLimit(c, d.Limiter) {
return
}
var req chunkInitRequest
if err := bindJSONOrForm(c, &req); err != nil {
respondError(c, err)
return
}
safeName := storage.SanitizeFileName(req.FileName)
if safeName == "" {
auditRecordFailed(c, d.AuditSvc, "文件名非法")
response.Fail(c, http.StatusBadRequest, "文件名非法")
return
}
// 文件类型白名单(无内容可校验,仅名称)
if err := validateFileMagic(d.Cfg, safeName, "", nil); err != nil {
auditUploadEntry(c, "", safeName, req.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "文件类型被拒绝")
respondError(c, err)
return
}
chunkSize := req.ChunkSize
if chunkSize <= 0 {
chunkSize = 5 * 1024 * 1024 // 默认 5MB(对齐参考 InitChunkUploadModel
}
// M3:单片全部读入内存后再落存储,必须限制单片大小(客户端声明的
// chunk_size 上界受策略约束,但策略允许至 10GiB → 显式封顶 32MB)。
if chunkSize > maxChunkSizeBytes {
auditRecordFailed(c, d.AuditSvc, "chunk_size 超过上限")
response.Fail(c, http.StatusBadRequest, fmt.Sprintf("chunk_size 过大,最大为 %d MB", maxChunkSizeBytes>>20))
return
}
if req.FileSize <= 0 {
auditRecordFailed(c, d.AuditSvc, "file_size 非法")
response.Fail(c, http.StatusBadRequest, "file_size 必须大于 0")
return
}
// 服务端按分片数上限校验总大小(防分片声明绕过)
totalChunks := (req.FileSize + chunkSize - 1) / chunkSize
maxPossible := totalChunks * chunkSize
// v2 需求 ④⑩:动态策略校验(max_file_size0=回落 uploadSize
if err := d.CurrentUploadPolicy().CheckSize(maxPossible); err != nil {
auditUploadEntry(c, "", safeName, req.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "文件大小超过限制")
respondError(c, err)
return
}
ctx := c.Request.Context()
// 断点续传:查找相同 hash+大小+文件名的未完成会话(chunk_index=-1 为会话头)
var existing model.UploadChunk
err := d.DB.WithContext(ctx).
Where("chunk_hash = ? AND chunk_index = -1 AND file_size = ? AND file_name = ?",
req.FileHash, req.FileSize, safeName).
First(&existing).Error
if err == nil {
if existing.SavePath == "" {
// 脏会话:清理后按新建处理
_ = d.DB.WithContext(ctx).
Where("upload_id = ?", existing.UploadID).
Delete(&model.UploadChunk{}).Error
releaseStorage(ctx, d.DB, "chunk:"+existing.UploadID)
} else {
if err := reserveStorage(ctx, d.DB, d.Cfg, "chunk:"+existing.UploadID, existing.FileSize, chunkExpireTTL); err != nil {
respondError(c, err)
return
}
uploaded := d.uploadedChunkIndexes(ctx, existing.UploadID)
auditUploadEntry(c, existing.UploadID, safeName, req.FileSize, 0)
auditRecordSuccess(c, d.AuditSvc)
response.OK(c, gin.H{
"existed": false,
"upload_id": existing.UploadID,
"chunk_size": existing.ChunkSize,
"total_chunks": existing.TotalChunks,
"uploaded_chunks": uploaded,
})
return
}
} else if !errors.Is(err, gorm.ErrRecordNotFound) {
respondError(c, errInternal("查询上传会话失败: "+err.Error()))
return
}
// 新建会话
uploadID := uuidHex()
resToken := "chunk:" + uploadID
if err := reserveStorage(ctx, d.DB, d.Cfg, resToken, req.FileSize, chunkExpireTTL); err != nil {
auditUploadEntry(c, "", safeName, req.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "容量预留失败")
respondError(c, err)
return
}
// M5:init 即计入上传限流(此前仅 complete 成功时计数,
// 恶意客户端可无限创建会话占用容量预留)
d.Limiter.Add(c, middleware.LimitUpload)
_, _, _, _, savePath := buildSavePath(d.Cfg, safeName, uploadID)
session := model.UploadChunk{
UploadID: uploadID,
ChunkIndex: -1,
TotalChunks: int(totalChunks),
FileSize: req.FileSize,
ChunkSize: int(chunkSize),
ChunkHash: req.FileHash,
FileName: safeName,
SavePath: savePath,
Engine: d.Store.CurrentName(), // v3:会话归属引擎(分片/合并全程走同一引擎)
}
if err := d.DB.WithContext(ctx).Create(&session).Error; err != nil {
releaseStorage(ctx, d.DB, resToken)
auditUploadEntry(c, "", safeName, req.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "会话创建失败")
respondError(c, errInternal("创建上传会话失败: "+err.Error()))
return
}
auditUploadEntry(c, uploadID, safeName, req.FileSize, 0)
auditRecordSuccess(c, d.AuditSvc)
response.OK(c, gin.H{
"existed": false,
"upload_id": uploadID,
"chunk_size": chunkSize,
"total_chunks": totalChunks,
"uploaded_chunks": []int{},
})
}
// uploadedChunkIndexes 查询会话中已完成分片的索引列表。
func (d *Deps) uploadedChunkIndexes(ctx context.Context, uploadID string) []int {
var rows []model.UploadChunk
if err := d.DB.WithContext(ctx).
Where("upload_id = ? AND completed = ?", uploadID, true).
Order("chunk_index ASC").Find(&rows).Error; err != nil {
return []int{}
}
out := make([]int, 0, len(rows))
for _, r := range rows {
out = append(out, r.ChunkIndex)
}
return out
}
// ============ POST /chunk/upload/{uploadID}/{index}(及扁平兼容)============
// chunkUploadFlat 扁平模式:POST /chunk/uploadupload_id/chunk_index 走表单或 query。
// 多文件字段(chunk/chunks)时按 base_chunk_index 顺序批量接收。
func (d *Deps) chunkUploadFlat(c *gin.Context) {
c.Params = append(c.Params, gin.Param{Key: "uploadID", Value: resolveUploadID(c)})
c.Params = append(c.Params, gin.Param{Key: "chunkIndex", Value: resolveChunkIndex(c)})
d.chunkUpload(c)
}
// resolveUploadID 解析 upload_id:路径参数 → multipart 表单 → query。
func resolveUploadID(c *gin.Context) string {
if v := c.Param("uploadID"); v != "" {
return v
}
if v := c.PostForm("upload_id"); v != "" {
return v
}
return c.Query("upload_id")
}
// resolveChunkIndex 解析 chunk_index:路径参数 → multipart 表单 → query。
func resolveChunkIndex(c *gin.Context) string {
if v := c.Param("chunkIndex"); v != "" {
return v
}
if v := c.PostForm("chunk_index"); v != "" {
return v
}
return c.Query("chunk_index")
}
// chunkUpload 上传单个(或批量)分片(对齐参考 upload_chunk)。
// multipart 文件字段:chunk(主)或 file(回退);批量用 chunk[]/chunks 数组 + chunk_index 为起始索引。
func (d *Deps) chunkUpload(c *gin.Context) {
if !d.requireChunkEnabled(c) {
return
}
if !d.requireShareLogin(c) {
return
}
uploadID := resolveUploadID(c)
ctx := c.Request.Context()
var session model.UploadChunk
if err := d.DB.WithContext(ctx).
Where("upload_id = ? AND chunk_index = -1", uploadID).
First(&session).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
auditUploadEntry(c, uploadID, "", 0, 0)
auditRecordFailed(c, d.AuditSvc, "上传会话不存在")
response.Fail(c, http.StatusNotFound, "上传会话不存在")
return
}
respondError(c, errInternal("查询上传会话失败: "+err.Error()))
return
}
if err := reserveStorage(ctx, d.DB, d.Cfg, "chunk:"+uploadID, session.FileSize, chunkExpireTTL); err != nil {
respondError(c, err)
return
}
// 收集分片文件:chunk(单)→ file(回退)→ chunk[]/chunks(批量)
form, err := c.MultipartForm()
if err != nil {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "multipart 解析失败")
response.Fail(c, http.StatusBadRequest, "multipart 表单解析失败")
return
}
files := form.File["chunk"]
single := len(files) == 0
if single {
files = form.File["file"]
}
if len(files) == 0 {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "缺少 chunk 分片字段")
response.Fail(c, http.StatusBadRequest, "缺少分片文件字段 chunk")
return
}
baseIndex, err := strconv.Atoi(resolveChunkIndex(c))
if err != nil {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "无效的分片索引")
response.Fail(c, http.StatusBadRequest, "无效的分片索引")
return
}
results := make([]gin.H, 0, len(files))
for i, fh := range files {
// 单分片模式严格使用请求索引;批量模式从 base 递增
idx := baseIndex
if !single && len(files) > 1 {
idx = baseIndex + i
}
res, status, msg := d.saveOneChunk(c, ctx, &session, idx, fh)
if status != 0 {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, msg)
response.Fail(c, status, msg)
return
}
results = append(results, res)
}
// 审计:传输字节数为本次请求分片总和
var transferred int64
for _, fh := range files {
transferred += fh.Size
}
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, transferred)
auditRecordSuccess(c, d.AuditSvc)
if len(results) == 1 {
response.OK(c, results[0])
return
}
response.OK(c, gin.H{"chunks": results})
}
// saveOneChunk 保存一个分片:查重→读数据→校验→存储→记录。
// 返回 (响应体, HTTP错误状态码, 错误信息);成功时状态码为 0。
func (d *Deps) saveOneChunk(c *gin.Context, ctx context.Context, session *model.UploadChunk, idx int, fh *multipart.FileHeader) (gin.H, int, string) {
if idx < 0 || idx >= session.TotalChunks {
return nil, http.StatusBadRequest, "无效的分片索引"
}
// 已上传分片:断点续传直接跳过
var existing model.UploadChunk
err := d.DB.WithContext(ctx).
Where("upload_id = ? AND chunk_index = ? AND completed = ?", session.UploadID, idx, true).
First(&existing).Error
if err == nil {
return gin.H{"chunk_hash": existing.ChunkHash, "skipped": true, "chunk_index": idx}, 0, ""
}
if !errors.Is(err, gorm.ErrRecordNotFound) {
return nil, http.StatusInternalServerError, "查询分片记录失败"
}
f, err := fh.Open()
if err != nil {
return nil, http.StatusBadRequest, "分片数据读取失败"
}
defer func() { _ = f.Close() }()
data, err := io.ReadAll(io.LimitReader(f, int64(session.ChunkSize)+1))
if err != nil {
return nil, http.StatusBadRequest, "分片数据读取失败"
}
// 校验分片大小不超过声明值
if int64(len(data)) > int64(session.ChunkSize) {
return nil, http.StatusBadRequest,
"分片大小超过声明值: 最大 " + strconv.Itoa(session.ChunkSize) + ", 实际 " + strconv.Itoa(len(data))
}
// 累计大小校验(已传分片数×chunk_size + 当前分片;动态策略上限)
var uploadedCount int64
_ = d.DB.WithContext(ctx).Model(&model.UploadChunk{}).
Where("upload_id = ? AND completed = ?", session.UploadID, true).
Count(&uploadedCount).Error
if err := d.CurrentUploadPolicy().CheckSize(uploadedCount*int64(session.ChunkSize) + int64(len(data))); err != nil {
return nil, http.StatusForbidden, err.Error()
}
// 首分片做 magic bytes 防伪造
if idx == 0 {
head := data
if len(head) > 64 {
head = head[:64]
}
if err := validateFileMagic(d.Cfg, session.FileName, "", head); err != nil {
return nil, http.StatusForbidden, "文件内容校验失败:" + err.Error()
}
}
sum := sha256.Sum256(data)
chunkHash := hex.EncodeToString(sum[:])
if _, err := d.Store.SaveChunk(ctx, session.UploadID, idx, bytes.NewReader(data), session.SavePath); err != nil {
return nil, http.StatusInternalServerError, "分片保存失败: " + err.Error()
}
// 保存成功后再记录(对齐参考:先存储后落库)。
// 注意:不能用结构体 Where 条件(GORM 会忽略零值字段,chunk_index=0 会被
// 丢弃从而误匹配 -1 会话行),必须用字符串条件 + 完整目标结构体。
rec := model.UploadChunk{
UploadID: session.UploadID,
ChunkIndex: idx,
ChunkHash: chunkHash,
Completed: true,
FileSize: session.FileSize,
TotalChunks: session.TotalChunks,
ChunkSize: session.ChunkSize,
FileName: session.FileName,
SavePath: session.SavePath,
Engine: session.Engine, // v3:继承会话引擎
}
if err := d.DB.WithContext(ctx).
Where("upload_id = ? AND chunk_index = ?", session.UploadID, idx).
FirstOrCreate(&rec).Error; err != nil {
return nil, http.StatusInternalServerError, "分片记录写入失败"
}
return gin.H{"chunk_hash": chunkHash, "chunk_index": idx}, 0, ""
}
// ============ GET /chunk/upload/status/{uploadID} ============
// chunkStatus 查询上传进度(对齐参考 get_upload_status)。
func (d *Deps) chunkStatus(c *gin.Context) {
if !d.requireShareLogin(c) {
return
}
uploadID := c.Param("uploadID")
if uploadID == "" {
uploadID = c.Query("upload_id")
}
ctx := c.Request.Context()
var session model.UploadChunk
if err := d.DB.WithContext(ctx).
Where("upload_id = ? AND chunk_index = -1", uploadID).
First(&session).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.Fail(c, http.StatusNotFound, "上传会话不存在")
return
}
respondError(c, errInternal("查询上传会话失败: "+err.Error()))
return
}
uploaded := d.uploadedChunkIndexes(ctx, uploadID)
var progress float64
if session.TotalChunks > 0 {
progress = float64(len(uploaded)) / float64(session.TotalChunks) * 100
}
response.OK(c, gin.H{
"upload_id": uploadID,
"file_name": session.FileName,
"file_size": session.FileSize,
"chunk_size": session.ChunkSize,
"total_chunks": session.TotalChunks,
"uploaded_chunks": uploaded,
"progress": progress,
})
}
// ============ POST /chunk/upload/complete/{uploadID} ============
// chunkCompleteRequest complete 请求体。
type chunkCompleteRequest struct {
ExpireValue int `json:"expire_value" form:"expire_value"`
ExpireStyle string `json:"expire_style" form:"expire_style"`
Code string `json:"code" form:"code"` // v3.1:自定义提取码(4-8 位字母数字,空=随机)
}
// chunkComplete 合并分片并创建分享(对齐参考 complete_upload)。
func (d *Deps) chunkComplete(c *gin.Context) {
if !d.requireChunkEnabled(c) {
return
}
if !d.requireShareLogin(c) {
return
}
if !requireUploadLimit(c, d.Limiter) {
return
}
uploadID := c.Param("uploadID")
if uploadID == "" {
uploadID = resolveUploadID(c)
}
ctx := c.Request.Context()
var session model.UploadChunk
if err := d.DB.WithContext(ctx).
Where("upload_id = ? AND chunk_index = -1", uploadID).
First(&session).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
auditUploadEntry(c, uploadID, "", 0, 0)
auditRecordFailed(c, d.AuditSvc, "上传会话不存在")
response.Fail(c, http.StatusNotFound, "上传会话不存在")
return
}
respondError(c, errInternal("查询上传会话失败: "+err.Error()))
return
}
var req chunkCompleteRequest
if err := bindJSONOrForm(c, &req); err != nil {
respondError(c, err)
return
}
exp, err := resolveExpire(d.Cfg, req.ExpireValue, req.ExpireStyle)
if err != nil {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "过期策略非法")
respondError(c, err)
return
}
// v3.1:自定义提取码(合并前校验,失败快速返回)
if err := validatePickupCode(req.Code); err != nil {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "提取码非法")
respondError(c, err)
return
}
if err := reserveStorage(ctx, d.DB, d.Cfg, "chunk:"+uploadID, session.FileSize, chunkExpireTTL); err != nil {
respondError(c, err)
return
}
// 分片完整性校验(chunk_index >= 0-1 为会话头,completed 恒为 false
var completed []model.UploadChunk
if err := d.DB.WithContext(ctx).
Where("upload_id = ? AND completed = ? AND chunk_index >= 0", uploadID, true).
Find(&completed).Error; err != nil {
respondError(c, errInternal("查询分片记录失败: "+err.Error()))
return
}
if len(completed) != session.TotalChunks {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "分片不完整")
response.Fail(c, http.StatusBadRequest, "分片不完整")
return
}
// 累计大小上限校验(超限清理会话,对齐参考;动态策略上限)
if err := d.CurrentUploadPolicy().CheckSize(int64(len(completed)) * int64(session.ChunkSize)); err != nil {
if cs, ce := d.storeFor(session.Engine); ce == nil {
_ = cs.CleanChunks(ctx, uploadID, session.SavePath)
}
_ = d.DB.WithContext(ctx).Where("upload_id = ?", uploadID).Delete(&model.UploadChunk{}).Error
releaseStorage(ctx, d.DB, "chunk:"+uploadID)
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "实际上传大小超过限制")
respondError(c, err)
return
}
// 合并(引擎负责按索引有序合并+SHA256 校验)
verifyHash := func(index int) (string, error) {
var rec model.UploadChunk
err := d.DB.WithContext(ctx).
Where("upload_id = ? AND chunk_index = ?", uploadID, index).
First(&rec).Error
if err != nil {
return "", err
}
return rec.ChunkHash, nil
}
// v3:合并走会话归属引擎(会话创建时的引擎,即使中途热切换也不受影响)
mergeStore, sErr := d.storeFor(session.Engine)
if sErr != nil {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, session.FileSize)
auditRecordFailed(c, d.AuditSvc, "存储引擎不可用: "+sErr.Error())
respondError(c, mapStorageError(sErr))
return
}
size, fileHash, err := mergeStore.MergeChunks(ctx, uploadID, session.TotalChunks, verifyHash, session.SavePath)
if err != nil {
_ = mergeStore.CleanChunks(ctx, uploadID, session.SavePath)
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, session.FileSize)
auditRecordFailed(c, d.AuditSvc, "文件合并失败")
respondError(c, mapStorageError(err))
return
}
// 创建分享记录(v3.1:支持自定义提取码)
code, err := pickCustomCode(ctx, d.DB, d.Cfg, req.Code)
if err == nil {
fc := model.FileCodes{
Code: code,
FileHash: &fileHash,
IsChunked: true,
UploadID: &uploadID,
Size: session.FileSize,
ExpiredAt: exp.ExpiredAt,
ExpiredCount: exp.ExpiredCount,
UsedCount: exp.UsedCount,
Engine: session.Engine, // v3:归属引擎戳
}
// 拆分路径与文件名(对齐参考:path=dirname(save_path), uuid=basename
dir, name := splitDirBase(session.SavePath)
ext := baseExt(name)
fc.FilePath = &dir
fc.UUIDFileName = &name
fc.Prefix = trimExt(name)
fc.Suffix = ext
err = d.DB.WithContext(ctx).Create(&fc).Error
err = mapCodeConflict(err) // v3.1
}
if err == nil {
// 成功:清理分片与记录(走归属引擎)
_ = mergeStore.CleanChunks(ctx, uploadID, session.SavePath)
_ = d.DB.WithContext(ctx).Where("upload_id = ?", uploadID).Delete(&model.UploadChunk{}).Error
}
releaseStorage(ctx, d.DB, "chunk:"+uploadID)
if err != nil {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, size)
auditRecordFailed(c, d.AuditSvc, "创建分享失败")
respondError(c, errInternal("创建分享失败: "+err.Error()))
return
}
d.Limiter.Add(c, middleware.LimitUpload)
auditUploadEntry(c, code, session.FileName, session.FileSize, size)
auditRecordSuccess(c, d.AuditSvc)
response.OK(c, gin.H{"code": code, "name": session.FileName})
}
// splitDirBase 拆分相对路径为目录与文件名。
func splitDirBase(p string) (dir, base string) {
for i := len(p) - 1; i >= 0; i-- {
if p[i] == '/' {
return p[:i], p[i+1:]
}
}
return "", p
}
// baseExt 提取扩展名(含点)。
func baseExt(name string) string {
for i := len(name) - 1; i >= 0; i-- {
if name[i] == '.' {
return name[i:]
}
if name[i] == '/' {
break
}
}
return ""
}
// trimExt 去除扩展名。
func trimExt(name string) string {
ext := baseExt(name)
return name[:len(name)-len(ext)]
}
// ============ DELETE /chunk/upload/{uploadID} 取消上传 ============
// chunkCancel 取消上传并清理临时文件(对齐参考 cancel_upload)。
func (d *Deps) chunkCancel(c *gin.Context) {
if !d.requireShareLogin(c) {
return
}
uploadID := c.Param("uploadID")
if uploadID == "" {
uploadID = c.Query("upload_id")
}
ctx := c.Request.Context()
var session model.UploadChunk
if err := d.DB.WithContext(ctx).
Where("upload_id = ? AND chunk_index = -1", uploadID).
First(&session).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.Fail(c, http.StatusNotFound, "上传会话不存在")
return
}
respondError(c, errInternal("查询上传会话失败: "+err.Error()))
return
}
if session.SavePath != "" {
if cs, ce := d.storeFor(session.Engine); ce == nil {
_ = cs.CleanChunks(ctx, uploadID, session.SavePath)
}
}
if err := d.DB.WithContext(ctx).
Where("upload_id = ?", uploadID).
Delete(&model.UploadChunk{}).Error; err != nil {
respondError(c, errInternal("取消上传失败: "+err.Error()))
return
}
releaseStorage(ctx, d.DB, "chunk:"+uploadID)
response.OK(c, gin.H{"message": "上传已取消"})
}
+798
View File
@@ -0,0 +1,798 @@
// Package api 提供文件快传的 HTTP API 层:
// 分享(share)、分片上传(chunk)、预签名直传(presign)、管理端(admin)、
// 初始化向导(setup)与前端静态资源(web)。
//
// 接口语义对齐参考实现 apps/base/views.py 与 apps/admin/views.py
// 响应统一 {"code":200,"msg":"","data":...}internal/response)。
// 所有上传/下载端点经 middleware.Audit 落审计日志(需求 ③)。
package api
import (
"bytes"
"context"
"crypto/hmac"
"crypto/rand"
"crypto/sha256"
"encoding/hex"
"errors"
"fmt"
"io"
"math/big"
"mime/multipart"
"net/http"
"net/url"
"path"
"regexp"
"strconv"
"strings"
"time"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"filecodebox/internal/audit"
"filecodebox/internal/config"
"filecodebox/internal/middleware"
"filecodebox/internal/model"
"filecodebox/internal/response"
"filecodebox/internal/storage"
)
// apiError 统一的业务错误:handler 返回该错误并由 respondError 映射 HTTP 状态。
type apiError struct {
Status int // HTTP 状态码(同时写入响应体 code)
Msg string // 错误描述(中文)
}
func (e *apiError) Error() string { return e.Msg }
func errBadRequest(msg string) error { return &apiError{Status: http.StatusBadRequest, Msg: msg} }
func errForbidden(msg string) error { return &apiError{Status: http.StatusForbidden, Msg: msg} }
func errNotFound(msg string) error { return &apiError{Status: http.StatusNotFound, Msg: msg} }
func errConflict(msg string) error { return &apiError{Status: http.StatusConflict, Msg: msg} }
func errLocked(msg string) error { return &apiError{Status: http.StatusLocked, Msg: msg} }
func errInternal(msg string) error {
return &apiError{Status: http.StatusInternalServerError, Msg: msg}
}
func errInsufficient(msg string) error {
return &apiError{Status: http.StatusInsufficientStorage, Msg: msg}
}
func errNotImpl(msg string) error { return &apiError{Status: http.StatusNotImplemented, Msg: msg} }
// mapStorageError 把存储层哨兵错误映射为 apiError(对齐 README 约定):
// ErrNotFound→404、ErrInvalidPath→400、ErrUnavailable→503、ErrNotSupported→501、
// ErrRangeNotSatisfiable→416、ErrHashMismatch→400(支持 %w 包装判定)。
func mapStorageError(err error) error {
if err == nil {
return nil
}
switch {
case errors.Is(err, storage.ErrNotFound):
return errNotFound("文件不存在")
case errors.Is(err, storage.ErrInvalidPath):
return errBadRequest("非法文件路径")
case errors.Is(err, storage.ErrUnavailable):
return &apiError{Status: http.StatusServiceUnavailable, Msg: "存储服务不可用,请稍后再试"}
case errors.Is(err, storage.ErrNotSupported):
return errNotImpl("当前存储引擎不支持该操作")
case errors.Is(err, storage.ErrRangeNotSatisfiable):
return &apiError{Status: http.StatusRequestedRangeNotSatisfiable, Msg: "请求范围超出文件大小"}
case errors.Is(err, storage.ErrHashMismatch):
return errBadRequest("分片哈希校验失败,请重新上传")
}
return errInternal("存储操作失败: " + err.Error())
}
// respondError 统一错误出口:apiError 按其状态码响应,其余 500。
func respondError(c *gin.Context, err error) {
var ae *apiError
if !errors.As(err, &ae) {
ae = &apiError{Status: http.StatusInternalServerError, Msg: err.Error()}
}
// 非 2xx 且未标记审计跳过时,兜底给出 result/errorMsg(中间件还会按状态兜底)
middleware.AuditSet(c, func(e *audit.Entry) {
if e.Result == "" {
if ae.Status == 401 || ae.Status == 403 || ae.Status == 423 || ae.Status == 429 || ae.Status == 428 {
e.Result = model.AuditResultDenied
} else {
e.Result = model.AuditResultFailed
}
}
if e.ErrorMsg == "" {
e.ErrorMsg = ae.Msg
}
})
response.Fail(c, ae.Status, ae.Msg)
}
// ============ 取件码 / 下载令牌 ============
// codeChars 随机字符码字符集(对齐参考 string.ascii_uppercase + digits)。
const codeChars = "ABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789"
// generateCode 生成取件码:secret→5 位大写字母+数字;number→5 位数字(10000~99999)。
// 与参考 get_random_code 一致:冲突时重试(由调用方控制上限)。
func generateCode(style string) string {
if style == "number" {
n, _ := rand.Int(rand.Reader, big.NewInt(90000))
return strconv.FormatInt(10000+n.Int64(), 10)
}
b := make([]byte, 5)
for i := range b {
n, _ := rand.Int(rand.Reader, big.NewInt(int64(len(codeChars))))
b[i] = codeChars[n.Int64()]
}
return string(b)
}
// ============ 自定义提取码(v3.1,防撞库)============
const (
// pickupCodeMinLen 最小长度:L3 由 4 提升至 5(4 位码空间仅 168 万,
// 多 IP 分布式撞库可在限流窗口内覆盖;5 位起 6000 万空间显著提高成本)
pickupCodeMinLen = 5
pickupCodeMaxLen = 8
)
// validatePickupCode 校验自定义提取码:
// - 长度 5-8 字符(下限保证码空间 36^5≈6000 万起,配合取件错误限流防撞库);
// - 仅字母/数字(排除空格与符号,规避中文场景误输与 URL 转义问题);
// - 空串合法(空 = 使用系统随机码)。
// - 兼容:仅约束新建码,历史 4 位码仍可正常取件。
func validatePickupCode(code string) error {
code = strings.TrimSpace(code)
if code == "" {
return nil
}
if l := len(code); l < pickupCodeMinLen || l > pickupCodeMaxLen {
return errBadRequest(fmt.Sprintf("提取码长度须为 %d-%d 位", pickupCodeMinLen, pickupCodeMaxLen))
}
for _, r := range code {
if !(r >= '0' && r <= '9') && !(r >= 'a' && r <= 'z') && !(r >= 'A' && r <= 'Z') {
return errBadRequest("提取码仅支持字母和数字")
}
}
return nil
}
// pickCustomCode 解析提取码:自定义非空→查重返回;空→随机码。
// 占用时 400(不静默改码——用户可能已把链接发出去,静默改码会导致取不到件)。
func pickCustomCode(ctx context.Context, db *gorm.DB, cfg *config.Config, custom string) (string, error) {
custom = strings.ToUpper(strings.TrimSpace(custom))
if custom == "" {
return randomCode(ctx, db, cfg)
}
var count int64
if err := db.WithContext(ctx).Model(&model.FileCodes{}).Where("code = ?", custom).Count(&count).Error; err != nil {
return "", errInternal("取件码查重失败: " + err.Error())
}
if count > 0 {
return "", errBadRequest("该提取码已被占用,请换一个")
}
return custom, nil
}
// mapCodeConflict v3.1:分享记录创建失败时,若是自定义码唯一索引冲突(并发兜底,
// pickCustomCode 的预查重未覆盖竞态),转为友好 400;其余错误原样返回。
func mapCodeConflict(err error) error {
if err == nil {
return nil
}
msg := err.Error()
if strings.Contains(msg, "已被占用") ||
strings.Contains(msg, "UNIQUE constraint") ||
strings.Contains(msg, "duplicate key") ||
strings.Contains(msg, "Duplicate entry") {
return errBadRequest("该提取码已被占用,请换一个")
}
return err
}
// ============ 站点对外域名(v3.1============
// SiteDomain 站点对外域名规范化(v3.1):空串合法(分享链接用当前访问地址)。
// 接受 http(s)://host[:port] 或裸 host[:port](自动补 http://,内网场景)。
// 拒绝路径/查询/片段/用户信息/非 http(s) 协议/非法主机字符(防 javascript: 注入分享链接)。
var siteDomainHostRe = regexp.MustCompile(`^[A-Za-z0-9]([A-Za-z0-9-]*[A-Za-z0-9])?(\.[A-Za-z0-9]([A-Za-z0-9-]*[A-Za-z0-9])?)*$`)
func normalizeSiteDomain(raw string) (string, error) {
s := strings.TrimSpace(raw)
if s == "" {
return "", nil
}
s = strings.TrimRight(s, "/")
if !strings.Contains(s, "://") {
s = "http://" + s
}
u, err := url.Parse(s)
if err != nil || u.Host == "" || (u.Scheme != "http" && u.Scheme != "https") {
return "", errBadRequest("站点域名格式:http(s)://主机[:端口],如 https://share.example.com")
}
if (u.Path != "" && u.Path != "/") || u.RawQuery != "" || u.Fragment != "" || u.User != nil {
return "", errBadRequest("站点域名只填主机与端口,不带路径,如 https://share.example.com")
}
if !siteDomainHostRe.MatchString(u.Hostname()) {
return "", errBadRequest("站点域名主机名仅支持字母、数字、点与连字符")
}
if p := u.Port(); p != "" {
if n, perr := strconv.Atoi(p); perr != nil || n < 1 || n > 65535 {
return "", errBadRequest("站点域名端口须为 1-65535")
}
}
return s, nil
}
// randomCode 生成唯一取件码(查库去重,最多尝试 20 次防止极端碰撞)。
// style 为空时取配置 code_generate_typesecret/string→secret,其余→number)。
func randomCode(ctx context.Context, db *gorm.DB, cfg *config.Config) (string, error) {
style := strings.TrimSpace(cfg.GetString("code_generate_type"))
if style == "string" {
style = "secret"
}
if style != "secret" && style != "number" {
style = "number"
}
for i := 0; i < 20; i++ {
code := generateCode(style)
var count int64
if err := db.WithContext(ctx).Model(&model.FileCodes{}).
Where("code = ?", code).Count(&count).Error; err != nil {
return "", errInternal("取件码生成失败: " + err.Error())
}
if count == 0 {
return code, nil
}
}
return "", errInternal("取件码生成失败,请重试")
}
// GetSelectToken 生成下载令牌(L2HMAC-SHA256 替换拼接哈希,消除拼接歧义;
// 密钥前置为 HMAC key,窗口语义不变):
// HMAC-SHA256(key=secret, msg=code|time_factor)time_factor = unix秒/1000 - offset。
// offset=0 当前窗口、offset=1 上一窗口——下载端点同时接受两个窗口,
// 避免 ~16.7 分钟窗口边界竞态导致偶发 403。
func GetSelectToken(code, secret string, offset int) string {
timeFactor := time.Now().Unix()/1000 - int64(offset)
mac := hmac.New(sha256.New, []byte(secret))
fmt.Fprintf(mac, "%s|%d", code, timeFactor)
return hex.EncodeToString(mac.Sum(nil))
}
// VerifySelectToken 常量时间校验下载令牌(当前与上一窗口任一匹配即通过)。
func VerifySelectToken(code, secret, key string) bool {
return hmac.Equal([]byte(key), []byte(GetSelectToken(code, secret, 0))) ||
hmac.Equal([]byte(key), []byte(GetSelectToken(code, secret, 1)))
}
// ============ 过期策略 ============
// expireResult 过期策略解析结果(对齐参考 get_expire_info)。
type expireResult struct {
ExpiredAt *time.Time // nil 表示永久
ExpiredCount int // <0 按时间过期;>0 按次数
UsedCount int
}
// resolveExpire 校验 expire_style 白名单并计算过期信息。
// 对齐参考:max_save_seconds>0 时为最长保存上限(超限 403),否则默认 7 天上限;
// v2 需求 ④:style=count 时 expire_value 不得超出 max_save_count0=不限制,超限 403)。
func resolveExpire(cfg *config.Config, expireValue int, expireStyle string) (*expireResult, error) {
allowed := cfg.ExpireStyle()
okStyle := false
for _, s := range allowed {
if s == expireStyle {
okStyle = true
break
}
}
if !okStyle {
return nil, errBadRequest("过期时间类型错误")
}
if expireValue <= 0 {
return nil, errBadRequest("过期时间值必须大于 0")
}
now := time.Now()
res := &expireResult{ExpiredCount: -1, UsedCount: 0}
var expiredAt time.Time
switch expireStyle {
case "day":
expiredAt = now.AddDate(0, 0, expireValue)
case "hour":
expiredAt = now.Add(time.Duration(expireValue) * time.Hour)
case "minute":
expiredAt = now.Add(time.Duration(expireValue) * time.Minute)
case "count":
// 保存次数策略(需求 ④):max_save_count>0 时为可取次数上限,超限 403
if maxCount := cfg.MaxSaveCount(); maxCount > 0 && expireValue > maxCount {
return nil, errForbidden(fmt.Sprintf("限制次数最多为 %d 次", maxCount))
}
// 按次数过期:固定保留 1 天时间兜底(对齐参考)
expiredAt = now.AddDate(0, 0, 1)
res.ExpiredCount = expireValue
case "forever":
res.ExpiredAt = nil
res.ExpiredCount = -1
return res, nil
default:
expiredAt = now.AddDate(0, 0, 1)
}
// 最长保存时间限制
maxSeconds := cfg.MaxSaveSeconds()
maxDelta := 7 * 24 * time.Hour
if maxSeconds > 0 {
maxDelta = time.Duration(maxSeconds) * time.Second
}
if expiredAt.Sub(now) > maxDelta {
return nil, errForbidden(fmt.Sprintf("限制最长时间为 %s,可换用其他方式", formatDurationCN(maxDelta)))
}
res.ExpiredAt = &expiredAt
return res, nil
}
// formatDurationCN 把时长格式化为中文描述(对齐参考 max_save_times_desc)。
func formatDurationCN(d time.Duration) string {
sec := int64(d.Seconds())
days := sec / 86400
hours := sec % 86400 / 3600
minutes := sec % 3600 / 60
seconds := sec % 60
var parts []string
if days > 0 {
parts = append(parts, fmt.Sprintf("%d天", days))
}
if hours > 0 {
parts = append(parts, fmt.Sprintf("%d小时", hours))
}
if minutes > 0 {
parts = append(parts, fmt.Sprintf("%d分钟", minutes))
}
if seconds > 0 {
parts = append(parts, fmt.Sprintf("%d秒", seconds))
}
if len(parts) == 0 {
return "0秒"
}
return strings.Join(parts, "")
}
// ============ 存储路径 / 容量预留 ============
// storeFor v3:按归属引擎取存储实例(空戳/未知名回落当前引擎,兼容历史数据)。
func (d *Deps) storeFor(engine string) (storage.Storage, error) {
if engine == "" || !storage.ValidEngine(engine) {
return d.Store, nil
}
if engine == d.Store.CurrentName() {
return d.Store, nil
}
return d.Store.EngineOf(engine)
}
// buildSavePath 生成上传文件的存储相对路径(对齐参考 get_file_path_name):
// [storage_path/]share/data/YYYY/MM/DD/<uuid>/<清理后文件名>。
func buildSavePath(cfg *config.Config, rawName string, fileUUID string) (dirPath, prefix, suffix, cleanName, savePath string) {
today := time.Now().Format("2006/01/02")
cleanName = storage.SanitizeFileName(rawName)
ext := path.Ext(cleanName)
prefix = strings.TrimSuffix(cleanName, ext)
suffix = ext
base := "share/data/" + today + "/" + fileUUID
if sp := strings.Trim(cfg.GetString("storage_path"), "/"); sp != "" {
base = sp + "/" + base
}
dirPath = base
savePath = base + "/" + cleanName
return
}
// reserveStorage 原子预留上传容量(对齐参考 quota.reserve_storage):
// storageLimit<=0 或 size=0 时不限制直接返回;
// 通过单条 INSERT...SELECT 条件写入保证 (已用+已预留+本次) <= limit,超限返回 507。
func reserveStorage(ctx context.Context, db *gorm.DB, cfg *config.Config, token string, size int64, ttl time.Duration) error {
limit := cfg.GetInt64("storageLimit")
if limit <= 0 || size <= 0 {
return nil
}
now := time.Now()
expiresAt := now.Add(ttl)
// 清理同 token 的过期预留
if err := db.WithContext(ctx).
Where("token = ? AND expires_at <= ?", token, now).
Delete(&model.StorageReservation{}).Error; err != nil {
return errInternal("容量预留失败: " + err.Error())
}
// 同 token 已有生效预留:大小一致则幂等返回,不一致报冲突(对齐参考 409)
var existing model.StorageReservation
err := db.WithContext(ctx).
Where("token = ? AND expires_at > ?", token, now).First(&existing).Error
if err == nil {
if existing.Size == size {
return nil
}
return errConflict("上传容量预留信息不一致")
}
if !errors.Is(err, gorm.ErrRecordNotFound) {
return errInternal("容量预留失败: " + err.Error())
}
// 原子条件插入:已用 + 生效预留 + 本次 <= 上限。
// Info5Postgres READ COMMITTED 下并发 INSERT..SELECT 可能同时读到相同快照
// 而轻微超额记账,故包事务并用事务级 advisory lock 串行化配额判定
// (SQLite 写本身串行,无需加锁)。
lockFn := func(tx *gorm.DB) error {
if tx.Dialector.Name() == config.DBDriverPostgres {
return tx.Exec(`SELECT pg_advisory_xact_lock(?)`, quotaLockKey).Error
}
return nil
}
var insertErr error
txErr := db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
if err := lockFn(tx); err != nil {
return err
}
res := tx.Exec(`
INSERT INTO storage_reservations (token, size, expires_at)
SELECT ?, ?, ?
WHERE (
COALESCE((SELECT COALESCE(SUM(size),0) FROM file_codes), 0)
+ COALESCE((SELECT COALESCE(SUM(size),0) FROM storage_reservations WHERE expires_at > ?), 0)
+ ?
) <= ?`,
token, size, expiresAt, now, size, limit)
insertErr = res.Error
if res.Error != nil {
return res.Error // 触发回滚(同 token 冲突分支在外层处理)
}
if res.RowsAffected == 0 {
return errInsufficient("存储空间已达到管理员设置的容量上限")
}
return nil
})
if txErr != nil {
// 并发冲突回退:检查是否已有同 token 同大小的生效预留(对齐参考并发分支)
var cnt int64
_ = db.WithContext(ctx).Model(&model.StorageReservation{}).
Where("token = ? AND size = ? AND expires_at > ?", token, size, now).
Count(&cnt).Error
if cnt > 0 {
return nil
}
var ins *apiError
if errors.As(txErr, &ins) && ins.Status == http.StatusInsufficientStorage {
return txErr // 507:真实容量不足
}
if insertErr != nil && errors.Is(insertErr, txErr) {
return errInternal("容量预留失败: " + insertErr.Error())
}
return errInternal("容量预留失败: " + txErr.Error())
}
return nil
}
// quotaLockKey Postgres advisory lock 键(配额判定的事务级串行化)。
const quotaLockKey int64 = 0x46434251 // "FCBQ"
// releaseStorage 释放容量预留(幂等)。
func releaseStorage(ctx context.Context, db *gorm.DB, token string) {
_ = db.WithContext(ctx).Where("token = ?", token).Delete(&model.StorageReservation{}).Error
}
// ============ 文件类型校验(对齐 apps/base/file_validation.py============
// fileKind 已知文件类型:扩展名 / MIME / magic bytes。
type fileKind struct {
name string
extensions []string
mimes []string
signatures [][]byte
}
var fileKinds = []fileKind{
{"png", []string{".png"}, []string{"image/png"}, [][]byte{{0x89, 'P', 'N', 'G', '\r', '\n', 0x1a, '\n'}}},
{"jpg", []string{".jpg", ".jpeg"}, []string{"image/jpeg"}, [][]byte{{0xff, 0xd8, 0xff}}},
{"gif", []string{".gif"}, []string{"image/gif"}, [][]byte{[]byte("GIF87a"), []byte("GIF89a")}},
{"webp", []string{".webp"}, []string{"image/webp"}, nil},
{"bmp", []string{".bmp"}, []string{"image/bmp", "image/x-ms-bmp"}, [][]byte{[]byte("BM")}},
{"pdf", []string{".pdf"}, []string{"application/pdf"}, [][]byte{[]byte("%PDF")}},
{"zip", []string{".zip", ".docx", ".xlsx", ".pptx", ".apk", ".jar"},
[]string{"application/zip", "application/x-zip-compressed",
"application/vnd.openxmlformats-officedocument.wordprocessingml.document",
"application/vnd.openxmlformats-officedocument.spreadsheetml.sheet",
"application/vnd.openxmlformats-officedocument.presentationml.presentation",
"application/java-archive", "application/vnd.android.package-archive"},
[][]byte{[]byte("PK\x03\x04"), []byte("PK\x05\x06"), []byte("PK\x07\x08")}},
{"rar", []string{".rar"}, []string{"application/x-rar-compressed", "application/vnd.rar"},
[][]byte{[]byte("Rar!\x1a\x07\x00"), []byte("Rar!\x1a\x07\x01\x00")}},
{"7z", []string{".7z"}, []string{"application/x-7z-compressed"}, [][]byte{{0x37, 0x7a, 0xbc, 0xaf, 0x27, 0x1c}}},
{"gz", []string{".gz", ".tgz"}, []string{"application/gzip", "application/x-gzip"}, [][]byte{{0x1f, 0x8b}}},
{"mp3", []string{".mp3"}, []string{"audio/mpeg"}, [][]byte{[]byte("ID3"), {0xff, 0xfb}, {0xff, 0xf3}, {0xff, 0xf2}}},
{"mp4", []string{".mp4", ".m4a", ".mov"}, []string{"video/mp4", "audio/mp4", "video/quicktime"}, nil},
{"exe", []string{".exe", ".dll", ".sys"}, []string{"application/x-msdownload", "application/x-dosexec"}, [][]byte{[]byte("MZ")}},
{"elf", []string{".elf", ".so", ".o"}, []string{"application/x-executable"}, [][]byte{{0x7f, 'E', 'L', 'F'}}},
}
// knownExtensions 全部已知扩展名集合。
var knownExtensions = func() map[string]bool {
m := map[string]bool{}
for _, k := range fileKinds {
for _, ext := range k.extensions {
m[ext] = true
}
}
return m
}()
// isTypeAllowed 判断文件是否在 allowed_file_types 白名单内("*"/*/* 放行全部)。
func isTypeAllowed(cfg *config.Config, fileName, contentType string) bool {
allowed := cfg.AllowedFileTypes()
if len(allowed) == 0 {
return true
}
name := strings.ToLower(strings.TrimSpace(fileName))
ct := strings.ToLower(strings.TrimSpace(contentType))
for _, rule := range allowed {
rule = strings.ToLower(strings.TrimSpace(rule))
switch {
case rule == "*" || rule == "*/*":
return true
case strings.Contains(rule, "/"):
if ok, _ := path.Match(rule, ct); ok {
return true
}
default:
if !strings.HasPrefix(rule, ".") {
rule = "." + rule
}
if strings.HasSuffix(name, rule) {
return true
}
}
}
return false
}
// detectFileKind 按文件头识别类型(对齐参考:RIFF/WEBP、ftyp/mp4 与前缀签名表)。
func detectFileKind(header []byte) *fileKind {
if len(header) == 0 {
return nil
}
if len(header) >= 12 && string(header[:4]) == "RIFF" && string(header[8:12]) == "WEBP" {
for i := range fileKinds {
if fileKinds[i].name == "webp" {
return &fileKinds[i]
}
}
}
if len(header) >= 12 && string(header[4:8]) == "ftyp" {
for i := range fileKinds {
if fileKinds[i].name == "mp4" {
return &fileKinds[i]
}
}
}
var best *fileKind
bestLen := 0
for i := range fileKinds {
for _, sig := range fileKinds[i].signatures {
if len(sig) > 0 && len(header) >= len(sig) && string(header[:len(sig)]) == string(sig) {
if len(sig) > bestLen {
bestLen = len(sig)
best = &fileKinds[i]
}
}
}
}
return best
}
// validateFileMagic 白名单 + magic bytes 防伪造(对齐参考 validate_file_magic)。
// header 为文件前 64 字节,可为空(空则只校验白名单)。
func validateFileMagic(cfg *config.Config, fileName, contentType string, header []byte) error {
if !isTypeAllowed(cfg, fileName, contentType) {
return errForbidden("不允许上传该类型文件")
}
if len(header) == 0 {
return nil
}
ext := strings.ToLower(path.Ext(fileName))
ct := strings.ToLower(strings.TrimSpace(contentType))
detected := detectFileKind(header)
if knownExtensions[ext] {
if detected == nil {
return errForbidden("文件内容与扩展名不匹配,疑似伪造类型")
}
matched := false
for _, e := range detected.extensions {
if e == ext {
matched = true
break
}
}
if !matched {
return errForbidden("文件内容与扩展名不匹配,疑似伪造类型")
}
}
if ct != "" {
for _, k := range fileKinds {
for _, m := range k.mimes {
if m == ct {
// 声明了已知 MIME:内容必须匹配
if detected == nil {
return errForbidden("文件内容与 Content-Type 不匹配,疑似伪造类型")
}
matched := false
for _, m2 := range detected.mimes {
if m2 == ct {
matched = true
break
}
}
if !matched {
return errForbidden("文件内容与 Content-Type 不匹配,疑似伪造类型")
}
break
}
}
}
}
return nil
}
// readMultipartHeader 读取上传文件前 n 字节并 seek 回起点(用于 magic 校验)。
func readMultipartHeader(f multipart.File, n int64) []byte {
if f == nil {
return nil
}
buf := make([]byte, n)
nread, _ := f.Read(buf)
_, _ = f.Seek(0, 0)
if nread <= 0 {
return nil
}
return buf[:nread]
}
// ============ 杂项 ============
// humanSize 把字节数转成人类可读描述(B/KB/MB/GB 自适应;
// v2 需求④:max_file_size 支持子 MB 上限,固定 MB 格式会显示 0.00 MB)。
func humanSize(n int64) string {
const kb, mb, gb = int64(1024), int64(1024 * 1024), int64(1024 * 1024 * 1024)
switch {
case n >= gb:
return fmt.Sprintf("%.2f GB", float64(n)/float64(gb))
case n >= mb:
return fmt.Sprintf("%.2f MB", float64(n)/float64(mb))
case n >= kb:
return fmt.Sprintf("%.2f KB", float64(n)/float64(kb))
default:
return fmt.Sprintf("%d B", n)
}
}
// contentDisposition 生成 RFC 5987 附件头(对齐参考 filename*=UTF-8” 格式)。
func contentDisposition(name string) string {
quoted := urlPathEscape(name)
return "attachment; filename*=UTF-8''" + quoted
}
// urlPathEscape RFC 5987 风格百分号编码(等价 urllib.parse.quote(safe=”))。
func urlPathEscape(s string) string {
var b strings.Builder
for _, r := range []byte(s) {
if (r >= 'A' && r <= 'Z') || (r >= 'a' && r <= 'z') || (r >= '0' && r <= '9') ||
r == '-' || r == '_' || r == '.' || r == '~' {
b.WriteByte(r)
} else {
fmt.Fprintf(&b, "%%%02X", r)
}
}
return b.String()
}
// parseISOTime 解析 ISO 8601 / RFC3339 / 日期时间字符串,失败返回错误。
func parseISOTime(s string) (time.Time, error) {
s = strings.TrimSpace(s)
if s == "" {
return time.Time{}, errors.New("空时间")
}
layouts := []string{
time.RFC3339Nano, time.RFC3339,
"2006-01-02T15:04:05", "2006-01-02 15:04:05", "2006-01-02",
}
for _, layout := range layouts {
if t, err := time.Parse(layout, s); err == nil {
return t, nil
}
}
return time.Time{}, fmt.Errorf("时间格式错误: %s", s)
}
// requireUploadLimit 上传限流入口检查(进入即校验,超限 423;成功后由 handler 显式 Add)。
// 对齐参考:FastAPI Depends(ip_limit["upload"]) 在进入时 check。
func requireUploadLimit(c *gin.Context, limiter *middleware.RateLimiter) bool {
if allowed, _ := limiter.Check(c, middleware.LimitUpload); !allowed {
response.Fail(c, http.StatusLocked, "请求次数过多,请稍后再试")
return false
}
return true
}
// bindJSONOrForm 兼容 JSON 与表单请求体绑定到结构体。
func bindJSONOrForm(c *gin.Context, obj any) error {
ct := c.GetHeader("Content-Type")
if strings.Contains(ct, "application/json") {
if err := c.ShouldBindJSON(obj); err != nil {
return errBadRequest("请求体格式错误: " + err.Error())
}
return nil
}
// v3.1.1 兼容归一化:旧前端 bundlefetch 字符串 body 默认 text/plain)发的
// 是 text/plain + urlencoded 格式。此类请求改写 Content-Type 后走表单绑定,
// 否则 ShouldBind 对 text/plain 不解析,非空字段全部丢失。
base := ct
if i := strings.IndexByte(ct, ';'); i >= 0 {
base = ct[:i]
}
if strings.EqualFold(strings.TrimSpace(base), "text/plain") &&
c.Request != nil && c.Request.Body != nil {
if raw, err := io.ReadAll(c.Request.Body); err == nil {
trimmed := bytes.TrimSpace(raw)
// JSON 形态(无头/误标 text/plain):改写后按 JSON 绑定(须先于 ParseQuery 判断,
// 否则形如 {"a":1} 的 JSON 会被 ParseQuery 误判为单键 urlencoded
if len(trimmed) > 0 && trimmed[0] == '{' {
c.Request.Header.Set("Content-Type", "application/json")
c.Request.Body = io.NopCloser(bytes.NewReader(raw))
return c.ShouldBindJSON(obj)
}
if vals, perr := url.ParseQuery(string(raw)); perr == nil && len(vals) > 0 {
// urlencoded 形态:改写 Content-Type 走表单绑定
c.Request.Header.Set("Content-Type", "application/x-www-form-urlencoded")
c.Request.Body = io.NopCloser(bytes.NewReader(raw))
return c.ShouldBind(obj)
}
// 其他形态:还原 body 让 ShouldBind 按原样处理
c.Request.Body = io.NopCloser(bytes.NewReader(raw))
}
}
if err := c.ShouldBind(obj); err != nil {
return errBadRequest("请求体格式错误: " + err.Error())
}
return nil
}
// auditUploadEntry 填充上传类审计业务字段的便捷函数。
func auditUploadEntry(c *gin.Context, code, name string, size, transferred int64) {
middleware.AuditSet(c, func(e *audit.Entry) {
e.FileCode = code
e.FileName = name
e.SizeBytes = size
e.TransferredBytes = transferred
})
}
// auditRecordSuccess / auditRecordFailed 显式落库便捷函数。
func auditRecordSuccess(c *gin.Context, svc *audit.Service) {
middleware.AuditRecordRequest(c, svc, model.AuditResultSuccess, "")
}
func auditRecordFailed(c *gin.Context, svc *audit.Service, msg string) {
middleware.AuditRecordRequest(c, svc, model.AuditResultFailed, msg)
}
// uuidHex 生成 32 位十六进制随机串(对齐参考 uuid4().hex)。
func uuidHex() string {
b := make([]byte, 16)
_, _ = rand.Read(b)
// 设置版本号与变体位以保持 uuid4 兼容格式
b[6] = (b[6] & 0x0f) | 0x40
b[8] = (b[8] & 0x3f) | 0x80
return hex.EncodeToString(b)
}
// uuidCanonical 生成带连字符的 UUID 字符串(upload_id 用)。
func uuidCanonical() string {
h := uuidHex()
return h[0:8] + "-" + h[8:12] + "-" + h[12:16] + "-" + h[16:20] + "-" + h[20:32]
}
+264
View File
@@ -0,0 +1,264 @@
package api
import (
"errors"
"fmt"
"testing"
"time"
"filecodebox/internal/config"
"filecodebox/internal/storage"
)
// newTestConfig 构造测试配置(defaults 基线,无 KV 覆盖;需求 ⑧ 默认 sqlite,无需真实数据库)。
func newTestConfig(t *testing.T) *config.Config {
t.Helper()
t.Setenv("FCB_DB_DRIVER", "sqlite")
t.Setenv("FCB_DB_DSN", "")
cfg, err := config.New()
if err != nil {
t.Fatalf("config.New 失败: %v", err)
}
return cfg
}
// TestGetSelectTokenWindow 验证下载令牌的窗口号逻辑(对齐参考 get_select_token)。
func TestGetSelectTokenWindow(t *testing.T) {
code := "AB12C"
secret := "test-secret"
tok0 := GetSelectToken(code, secret, 0)
tok1 := GetSelectToken(code, secret, 1)
if tok0 == "" || tok1 == "" {
t.Fatal("令牌不应为空")
}
// 同一窗口内 offset=0 的两次生成必须一致(确定性问题)
if tok0 != GetSelectToken(code, secret, 0) {
t.Fatal("同窗口令牌应确定一致")
}
// 不同 offset 的令牌必然不同(窗口号不同)
if tok0 == tok1 {
t.Fatal("offset=0 与 offset=1 的令牌应不同")
}
// 不同 code / secret 的令牌不同
if tok0 == GetSelectToken("ZZ999", secret, 0) {
t.Fatal("不同取件码的令牌应不同")
}
if tok0 == GetSelectToken(code, "other-secret", 0) {
t.Fatal("不同密钥的令牌应不同")
}
// 令牌为 64 位十六进制(sha256 hex
if len(tok0) != 64 {
t.Fatalf("令牌长度应为 64,实际 %d", len(tok0))
}
// 窗口号公式:unix/1000 - offset(秒级窗口约 16.7 分钟)
now := time.Now().Unix()
if now/1000 == (now+1100)/1000 {
// 仅当测试跨越窗口边界才跳过该断言(罕见,容忍)
t.Log("测试跨越窗口边界,跳过窗口公式断言")
}
}
// TestMapStorageError 验证存储哨兵错误→HTTP 状态映射表。
func TestMapStorageError(t *testing.T) {
cases := []struct {
name string
err error
expect int
}{
{"NotFound 直返", storage.ErrNotFound, 404},
{"NotFound 包装", fmt.Errorf("引擎内层: %w", storage.ErrNotFound), 404},
{"InvalidPath", storage.ErrInvalidPath, 400},
{"InvalidPath 包装", fmt.Errorf("webdav: %w", storage.ErrInvalidPath), 400},
{"Unavailable", storage.ErrUnavailable, 503},
{"Unavailable 包装", fmt.Errorf("s3: %w", storage.ErrUnavailable), 503},
{"NotSupported", storage.ErrNotSupported, 501},
{"NotSupported 包装", fmt.Errorf("local: %w", storage.ErrNotSupported), 501},
{"RangeNotSatisfiable", storage.ErrRangeNotSatisfiable, 416},
{"HashMismatch", storage.ErrHashMismatch, 400},
{"未知错误", errors.New("其他错误"), 500},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
mapped := mapStorageError(tc.err)
var ae *apiError
if !errors.As(mapped, &ae) {
t.Fatalf("应映射为 apiError,得到 %T", mapped)
}
if ae.Status != tc.expect {
t.Fatalf("状态码应为 %d,实际 %d", tc.expect, ae.Status)
}
})
}
}
// TestResolveExpire 验证过期策略解析(白名单/上限/各 style)。
func TestResolveExpire(t *testing.T) {
cfg := newTestConfig(t)
// 非法 style
if _, err := resolveExpire(cfg, 1, "year"); err == nil {
t.Fatal("非白名单 style 应报错")
}
// 非法 value
if _, err := resolveExpire(cfg, 0, "day"); err == nil {
t.Fatal("expire_value<=0 应报错")
}
// day7 天内合法
res, err := resolveExpire(cfg, 3, "day")
if err != nil {
t.Fatalf("3 天应合法: %v", err)
}
if res.ExpiredAt == nil || res.ExpiredCount != -1 {
t.Fatal("day 类型应有 expired_at 且 expired_count=-1")
}
// 超过 7 天上限
if _, err := resolveExpire(cfg, 30, "day"); err == nil {
t.Fatal("超过 7 天上限应报错")
}
// count:按次数
res, err = resolveExpire(cfg, 5, "count")
if err != nil {
t.Fatalf("count 应合法: %v", err)
}
if res.ExpiredCount != 5 {
t.Fatalf("count 类型 expired_count 应为 5,实际 %d", res.ExpiredCount)
}
// forever:永久
res, err = resolveExpire(cfg, 1, "forever")
if err != nil {
t.Fatalf("forever 应合法: %v", err)
}
if res.ExpiredAt != nil || res.ExpiredCount != -1 {
t.Fatal("forever 应为 expired_at=nil 且 expired_count=-1")
}
}
// TestGenerateCode 验证取件码格式。
func TestGenerateCode(t *testing.T) {
for i := 0; i < 50; i++ {
num := generateCode("number")
if len(num) != 5 {
t.Fatalf("数字码应为 5 位,实际 %q", num)
}
for _, ch := range num {
if ch < '0' || ch > '9' {
t.Fatalf("数字码含非数字字符: %q", num)
}
}
secret := generateCode("secret")
if len(secret) != 5 {
t.Fatalf("字符码应为 5 位,实际 %q", secret)
}
}
}
// TestParseRangeHeader 验证 Range 头解析(对齐 HTTP 语义)。
func TestParseRangeHeader(t *testing.T) {
// 全量(无 Range
if parseRangeHeader("", 1000) != nil {
t.Fatal("无 Range 头应返回 nil")
}
// 标准区间
r := parseRangeHeader("bytes=0-99", 1000)
if r == nil || r.Start != 0 || r.End != 99 {
t.Fatalf("bytes=0-99 解析错误: %+v", r)
}
// 开区间到末尾
r = parseRangeHeader("bytes=500-", 1000)
if r == nil || r.Start != 500 || r.End != -1 {
t.Fatalf("bytes=500- 解析错误: %+v", r)
}
// 后缀区间(最后 100 字节)
r = parseRangeHeader("bytes=-100", 1000)
if r == nil || r.Start != 900 || r.End != -1 {
t.Fatalf("bytes=-100 解析错误: %+v", r)
}
// 后缀超长:截断到全文件
r = parseRangeHeader("bytes=-5000", 1000)
if r == nil || r.Start != 0 {
t.Fatalf("bytes=-5000 应从头开始: %+v", r)
}
// 多区间不支持→回退全量
if parseRangeHeader("bytes=0-1,5-6", 1000) != nil {
t.Fatal("多区间应返回 nil(回退全量)")
}
// 非法格式
if parseRangeHeader("items=0-1", 1000) != nil {
t.Fatal("非 bytes 单位应返回 nil")
}
if parseRangeHeader("bytes=abc-", 1000) != nil {
t.Fatal("非法数字应返回 nil")
}
}
// TestParseISOTime 验证时间解析的多格式兼容。
func TestParseISOTime(t *testing.T) {
valid := []string{
"2025-01-01T00:00:00Z",
"2025-01-01T08:00:00+08:00",
"2025-01-01 08:00:00",
"2025-01-01",
}
for _, s := range valid {
if _, err := parseISOTime(s); err != nil {
t.Fatalf("%q 应解析成功: %v", s, err)
}
}
if _, err := parseISOTime("not-a-time"); err == nil {
t.Fatal("非法时间应报错")
}
if _, err := parseISOTime(""); err == nil {
t.Fatal("空串应报错")
}
}
// TestFormatDurationCN 验证中文时长描述。
func TestFormatDurationCN(t *testing.T) {
cases := []struct {
d time.Duration
expect string
}{
{7 * 24 * time.Hour, "7天"},
{90 * time.Minute, "1小时30分钟"},
{45 * time.Second, "45秒"},
}
for _, tc := range cases {
if got := formatDurationCN(tc.d); got != tc.expect {
t.Fatalf("formatDurationCN(%v)=%q,期望 %q", tc.d, got, tc.expect)
}
}
}
// TestFileMagicValidation 验证 magic bytes 防伪造。
func TestFileMagicValidation(t *testing.T) {
cfg := newTestConfig(t)
// 白名单 * 全放行
if err := validateFileMagic(cfg, "a.txt", "", nil); err != nil {
t.Fatalf("白名单 * 应放行: %v", err)
}
// PNG 内容 + .exe 扩展名 → 拒绝(伪造)
if err := validateFileMagic(cfg, "evil.exe", "", []byte{0x89, 'P', 'N', 'G', '\r', '\n', 0x1a, '\n'}); err == nil {
t.Fatal("PNG 内容伪装 exe 应拒绝")
}
// PNG 内容 + .png 扩展名 → 通过
if err := validateFileMagic(cfg, "ok.png", "image/png", []byte{0x89, 'P', 'N', 'G', '\r', '\n', 0x1a, '\n'}); err != nil {
t.Fatalf("真 PNG 应通过: %v", err)
}
// 文本内容 + .png 扩展名 → 拒绝
if err := validateFileMagic(cfg, "fake.png", "", []byte("hello world, this is text")); err == nil {
t.Fatal("文本伪装 png 应拒绝")
}
}
// TestSanitizePathBuild 验证存储路径构造不含穿越。
func TestSanitizePathBuild(t *testing.T) {
cfg := newTestConfig(t)
_, _, _, clean, savePath := buildSavePath(cfg, "../../etc/passwd", "uuid-123")
if clean != "etc_passwd" && clean != "passwd" {
t.Logf("清理后的文件名: %q", clean)
}
if _, ok := storage.SanitizePath(savePath); !ok {
t.Fatalf("构造的 savePath 应通过安全校验: %q", savePath)
}
}
+45
View File
@@ -0,0 +1,45 @@
// policy.go — v2 上传策略统一读取与校验(需求 ④⑩)。
//
// 管理端在后台设置页修改策略(settings KVt1 schema)后,上传链路
// share/file、chunk、presign)每次请求实时读取当前策略并动态校验:
// - 大小上限:max_file_size0=回落 uploadSize,语义见 config.MaxFileSize);
// - 类型白名单:allowed_file_types"*" 不限制),由 validateFileMagic 统一执行;
// - 保存策略:expire_style 白名单 + max_save_seconds 时间上限 + max_save_count
// 次数上限,统一在 resolveExpirehelpers.go)执行。
//
// 超限返回 403(超出策略限制)/400(参数非法),错误信息为中文。
package api
import (
"fmt"
)
// UploadPolicy 当前生效的上传策略快照(每次上传请求实时读取,管理端改动立即生效)。
type UploadPolicy struct {
MaxFileSize int64 // 单文件大小上限(字节),0=不限制
AllowedTypes []string // 类型白名单,"*" 不限制
ExpireStyles []string // 允许的过期方式白名单
MaxSaveSeconds int64 // 最长保存秒数,0=不限制(默认 7 天兜底)
MaxSaveCount int // 单次分享最大可取次数上限,0=不限制
}
// CurrentUploadPolicy 读取当前上传策略快照。
// 上传页亦通过 GET /api/v1/config 的 policy 字段读取同一组值做动态渲染。
func (d *Deps) CurrentUploadPolicy() UploadPolicy {
cfg := d.Cfg
return UploadPolicy{
MaxFileSize: cfg.MaxFileSize(),
AllowedTypes: cfg.AllowedFileTypes(),
ExpireStyles: cfg.ExpireStyle(),
MaxSaveSeconds: cfg.MaxSaveSeconds(),
MaxSaveCount: cfg.MaxSaveCount(),
}
}
// CheckSize 校验单文件大小是否超出策略上限(超出返回 403,文案对齐参考实现)。
func (p UploadPolicy) CheckSize(size int64) error {
if p.MaxFileSize > 0 && size > p.MaxFileSize {
return errForbidden(fmt.Sprintf("大小超过限制,最大为%s", humanSize(p.MaxFileSize)))
}
return nil
}
+571
View File
@@ -0,0 +1,571 @@
package api
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"mime/multipart"
"net/http"
"net/http/httptest"
"path/filepath"
"strings"
"testing"
"github.com/gin-gonic/gin"
"filecodebox/internal/audit"
"filecodebox/internal/cache"
"filecodebox/internal/config"
"filecodebox/internal/database"
"filecodebox/internal/middleware"
"filecodebox/internal/settings"
"filecodebox/internal/storage"
)
// ============ 测试环境装配(真实 sqlite + 内存缓存 + 本地存储)============
// newPolicyTestDeps 构造带真实依赖的 Depssqlite 文件库(t.TempDir)、
// 本地存储引擎、内存缓存限流器与审计服务(需求 ⑧ 默认形态)。
func newPolicyTestDeps(t *testing.T) *Deps {
t.Helper()
gin.SetMode(gin.TestMode)
dir := t.TempDir()
t.Setenv("FCB_DB_DRIVER", "sqlite")
t.Setenv("FCB_DB_DSN", filepath.Join(dir, "test.db"))
cfg, err := config.New()
if err != nil {
t.Fatalf("config.New: %v", err)
}
ctx := context.Background()
db, err := database.Open(ctx, database.Options{Driver: config.DBDriverSQLite, DSN: filepath.Join(dir, "test.db")})
if err != nil {
t.Fatalf("database.Open: %v", err)
}
t.Cleanup(func() { _ = database.Close(db) })
if err := database.Migrate(ctx, db); err != nil {
t.Fatalf("database.Migrate: %v", err)
}
mgr, err := settings.NewManager(ctx, db, cfg)
if err != nil {
t.Fatalf("settings.NewManager: %v", err)
}
store, err := storage.NewLocalStorage(filepath.Join(dir, "storage"))
if err != nil {
t.Fatalf("storage.NewLocalStorage: %v", err)
}
// v3:包装为 Managerbuild 直接返回 local 实例,测试无需真实多引擎)
storeMgr := storage.NewManager("local", store, func(string) (storage.Storage, error) {
return storage.NewLocalStorage(filepath.Join(dir, "storage"))
})
return &Deps{
DB: db,
Cfg: cfg,
Mgr: mgr,
AuditSvc: audit.NewService(audit.NewDBSink(db)),
Limiter: middleware.NewRateLimiter(cache.NewMemory(), nil),
Store: storeMgr,
Version: "test",
}
}
// ============ 请求构造辅助 ============
// invoke 以给定请求调用 handler 并返回响应。
func invoke(handler gin.HandlerFunc, req *http.Request) *httptest.ResponseRecorder {
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = req
handler(c)
return w
}
// patchConfig 以 JSON 调用 PATCH /admin/config/update。
func patchConfig(d *Deps, patch map[string]any) *httptest.ResponseRecorder {
raw, _ := json.Marshal(patch)
req := httptest.NewRequest(http.MethodPatch, "/admin/config/update", bytes.NewReader(raw))
req.Header.Set("Content-Type", "application/json")
return invoke(d.adminConfigUpdate, req)
}
// getConfig 调用 GET /admin/config/get。
func getConfig(d *Deps) *httptest.ResponseRecorder {
return invoke(d.adminConfigGet, httptest.NewRequest(http.MethodGet, "/admin/config/get", nil))
}
// getPublicConfig 调用 GET /api/v1/config。
func getPublicConfig(d *Deps) *httptest.ResponseRecorder {
return invoke(d.publicConfig, httptest.NewRequest(http.MethodGet, "/api/v1/config", nil))
}
// uploadFile 以 multipart 表单调用 POST /share/file。
func uploadFile(d *Deps, name string, content []byte, fields map[string]string) *httptest.ResponseRecorder {
body := &bytes.Buffer{}
mw := multipart.NewWriter(body)
fw, err := mw.CreateFormFile("file", name)
if err != nil {
panic(err)
}
_, _ = fw.Write(content)
for k, v := range fields {
_ = mw.WriteField(k, v)
}
_ = mw.Close()
req := httptest.NewRequest(http.MethodPost, "/share/file", body)
req.Header.Set("Content-Type", mw.FormDataContentType())
return invoke(d.shareFile, req)
}
// chunkInitJSON 以 JSON 调用 POST /chunk/upload/init。
func chunkInitJSON(d *Deps, payload string) *httptest.ResponseRecorder {
req := httptest.NewRequest(http.MethodPost, "/chunk/upload/init", strings.NewReader(payload))
req.Header.Set("Content-Type", "application/json")
return invoke(d.chunkInit, req)
}
// respBody 解析统一响应体。
func respBody(t *testing.T, w *httptest.ResponseRecorder) (code int, data map[string]any) {
t.Helper()
var body struct {
Code int `json:"code"`
Msg string `json:"msg"`
Data map[string]any `json:"data"`
}
if err := json.Unmarshal(w.Body.Bytes(), &body); err != nil {
t.Fatalf("响应解析失败: %v; body=%s", err, w.Body.String())
}
return body.Code, body.Data
}
// pngMagic 最小合法 PNG 头(magic 校验可识别)。
var pngMagic = []byte{0x89, 'P', 'N', 'G', '\r', '\n', 0x1a, '\n'}
// ============ ① 公开 configv2 展示与策略字段下发 ============
// TestPublicConfigV2Fields 验证 /api/v1/config 下发背景/页脚/备案/通知与策略范围,
// 且响应不包含任何敏感键(admin_token/jwt_secret)。
func TestPublicConfigV2Fields(t *testing.T) {
d := newPolicyTestDeps(t)
// 管理端先设置 v2 展示字段
if w := patchConfig(d, map[string]any{
"background_url": "https://cdn.example.com/bg.png",
"footer_text": "自定义页脚内容",
"footer_beian": "京ICP备2024xxxxxx号-1",
"notify_enabled": 0,
"max_save_count": 5,
}); w.Code != 200 {
t.Fatalf("patchConfig 失败: %d %s", w.Code, w.Body.String())
}
w := getPublicConfig(d)
code, data := respBody(t, w)
if code != 200 {
t.Fatalf("publicConfig code=%d", code)
}
cfgMap, _ := data["config"].(map[string]any)
if cfgMap == nil {
t.Fatal("响应缺少 config 对象")
}
for key, want := range map[string]any{
"background_url": "https://cdn.example.com/bg.png",
"footer_text": "自定义页脚内容",
"footer_beian": "京ICP备2024xxxxxx号-1",
"notify_enabled": float64(0),
"notify_title": "系统通知",
} {
if got := cfgMap[key]; got != want {
t.Fatalf("config.%s = %v, 期望 %v", key, got, want)
}
}
// 策略范围
if _, ok := cfgMap["max_file_size"]; !ok {
t.Fatal("config 缺少 max_file_size(存储策略)")
}
if _, ok := cfgMap["max_save_seconds"]; !ok {
t.Fatal("config 缺少 max_save_seconds(保存时间策略)")
}
if got := cfgMap["max_save_count"]; got != float64(5) {
t.Fatalf("config.max_save_count = %v, 期望 5", got)
}
if _, ok := cfgMap["allowedFileTypes"]; !ok {
t.Fatal("config 缺少 allowedFileTypes")
}
if _, ok := cfgMap["expireStyle"]; !ok {
t.Fatal("config 缺少 expireStyle")
}
if _, ok := cfgMap["uploadSize"]; !ok {
t.Fatal("config 缺少 uploadSize")
}
// 敏感键绝不下发
raw := w.Body.String()
if strings.Contains(raw, "admin_token") || strings.Contains(raw, "jwt_secret") {
t.Fatal("公开 config 响应包含敏感键")
}
}
// ============ ② 管理端 get/updatev2 键全链路 + 类型范围校验 ============
// TestAdminConfigV2RoundTrip 验证 v2 新键 update → get → public 的往返,
// 且 admin_token 屏蔽、jwt_secret 不下发。
func TestAdminConfigV2RoundTrip(t *testing.T) {
d := newPolicyTestDeps(t)
patch := map[string]any{
"background_url": "https://cdn.example.com/bg.png",
"footer_text": "页脚 HTML 片段",
"footer_beian": "京ICP备20240001号",
"notify_enabled": 0,
"max_save_count": 20,
"max_file_size": 5242880,
"max_save_seconds": 86400,
}
w := patchConfig(d, patch)
if w.Code != 200 {
t.Fatalf("update 失败: %d %s", w.Code, w.Body.String())
}
// admin config get:新键可见 + 敏感键屏蔽
w = getConfig(d)
code, data := respBody(t, w)
if code != 200 {
t.Fatalf("get code=%d", code)
}
for key, want := range map[string]any{
"background_url": "https://cdn.example.com/bg.png",
"footer_text": "页脚 HTML 片段",
"footer_beian": "京ICP备20240001号",
"notify_enabled": float64(0),
"max_save_count": float64(20),
"max_file_size": float64(5242880),
"max_save_seconds": float64(86400),
} {
if got := data[key]; got != want {
t.Fatalf("admin get %s = %v, 期望 %v", key, got, want)
}
}
// v1 既有设计:admin_token 不在 configKeys 白名单(响应中不存在即屏蔽);
// 兼容两种形态:键缺失或空串均算通过
if got, present := data["admin_token"]; present && got != "" {
t.Fatalf("admin_token 应屏蔽(缺失或空串),实际 %v", got)
}
rawGet := w.Body.String()
if strings.Contains(rawGet, `"jwt_secret"`) {
t.Fatal("admin get 不应下发 jwt_secret")
}
// public config 立即反映(改策略 → 公开 config 即时更新)
w = getPublicConfig(d)
_, data = respBody(t, w)
cfgMap := data["config"].(map[string]any)
if got := cfgMap["max_file_size"]; got != float64(5242880) {
t.Fatalf("public max_file_size = %v, 期望 5242880", got)
}
if got := cfgMap["footer_beian"]; got != "京ICP备20240001号" {
t.Fatalf("public footer_beian = %v", got)
}
}
// TestAdminConfigV2Validation 验证新键的类型与范围校验(400 + 中文错误)。
func TestAdminConfigV2Validation(t *testing.T) {
d := newPolicyTestDeps(t)
cases := []struct {
name string
patch map[string]any
}{
{"max_file_size 负数", map[string]any{"max_file_size": -1}},
{"max_file_size 超上限", map[string]any{"max_file_size": config.MaxFileSizeMax + 1}},
{"max_save_count 超上限", map[string]any{"max_save_count": config.MaxSaveCountMax + 1}},
{"notify_enabled 越界", map[string]any{"notify_enabled": 2}},
{"footer_beian 超长", map[string]any{"footer_beian": strings.Repeat("备", config.FooterBeianMaxLen+1)}},
{"footer_text 超长", map[string]any{"footer_text": strings.Repeat("页", config.FooterTextMaxLen+1)}},
{"background_url 非法协议", map[string]any{"background_url": "javascript:alert(1)"}},
{"max_save_seconds 超上限", map[string]any{"max_save_seconds": config.MaxSaveSecondsMax + 1}},
{"allowed_file_types 类型错误", map[string]any{"allowed_file_types": 123}},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
w := patchConfig(d, tc.patch)
if w.Code != http.StatusBadRequest {
t.Fatalf("应 400,实际 %d %s", w.Code, w.Body.String())
}
code, _ := respBody(t, w)
if code != http.StatusBadRequest {
t.Fatalf("响应 code 应为 400,实际 %d", code)
}
})
}
// 合法值不受影响
if w := patchConfig(d, map[string]any{
"max_save_count": 0, // 0 = 不限制
"notify_enabled": 1,
"background_url": "data:image/png;base64,AAA",
"max_file_size": 1024,
"max_save_seconds": 0,
}); w.Code != 200 {
t.Fatalf("合法 patch 应 200: %d %s", w.Code, w.Body.String())
}
}
// ============ ③ 上传动态校验:admin 改策略 → 上传行为即时变化 ============
// TestUploadPolicyDynamicEnforcement 全链路:默认可传 → 改 max_file_size/白名单/
// 保存策略后 → 公开 config 反映 → 上传被新策略拒绝(403/400)。
func TestUploadPolicyDynamicEnforcement(t *testing.T) {
d := newPolicyTestDeps(t)
// 默认策略:小 PNG 上传成功
w := uploadFile(d, "ok.png", pngMagic, map[string]string{"expire_value": "1", "expire_style": "day"})
if code, _ := respBody(t, w); code != 200 {
t.Fatalf("默认策略上传应成功: %d %s", w.Code, w.Body.String())
}
// —— 大小上限:max_file_size=100 → 200B 文件 403 ——
if w = patchConfig(d, map[string]any{"max_file_size": 100}); w.Code != 200 {
t.Fatalf("patch max_file_size: %d %s", w.Code, w.Body.String())
}
w = getPublicConfig(d)
_, data := respBody(t, w)
if got := data["config"].(map[string]any)["max_file_size"]; got != float64(100) {
t.Fatalf("公开 config 未即时反映 max_file_size=100: %v", got)
}
w = uploadFile(d, "big.png", append(pngMagic, bytes.Repeat([]byte{0}, 200)...),
map[string]string{"expire_value": "1", "expire_style": "day"})
code, _ := respBody(t, w)
if code != http.StatusForbidden {
t.Fatalf("超限上传应 403: %d %s", w.Code, w.Body.String())
}
// —— 类型白名单:allowed_file_types=[.png] → .txt 403 ——
if w = patchConfig(d, map[string]any{"allowed_file_types": []string{".png"}}); w.Code != 200 {
t.Fatalf("patch allowed_file_types: %d %s", w.Code, w.Body.String())
}
w = uploadFile(d, "note.txt", []byte("hello"), map[string]string{"expire_value": "1", "expire_style": "day"})
if code, _ := respBody(t, w); code != http.StatusForbidden {
t.Fatalf("非白名单类型应 403: %d %s", w.Code, w.Body.String())
}
// —— 保存时间:max_save_seconds=3600 → expire 1 天 403 ——
if w = patchConfig(d, map[string]any{"max_save_seconds": 3600}); w.Code != 200 {
t.Fatalf("patch max_save_seconds: %d %s", w.Code, w.Body.String())
}
w = uploadFile(d, "timed.png", pngMagic, map[string]string{"expire_value": "1", "expire_style": "day"})
if code, _ := respBody(t, w); code != http.StatusForbidden {
t.Fatalf("保存时间超范围应 403: %d %s", w.Code, w.Body.String())
}
// —— 保存次数:max_save_count=5 → count=10 403(重置时间策略避免交叉影响)——
if w = patchConfig(d, map[string]any{"max_save_count": 5, "max_save_seconds": 0}); w.Code != 200 {
t.Fatalf("patch max_save_count: %d %s", w.Code, w.Body.String())
}
w = uploadFile(d, "counted.png", pngMagic, map[string]string{"expire_value": "10", "expire_style": "count"})
if code, _ := respBody(t, w); code != http.StatusForbidden {
t.Fatalf("保存次数超上限应 403: %d %s", w.Code, w.Body.String())
}
// 次数在上限内合法
w = uploadFile(d, "counted.png", pngMagic, map[string]string{"expire_value": "3", "expire_style": "count"})
if code, _ := respBody(t, w); code != 200 {
t.Fatalf("次数在上限内应成功: %d %s", w.Code, w.Body.String())
}
// —— 过期方式白名单收窄:expireStyle=[day] → hour 400 ——
if w = patchConfig(d, map[string]any{"expireStyle": []string{"day"}}); w.Code != 200 {
t.Fatalf("patch expireStyle: %d %s", w.Code, w.Body.String())
}
w = uploadFile(d, "hour.png", pngMagic, map[string]string{"expire_value": "2", "expire_style": "hour"})
if code, _ := respBody(t, w); code != http.StatusBadRequest {
t.Fatalf("非白名单 expire_style 应 400: %d %s", w.Code, w.Body.String())
}
// 恢复后可用(说明策略动态读取)
if w = patchConfig(d, map[string]any{"expireStyle": []string{"day", "hour", "minute", "forever", "count"}}); w.Code != 200 {
t.Fatal("恢复 expireStyle 失败")
}
w = uploadFile(d, "hour.png", pngMagic, map[string]string{"expire_value": "2", "expire_style": "hour"})
if code, _ := respBody(t, w); code != 200 {
t.Fatalf("白名单恢复后应成功: %d %s", w.Code, w.Body.String())
}
}
// TestChunkUploadPolicyEnforcement 验证分片上传链路接入动态策略。
func TestChunkUploadPolicyEnforcement(t *testing.T) {
d := newPolicyTestDeps(t)
// L4:后端强制 enableChunk 开关,本测试前置开启
if w := patchConfig(d, map[string]any{"enableChunk": 1}); w.Code != 200 {
t.Fatalf("patch enableChunk: %d %s", w.Code, w.Body.String())
}
// 大小:max_file_size=1000 → file_size 5000 拒绝
if w := patchConfig(d, map[string]any{"max_file_size": 1000}); w.Code != 200 {
t.Fatalf("patch max_file_size: %d %s", w.Code, w.Body.String())
}
w := chunkInitJSON(d, `{"file_name":"a.png","file_size":5000,"chunk_size":1024,"file_hash":"h"}`)
if code, _ := respBody(t, w); code != http.StatusForbidden {
t.Fatalf("分片总大小超限应 403: %d %s", w.Code, w.Body.String())
}
// 类型:allowed_file_types=[.png] → b.txt 拒绝(max_file_size 重置为回落,隔离类型断言)
if w := patchConfig(d, map[string]any{"allowed_file_types": []string{".png"}, "max_file_size": 0}); w.Code != 200 {
t.Fatalf("patch allowed_file_types: %d %s", w.Code, w.Body.String())
}
w = chunkInitJSON(d, `{"file_name":"b.txt","file_size":500,"chunk_size":1024,"file_hash":"h"}`)
if code, _ := respBody(t, w); code != http.StatusForbidden {
t.Fatalf("分片文件类型非白名单应 403: %d %s", w.Code, w.Body.String())
}
// 白名单内 + 大小内 → 会话创建成功
w = chunkInitJSON(d, `{"file_name":"c.png","file_size":500,"chunk_size":1024,"file_hash":"h"}`)
if code, _ := respBody(t, w); code != 200 {
t.Fatalf("合法分片初始化应 200: %d %s", w.Code, w.Body.String())
}
}
// TestPresignPolicyEnforcement 验证预签名直传链路接入动态大小策略。
func TestPresignPolicyEnforcement(t *testing.T) {
d := newPolicyTestDeps(t)
if w := patchConfig(d, map[string]any{"max_file_size": 1000}); w.Code != 200 {
t.Fatalf("patch max_file_size: %d %s", w.Code, w.Body.String())
}
payload := `{"file_name":"a.png","file_size":5000,"expire_value":1,"expire_style":"day"}`
req := httptest.NewRequest(http.MethodPost, "/presign/upload/init", strings.NewReader(payload))
req.Header.Set("Content-Type", "application/json")
w := invoke(d.presignInit, req)
if code, _ := respBody(t, w); code != http.StatusForbidden {
t.Fatalf("预签名直传超限应 403: %d %s", w.Code, w.Body.String())
}
}
// TestSensitiveKeysNeverInPublicConfig 额外兜底:公开 config 任意策略下都无敏感键。
func TestSensitiveKeysNeverInPublicConfig(t *testing.T) {
d := newPolicyTestDeps(t)
// 写入敏感 KV(模拟已初始化实例),公开端点依旧不能带出
ctx := context.Background()
if err := d.Mgr.UpdateKV(ctx, map[string]any{
"jwt_secret": "super-secret-value",
"admin_token": settings.HashPassword("password-123456"),
}); err != nil {
t.Fatalf("UpdateKV: %v", err)
}
if err := d.Mgr.Reload(ctx); err != nil {
t.Fatalf("Reload: %v", err)
}
raw := getPublicConfig(d).Body.String()
if strings.Contains(raw, "super-secret-value") || strings.Contains(raw, "jwt_secret") {
t.Fatal("公开 config 泄露 jwt_secret")
}
if strings.Contains(raw, "admin_token") {
t.Fatal("公开 config 泄露 admin_token")
}
}
// TestPolicySnapshotMatchesConfig 验证策略快照与 config 一致(单一读取口径)。
func TestPolicySnapshotMatchesConfig(t *testing.T) {
d := newPolicyTestDeps(t)
if w := patchConfig(d, map[string]any{"max_file_size": 2048, "max_save_count": 9, "max_save_seconds": 7200}); w.Code != 200 {
t.Fatal("patch 失败")
}
pol := d.CurrentUploadPolicy()
if pol.MaxFileSize != 2048 || pol.MaxSaveCount != 9 || pol.MaxSaveSeconds != 7200 {
t.Fatalf("策略快照不一致: %+v", pol)
}
if err := pol.CheckSize(2048); err != nil {
t.Fatalf("边界值应放行: %v", err)
}
if err := pol.CheckSize(2049); err == nil {
t.Fatal("超限应拒绝")
}
// 0=回落 uploadSize
if w := patchConfig(d, map[string]any{"max_file_size": 0}); w.Code != 200 {
t.Fatal("patch 失败")
}
if got := d.CurrentUploadPolicy().MaxFileSize; got != d.Cfg.UploadSize() {
t.Fatalf("max_file_size=0 应回落 uploadSize: %d vs %d", got, d.Cfg.UploadSize())
}
}
// 编译期保证 fmt 被使用(测试辅助函数中错误路径占位)。
var _ = fmt.Sprintf
// ============ v3 存储引擎热切换 ============
// switchEngine 调用 POST /admin/storage/switch。
func switchEngine(d *Deps, engine string) *httptest.ResponseRecorder {
raw, _ := json.Marshal(map[string]any{"engine": engine})
req := httptest.NewRequest(http.MethodPost, "/admin/storage/switch", bytes.NewReader(raw))
req.Header.Set("Content-Type", "application/json")
return invoke(d.adminStorageSwitch, req)
}
// TestAdminStorageSwitchLocal 本地引擎切换(测试 build 只产 local,切 local 恒成功)。
func TestAdminStorageSwitchLocal(t *testing.T) {
d := newPolicyTestDeps(t)
w := switchEngine(d, "local")
if w.Code != 200 {
t.Fatalf("switch local 失败: %d %s", w.Code, w.Body.String())
}
// KV 持久化:admin get 可见
if got := getConfig(d); got.Code != 200 {
t.Fatal("get 失败")
}
// 非法引擎名 400
if w := switchEngine(d, "ftp"); w.Code != http.StatusBadRequest {
t.Fatalf("非法引擎应 400,实际 %d", w.Code)
}
}
// TestAdminConfigEngineSwitchFailure 测试环境下切换到不可用引擎保持原引擎(503)。
// 测试 Manager 的 build 返回 local;这里通过直接操作 Manager 验证 503 路径的响应格式。
func TestAdminConfigEngineSwitchFailure(t *testing.T) {
d := newPolicyTestDeps(t)
// 用一个恒失败的 Manager 替换(模拟 s3/webdav 健康检查不过)
d.Store = storage.NewManager("local", mustLocal(t), func(string) (storage.Storage, error) {
return nil, errors.New("连接失败")
})
w := switchEngine(d, "s3")
if w.Code != http.StatusServiceUnavailable {
t.Fatalf("不可用引擎应 503,实际 %d %s", w.Code, w.Body.String())
}
if !strings.Contains(w.Body.String(), "已保持原引擎") {
t.Fatal("错误信息应包含「已保持原引擎」")
}
// 失败后当前引擎不变
if d.Store.CurrentName() != "local" {
t.Fatalf("失败后应保持 local,实际 %s", d.Store.CurrentName())
}
}
// TestAdminConfigMaskedSecrets 敏感引擎凭据:get 掩码、update 空/掩码不落库。
func TestAdminConfigMaskedSecrets(t *testing.T) {
d := newPolicyTestDeps(t)
// 先写入真实凭据
if w := patchConfig(d, map[string]any{"webdav_password": "real-secret", "s3_secret_access_key": "sk-real"}); w.Code != 200 {
t.Fatalf("写凭据失败: %s", w.Body.String())
}
// get 应为掩码
_, data := respBody(t, getConfig(d))
if got := data["webdav_password"]; got != settings.SensitiveMaskValue {
t.Fatalf("webdav_password 应掩码,实际 %v", got)
}
if got := data["s3_secret_access_key"]; got != settings.SensitiveMaskValue {
t.Fatalf("s3_secret_access_key 应掩码,实际 %v", got)
}
// 提交掩码(模拟前端回显原样提交)→ 不应覆盖为掩码串
if w := patchConfig(d, map[string]any{"webdav_password": settings.SensitiveMaskValue}); w.Code != 200 {
t.Fatalf("掩码提交应 200: %s", w.Body.String())
}
// 提交空串 → 不修改
if w := patchConfig(d, map[string]any{"s3_secret_access_key": ""}); w.Code != 200 {
t.Fatalf("空串提交应 200: %s", w.Body.String())
}
// 公开 config 绝不含引擎凭据
_, pub := respBody(t, getPublicConfig(d))
rawPub, _ := json.Marshal(pub)
for _, sk := range []string{"webdav_password", "s3_secret_access_key", "aws_session_token", "jwt_secret"} {
if strings.Contains(string(rawPub), sk) {
t.Fatalf("公开 config 不应包含 %s", sk)
}
}
}
func mustLocal(t *testing.T) storage.Storage {
t.Helper()
s, err := storage.NewLocalStorage(t.TempDir())
if err != nil {
t.Fatal(err)
}
return s
}
+511
View File
@@ -0,0 +1,511 @@
package api
import (
"errors"
"fmt"
"net/http"
"time"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"filecodebox/internal/middleware"
"filecodebox/internal/model"
"filecodebox/internal/response"
"filecodebox/internal/storage"
)
// presignSessionExpires 预签名会话有效期(对齐参考 PRESIGN_SESSION_EXPIRES=900 秒)。
const presignSessionExpires = 900
// getValidSession 校验并返回预签名会话(对齐参考 _get_valid_session):
// 不存在 404、已过期删除后 404、mode 不符 400。
func (d *Deps) getValidSession(c *gin.Context, uploadID, expectedMode string) (*model.PresignUploadSession, error) {
ctx := c.Request.Context()
var session model.PresignUploadSession
err := d.DB.WithContext(ctx).
Where("upload_id = ?", uploadID).First(&session).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, errNotFound("上传会话不存在")
}
if err != nil {
return nil, errInternal("查询上传会话失败: " + err.Error())
}
if session.IsExpired(time.Now()) {
_ = d.DB.WithContext(ctx).Where("upload_id = ?", uploadID).
Delete(&model.PresignUploadSession{}).Error
releaseStorage(ctx, d.DB, "presign:"+uploadID)
return nil, errNotFound("上传会话已过期")
}
if expectedMode != "" && session.Mode != expectedMode {
return nil, errBadRequest("此会话不支持" + expectedMode + "模式")
}
return &session, nil
}
// ============ POST /presign/upload/init 初始化预签名上传 ============
// presignInitRequest init 请求体(对齐参考 PresignUploadInitRequest)。
type presignInitRequest struct {
FileName string `json:"file_name" form:"file_name"`
FileSize int64 `json:"file_size" form:"file_size"`
ExpireValue int `json:"expire_value" form:"expire_value"`
ExpireStyle string `json:"expire_style" form:"expire_style"`
Code string `json:"code" form:"code"` // v3.1:自定义提取码(init 时校验,完成时落库)
}
// presignInit 初始化预签名上传(对齐参考 presign_upload_init):
// 引擎支持直链(S3)返回 direct + 预签名 PUT URL;否则返回 proxy + 代理地址。
func (d *Deps) presignInit(c *gin.Context) {
if !d.requireShareLogin(c) {
return
}
if !requireUploadLimit(c, d.Limiter) {
return
}
var req presignInitRequest
if err := bindJSONOrForm(c, &req); err != nil {
respondError(c, err)
return
}
safeName := storage.SanitizeFileName(req.FileName)
if safeName == "" {
auditRecordFailed(c, d.AuditSvc, "文件名非法")
response.Fail(c, http.StatusBadRequest, "文件名非法")
return
}
// v3.1:自定义提取码提前校验(init 时快速失败;完成请求须再次携带)
if err := validatePickupCode(req.Code); err != nil {
auditRecordFailed(c, d.AuditSvc, "提取码非法")
respondError(c, err)
return
}
if err := validateFileMagic(d.Cfg, safeName, "", nil); err != nil {
auditUploadEntry(c, "", safeName, req.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "文件类型被拒绝")
respondError(c, err)
return
}
// v2 需求 ④⑩:动态策略校验(max_file_size0=回落 uploadSize
if err := d.CurrentUploadPolicy().CheckSize(req.FileSize); err != nil {
auditUploadEntry(c, "", safeName, req.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "文件大小超过限制")
respondError(c, err)
return
}
// M2:直传 confirm 的实际大小校验依赖真实对象,0/负值声明直接拒绝
if req.FileSize <= 0 {
auditRecordFailed(c, d.AuditSvc, "file_size 非法")
response.Fail(c, http.StatusBadRequest, "file_size 必须大于 0")
return
}
if req.ExpireValue <= 0 {
req.ExpireValue = 1
}
if req.ExpireStyle == "" {
req.ExpireStyle = "day"
}
if _, err := resolveExpire(d.Cfg, req.ExpireValue, req.ExpireStyle); err != nil {
auditUploadEntry(c, "", safeName, req.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "过期策略非法")
respondError(c, err)
return
}
ctx := c.Request.Context()
uploadID := uuidHex()
resToken := "presign:" + uploadID
if err := reserveStorage(ctx, d.DB, d.Cfg, resToken, req.FileSize, presignSessionExpires*time.Second); err != nil {
auditUploadEntry(c, "", safeName, req.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "容量预留失败")
respondError(c, err)
return
}
dirPath, _, _, _, savePath := buildSavePath(d.Cfg, safeName, uploadID)
mode := "proxy"
uploadURL := "/presign/upload/proxy/" + uploadID
putURL, err := d.Store.PresignPutURL(ctx, savePath, presignSessionExpires)
switch {
case err == nil:
mode = "direct"
uploadURL = putURL
case errors.Is(err, storage.ErrNotSupported):
// 引擎不支持直传:代理模式
default:
releaseStorage(ctx, d.DB, resToken)
auditUploadEntry(c, "", safeName, req.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "生成预签名失败")
respondError(c, mapStorageError(err))
return
}
session := model.PresignUploadSession{
UploadID: uploadID,
FileName: safeName,
FileSize: req.FileSize,
SavePath: savePath,
Mode: mode,
Engine: d.Store.CurrentName(), // v3:会话归属引擎
ExpireValue: req.ExpireValue,
ExpireStyle: req.ExpireStyle,
ExpiresAt: time.Now().Add(presignSessionExpires * time.Second),
}
if err := d.DB.WithContext(ctx).Create(&session).Error; err != nil {
releaseStorage(ctx, d.DB, resToken)
auditUploadEntry(c, "", safeName, req.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "会话创建失败")
respondError(c, errInternal("创建上传会话失败: "+err.Error()))
return
}
d.Limiter.Add(c, middleware.LimitUpload)
auditUploadEntry(c, uploadID, safeName, req.FileSize, 0)
auditRecordSuccess(c, d.AuditSvc)
detail := gin.H{
"upload_id": uploadID,
"upload_url": uploadURL,
"mode": mode,
"expires_in": presignSessionExpires,
"file_path": dirPath,
}
if mode == "proxy" {
detail["proxy_upload_url"] = uploadURL
detail["legacy_proxy_upload_url"] = "/api" + uploadURL
}
response.OK(c, detail)
}
// ============ PUT /presign/upload/proxy/{uploadID} 代理上传 ============
// presignProxy 代理模式上传(对齐参考 presign_upload_proxy):
// 服务器接收文件并转存到存储引擎,随后立即创建分享记录。
func (d *Deps) presignProxy(c *gin.Context) {
if !d.requireShareLogin(c) {
return
}
if !requireUploadLimit(c, d.Limiter) {
return
}
uploadID := c.Param("uploadID")
session, err := d.getValidSession(c, uploadID, "proxy")
if err != nil {
auditUploadEntry(c, uploadID, "", 0, 0)
auditRecordFailed(c, d.AuditSvc, err.Error())
respondError(c, err)
return
}
// v3.1:自定义提取码随代理上传表单携带(init 时已预校验)
if err := validatePickupCode(c.PostForm("code")); err != nil {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "提取码非法")
respondError(c, err)
return
}
ctx := c.Request.Context()
if err := reserveStorage(ctx, d.DB, d.Cfg, "presign:"+uploadID, session.FileSize, presignSessionExpires*time.Second); err != nil {
respondError(c, err)
return
}
fh, err := c.FormFile("file")
if err != nil {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "缺少 file 字段")
response.Fail(c, http.StatusBadRequest, "缺少上传文件 file 字段")
return
}
// 动态策略快照(与 share/chunk 上传路径一致,消除会话窗口内的策略滞后)
maxSize := d.CurrentUploadPolicy().MaxFileSize
if maxSize > 0 && fh.Size > maxSize {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "大小超过限制")
response.Fail(c, http.StatusForbidden, fmt.Sprintf("大小超过限制,最大为%s", humanSize(maxSize)))
return
}
// 文件大小与声明不符(±1KB 容差,对齐参考)
if abs64(fh.Size-session.FileSize) > 1024 {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "文件大小与声明不符")
response.Fail(c, http.StatusBadRequest, "文件大小与声明不符")
return
}
f, err := fh.Open()
if err == nil {
defer func() { _ = f.Close() }()
if err = validateFileMagic(d.Cfg, session.FileName, fh.Header.Get("Content-Type"), readMultipartHeader(f, 64)); err == nil {
// v3:落盘走会话归属引擎
var ps storage.Storage
ps, sErr := d.storeFor(session.Engine)
if sErr != nil {
err = sErr
} else if _, err = ps.SaveFile(ctx, f, session.SavePath); err != nil {
// 落盘失败,err 交给统一错误处理
}
}
}
if err != nil {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "文件保存失败")
if isStorageErr(err) {
respondError(c, mapStorageError(err))
} else {
respondError(c, errInternal("文件保存失败: "+err.Error()))
}
return
}
code, err := d.createRecordFromSession(c, session, c.PostForm("code"))
releaseStorage(ctx, d.DB, "presign:"+uploadID)
if err != nil {
if ps, sErr := d.storeFor(session.Engine); sErr == nil {
_ = ps.DeleteFile(ctx, session.SavePath) // v3:清理走归属引擎
}
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "创建分享失败")
respondError(c, err)
return
}
_ = d.DB.WithContext(ctx).Where("upload_id = ?", uploadID).
Delete(&model.PresignUploadSession{}).Error
d.Limiter.Add(c, middleware.LimitUpload)
auditUploadEntry(c, code, session.FileName, session.FileSize, fh.Size)
auditRecordSuccess(c, d.AuditSvc)
response.OK(c, gin.H{"code": code, "name": session.FileName})
}
// ============ POST /presign/upload/confirm/{uploadID} 直传确认 ============
// presignConfirm 直传确认(对齐参考 presign_upload_confirm):
// 客户端完成 S3 直传后调用,校验文件已存在并创建分享记录。
func (d *Deps) presignConfirm(c *gin.Context) {
if !d.requireShareLogin(c) {
return
}
if !requireUploadLimit(c, d.Limiter) {
return
}
uploadID := c.Param("uploadID")
session, err := d.getValidSession(c, uploadID, "direct")
if err != nil {
auditUploadEntry(c, uploadID, "", 0, 0)
auditRecordFailed(c, d.AuditSvc, err.Error())
respondError(c, err)
return
}
ctx := c.Request.Context()
if err := reserveStorage(ctx, d.DB, d.Cfg, "presign:"+uploadID, session.FileSize, presignSessionExpires*time.Second); err != nil {
respondError(c, err)
return
}
// v3:直传文件存在性按会话归属引擎检查(直传可能落在旧引擎)
psCheck, sErr := d.storeFor(session.Engine)
if sErr != nil {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "存储引擎不可用: "+sErr.Error())
respondError(c, mapStorageError(sErr))
return
}
// v3.1:自定义提取码随确认请求携带(query 或 JSON/form body,均可选)
customCode := c.Query("code")
if customCode == "" && c.Request.Body != nil && c.Request.ContentLength != 0 {
var fin struct {
Code string `json:"code" form:"code"`
}
if err := bindJSONOrForm(c, &fin); err != nil {
respondError(c, err)
return
}
customCode = fin.Code
}
exists, err := psCheck.FileExists(ctx, session.SavePath)
if err == nil && !exists {
err = errNotFound("文件未上传或上传失败")
}
if err != nil {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "文件未上传或上传失败")
respondError(c, err)
return
}
// M2 修复:直传内容不经过服务器,confirm 必须核实实际大小与内容类型。
meta, head, hErr := psCheck.HeadMeta(ctx, session.SavePath, 64)
if hErr != nil {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "文件信息读取失败")
respondError(c, mapStorageError(hErr))
return
}
if meta == nil {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "文件未上传或上传失败")
respondError(c, errNotFound("文件未上传或上传失败"))
return
}
// 大小上限:实际大小超过策略上限 → 删除对象并 403(防绕过 max_file_size / storageLimit
maxSize := d.CurrentUploadPolicy().MaxFileSize
if maxSize > 0 && meta.Size > maxSize {
_ = psCheck.DeleteFile(ctx, session.SavePath)
releaseStorage(ctx, d.DB, "presign:"+uploadID)
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "实际文件大小超过限制")
respondError(c, errForbidden(fmt.Sprintf("大小超过限制,最大为%s", humanSize(maxSize))))
return
}
// 大小与声明不符(±1KB 容差,对齐 proxy 模式):超差删除对象并 400
if abs64(meta.Size-session.FileSize) > 1024 {
_ = psCheck.DeleteFile(ctx, session.SavePath)
releaseStorage(ctx, d.DB, "presign:"+uploadID)
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "文件大小与声明不符")
respondError(c, errBadRequest("文件大小与声明不符"))
return
}
// 内容类型防伪造(对齐 proxy 模式 magic bytes 校验)
if err := validateFileMagic(d.Cfg, session.FileName, "", head); err != nil {
_ = psCheck.DeleteFile(ctx, session.SavePath)
releaseStorage(ctx, d.DB, "presign:"+uploadID)
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "文件内容校验失败")
respondError(c, err)
return
}
code, err := d.createRecordFromSession(c, session, customCode)
releaseStorage(ctx, d.DB, "presign:"+uploadID)
if err != nil {
if ps, sErr := d.storeFor(session.Engine); sErr == nil {
_ = ps.DeleteFile(ctx, session.SavePath) // v3:清理走归属引擎
}
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "创建分享失败")
respondError(c, err)
return
}
_ = d.DB.WithContext(ctx).Where("upload_id = ?", uploadID).
Delete(&model.PresignUploadSession{}).Error
d.Limiter.Add(c, middleware.LimitUpload)
auditUploadEntry(c, code, session.FileName, session.FileSize, session.FileSize)
auditRecordSuccess(c, d.AuditSvc)
response.OK(c, gin.H{"code": code, "name": session.FileName})
}
// createRecordFromSession 依据预签名会话创建分享记录(对齐参考 create_file_record)。
func (d *Deps) createRecordFromSession(c *gin.Context, session *model.PresignUploadSession, customCode string) (string, error) {
exp, err := resolveExpire(d.Cfg, session.ExpireValue, session.ExpireStyle)
if err != nil {
return "", err
}
// v3.1:完成请求的自定义提取码兜底校验(init 已验,防只发完成请求绕过)
if err := validatePickupCode(customCode); err != nil {
return "", err
}
ctx := c.Request.Context()
code, err := pickCustomCode(ctx, d.DB, d.Cfg, customCode)
if err != nil {
return "", err
}
dir, name := splitDirBase(session.SavePath)
ext := baseExt(name)
fc := model.FileCodes{
Code: code,
Prefix: trimExt(name),
Suffix: ext,
UUIDFileName: &name,
FilePath: &dir,
Size: session.FileSize,
ExpiredAt: exp.ExpiredAt,
ExpiredCount: exp.ExpiredCount,
UsedCount: exp.UsedCount,
Engine: session.Engine, // v3:归属引擎戳
}
if err := d.DB.WithContext(ctx).Create(&fc).Error; err != nil {
return "", mapCodeConflict(err) // v3.1:并发占用自定义码 → 友好 400
}
return code, nil
}
// ============ GET /presign/upload/status/{uploadID} ============
// presignStatus 查询预签名会话状态(对齐参考 presign_upload_status)。
func (d *Deps) presignStatus(c *gin.Context) {
if !d.requireShareLogin(c) {
return
}
uploadID := c.Param("uploadID")
ctx := c.Request.Context()
var session model.PresignUploadSession
if err := d.DB.WithContext(ctx).
Where("upload_id = ?", uploadID).First(&session).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.Fail(c, http.StatusNotFound, "上传会话不存在")
return
}
respondError(c, errInternal("查询上传会话失败: "+err.Error()))
return
}
response.OK(c, gin.H{
"upload_id": session.UploadID,
"file_name": session.FileName,
"file_size": session.FileSize,
"mode": session.Mode,
"created_at": session.CreatedAt.Format(time.RFC3339),
"expires_at": session.ExpiresAt.Format(time.RFC3339),
"is_expired": session.IsExpired(time.Now()),
})
}
// ============ DELETE /presign/upload/{uploadID} 取消会话 ============
// presignCancel 取消预签名上传会话(对齐参考 presign_upload_cancel):
// 直传模式尽力清理已直传的文件。
func (d *Deps) presignCancel(c *gin.Context) {
if !d.requireShareLogin(c) {
return
}
uploadID := c.Param("uploadID")
ctx := c.Request.Context()
session, err := d.getValidSession(c, uploadID, "")
if err != nil {
respondError(c, err)
return
}
if session.Mode == "direct" {
// v3:清理走会话归属引擎
if ps, sErr := d.storeFor(session.Engine); sErr == nil {
if exists, eErr := ps.FileExists(ctx, session.SavePath); eErr == nil && exists {
_ = ps.DeleteFile(ctx, session.SavePath)
}
}
}
if err := d.DB.WithContext(ctx).Where("upload_id = ?", uploadID).
Delete(&model.PresignUploadSession{}).Error; err != nil {
respondError(c, errInternal("取消上传会话失败: "+err.Error()))
return
}
releaseStorage(ctx, d.DB, "presign:"+uploadID)
response.OK(c, gin.H{"message": "上传会话已取消"})
}
// ============ 杂项 ============
// abs64 绝对值。
func abs64(n int64) int64 {
if n < 0 {
return -n
}
return n
}
// isStorageErr 判断是否为存储层哨兵错误(含 %w 包装)。
func isStorageErr(err error) bool {
return err != nil && (errors.Is(err, storage.ErrNotFound) ||
errors.Is(err, storage.ErrInvalidPath) ||
errors.Is(err, storage.ErrUnavailable) ||
errors.Is(err, storage.ErrNotSupported) ||
errors.Is(err, storage.ErrRangeNotSatisfiable) ||
errors.Is(err, storage.ErrHashMismatch))
}
+187
View File
@@ -0,0 +1,187 @@
package api
import (
"net/http"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"filecodebox/internal/audit"
"filecodebox/internal/config"
"filecodebox/internal/middleware"
"filecodebox/internal/response"
"filecodebox/internal/settings"
"filecodebox/internal/storage"
)
// Deps API 层共享依赖(main.go 装配后注入)。
type Deps struct {
DB *gorm.DB
Cfg *config.Config
Mgr *settings.Manager
AuditSvc *audit.Service
Limiter *middleware.RateLimiter
Store *storage.Manager // v3:可热切换引擎管理器(实现 Storage 接口)
Version string
}
// jwtSecret 当前 JWT 签名密钥(settings KV 运行时可变)。
func (d *Deps) jwtSecret() string { return d.Mgr.SecretProvider()() }
// Register 注册全部 API 路由与前端静态资源回退。
// 业务路由挂根路径(/share /chunk /presign /admin),与审计中间件
// DefaultClassifier 的路由模式一致(t1 冻结契约);公共接口保留
// /api/v1/health 与 /api/v1/config(对齐 t1 骨架)。
func Register(r *gin.Engine, d *Deps) {
// —— 公共接口 ——
r.GET("/api/v1/health", d.health)
r.GET("/api/v1/config", d.publicConfig)
// Info3robotsText 配置键此前无路由承接,补上(内容可在管理端自定义)
r.GET("/robots.txt", d.robotsText)
// —— 初始化向导(未初始化时唯一可用入口,GuardNotInitialized 白名单)——
registerSetup(r, d)
// —— 分享 ——
share := r.Group("/share")
{
share.POST("/text", d.shareText)
share.POST("/file", d.shareFile)
// metadata:每次访问即计数(RequireRateLimit=进入检查+完成计数)
share.GET("/metadata", d.Limiter.RequireRateLimit(middleware.LimitMeta), d.shareMetadata)
share.POST("/metadata", d.Limiter.RequireRateLimit(middleware.LimitMeta), d.shareMetadataPost)
share.GET("/select", d.shareSelect)
share.POST("/select", d.shareSelectPost)
share.GET("/download", d.shareDownload)
}
// —— 分片上传 ——
chunk := r.Group("/chunk")
{
chunk.POST("/upload/init", d.chunkInit)
// 主路径(参考语义):/chunk/upload/{uploadID}/{index}
// 扁平兼容:/chunk/upload + 表单/query 传 upload_id/chunk_index
chunk.POST("/upload/:uploadID/:chunkIndex", d.chunkUpload)
chunk.POST("/upload", d.chunkUploadFlat)
chunk.GET("/upload/status/:uploadID", d.chunkStatus)
chunk.POST("/upload/complete/:uploadID", d.chunkComplete)
chunk.DELETE("/upload/:uploadID", d.chunkCancel)
}
// —— 预签名直传 ——
presign := r.Group("/presign")
{
presign.POST("/upload/init", d.presignInit)
presign.PUT("/upload/proxy/:uploadID", d.presignProxy)
presign.POST("/upload/confirm/:uploadID", d.presignConfirm)
presign.GET("/upload/status/:uploadID", d.presignStatus)
presign.DELETE("/upload/:uploadID", d.presignCancel)
}
// —— 管理端(login 公开,其余需管理员 JWT)——
registerAdmin(r, d)
// —— 前端静态资源 + SPA 回退(须最后注册)——
registerWeb(r, d)
}
// health 健康检查(对齐 t1 骨架,保持 /api/v1/health 语义不变)。
func (d *Deps) health(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{
"code": 200,
"msg": "ok",
"data": gin.H{
"status": "ok",
"version": d.Version,
"storage": d.Cfg.Engine(),
"time": nowRFC3339(),
},
})
}
// robotsText 输出管理端可配置的 robots.txt 内容。
func (d *Deps) robotsText(c *gin.Context) {
c.Data(http.StatusOK, "text/plain; charset=utf-8", []byte(d.Cfg.GetString("robotsText")))
}
// publicConfig 公共配置(前端首页/上传页所需;v2 需求 ①②③④⑩ 扩展):
// - 展示字段:站点信息、Logo/favicon、背景图、页脚文案/备案号、通知;
// - 策略范围(上传页动态渲染):大小上限、类型白名单、过期方式、保存
// 时间/次数上限、上传频率(仅范围,不含内部实现键)。
//
// 敏感键(admin_token/jwt_secretsettings.SensitiveKeys)与本端点无关:
// 下发字段为白名单显式构造,任何敏感键均不会出现在响应中。
func (d *Deps) publicConfig(c *gin.Context) {
cfg := d.Cfg
policy := d.CurrentUploadPolicy()
uploadCount := cfg.GetInt("uploadCount")
uploadMinute := cfg.GetInt("uploadMinute")
// uploadSize 为参考语义的回落上限,单独下发供管理端联动展示
c.JSON(http.StatusOK, gin.H{
"code": 200,
"msg": "ok",
"data": gin.H{
"config": gin.H{
"name": cfg.SiteName(),
"description": cfg.GetString("description"),
"explain": cfg.GetString("page_explain"),
// 需求 ①:Logo/favicon/背景图
"logo_url": cfg.LogoURL(),
"favicon_url": cfg.FaviconURL(),
"background_url": cfg.BackgroundURL(),
// 需求 ②:页脚自定义内容与备案号
"footer_text": cfg.FooterText(),
"footer_beian": cfg.FooterBeian(),
// v3:当前存储引擎名(仅名称,任何引擎参数/凭据不下发)
"storage_engine": d.Store.CurrentName(),
"site_domain": d.Cfg.SiteDomain(),
// 需求 ③:系统通知(开关 + 内容,前台右上角悬浮窗)
// L7:读取侧再做一次白名单净化,覆盖历史存量与直改库的数据
"notify_enabled": boolToInt(cfg.NotifyEnabled()),
"notify_title": cfg.GetString("notify_title"),
"notify_content": settings.SanitizeInlineHTML(cfg.GetString("notify_content")),
// 策略范围(需求 ④⑩):上传页动态读取并在范围内选择
"uploadSize": cfg.UploadSize(),
"max_file_size": policy.MaxFileSize,
"maxFileSize": policy.MaxFileSize,
"allowedFileTypes": policy.AllowedTypes,
"expireStyle": policy.ExpireStyles,
"max_save_seconds": policy.MaxSaveSeconds,
"maxSaveSeconds": policy.MaxSaveSeconds,
"max_save_count": policy.MaxSaveCount,
"maxSaveCount": policy.MaxSaveCount,
"uploadCount": uploadCount,
"uploadMinute": uploadMinute,
"enableChunk": cfg.EnableChunk(),
"openUpload": cfg.OpenUpload(),
},
"meta": gin.H{
"version": d.Version,
"features": gin.H{
"chunkUpload": cfg.EnableChunk(),
"guestUpload": cfg.OpenUpload(),
},
},
},
})
}
// requireShareLogin 分享上传权限(对齐参考 share_required_login):
// openUpload 开启时游客可传;关闭时要求管理员 Bearer token403)。
func (d *Deps) requireShareLogin(c *gin.Context) bool {
if d.Cfg.OpenUpload() {
return true
}
header := c.GetHeader("Authorization")
const prefix = "Bearer "
if len(header) <= len(prefix) || header[:len(prefix)] != prefix {
response.Fail(c, http.StatusForbidden, "本站未开启游客上传,如需上传请先登录后台")
return false
}
token := header[len(prefix):]
if _, err := middleware.VerifyAdminToken(d.jwtSecret(), token); err != nil {
response.Fail(c, http.StatusForbidden, "本站未开启游客上传,如需上传请先登录后台")
return false
}
return true
}
+193
View File
@@ -0,0 +1,193 @@
// security_fixes_test.go — 安全审计修复项行为测试:
// L4 enableChunk 强制、M2 presign 大小/类型校验、L3 提码长度、M3 chunk_size 上限。
package api
import (
"bytes"
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/gin-gonic/gin"
"filecodebox/internal/model"
"filecodebox/internal/settings"
)
// postJSON 以 JSON body 调用 POST 端点。
func postJSON(d *Deps, path string, body any) *httptest.ResponseRecorder {
var reader *bytes.Reader
if body == nil {
reader = bytes.NewReader(nil)
} else {
raw, _ := json.Marshal(body)
reader = bytes.NewReader(raw)
}
req := httptest.NewRequest(http.MethodPost, path, reader)
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = req
c.Params = append(c.Params, gin.Param{Key: "uploadID", Value: req.URL.Path[len("/presign/upload/confirm/"):]})
d.presignConfirm(c)
return w
}
// sha256LegacyHash 构造旧版 sha256$salt$hash 格式(M1 迁移测试用)。
func sha256LegacyHash(password string) string {
salt := make([]byte, 16)
for i := range salt {
salt[i] = byte(i)
}
saltHex := hex.EncodeToString(salt)
sum := sha256.Sum256([]byte(saltHex + password))
return "sha256$" + saltHex + "$" + hex.EncodeToString(sum[:])
}
// TestChunkToggleEnforced L4enableChunk=0 时 /chunk 相关端点一律 403。
func TestChunkToggleEnforced(t *testing.T) {
d := newPolicyTestDeps(t)
// 默认 enableChunk=0
w := chunkInitJSON(d, `{"file_name":"a.png","file_size":500,"chunk_size":1024,"file_hash":"h"}`)
if code, _ := respBody(t, w); code != http.StatusForbidden {
t.Fatalf("enableChunk=0 时 init 应 403: %d %s", code, w.Body.String())
}
// 开启后放行
if w := patchConfig(d, map[string]any{"enableChunk": 1}); w.Code != 200 {
t.Fatalf("patch enableChunk: %d", w.Code)
}
w = chunkInitJSON(d, `{"file_name":"a.png","file_size":500,"chunk_size":1024,"file_hash":"h"}`)
if code, _ := respBody(t, w); code != 200 {
t.Fatalf("enableChunk=1 时 init 应 200: %d %s", code, w.Body.String())
}
}
// TestChunkSizeCap M3chunk_size 超过 32MB 上限时 400。
func TestChunkSizeCap(t *testing.T) {
d := newPolicyTestDeps(t)
if w := patchConfig(d, map[string]any{"enableChunk": 1}); w.Code != 200 {
t.Fatalf("patch enableChunk: %d", w.Code)
}
w := chunkInitJSON(d, `{"file_name":"a.bin","file_size":70000000000,"chunk_size":34000000,"file_hash":"h"}`)
if code, _ := respBody(t, w); code != http.StatusBadRequest {
t.Fatalf("chunk_size 超上限应 400: %d %s", code, w.Body.String())
}
}
// TestPickupCodeMinLen L34 位自定义码拒绝、5 位通过。
func TestPickupCodeMinLen(t *testing.T) {
if err := validatePickupCode("abcd"); err == nil {
t.Fatal("4 位码应被拒绝")
}
if err := validatePickupCode("abcde"); err != nil {
t.Fatalf("5 位码应通过: %v", err)
}
}
// TestPresignConfirmRejectsOversizeObject M2
// 直传会话 confirm 时,若对象实际大小超过策略上限,应删除对象并 403。
func TestPresignConfirmRejectsOversizeObject(t *testing.T) {
d := newPolicyTestDeps(t)
ctx := context.Background()
// 声明 10 字节、策略上限 100 → 实际 PUT 500 字节对象
if err := d.Mgr.UpdateKV(ctx, map[string]any{"max_file_size": 100}); err != nil {
t.Fatalf("UpdateKV: %v", err)
}
if err := d.Mgr.Reload(ctx); err != nil {
t.Fatalf("Reload: %v", err)
}
uploadID := "test-oversize-confirm"
savePath := "share/data/presign_test.bin"
if _, err := d.Store.SaveFile(ctx, bytes.NewReader(make([]byte, 500)), savePath); err != nil {
t.Fatalf("SaveFile: %v", err)
}
sess := model.PresignUploadSession{
UploadID: uploadID, FileName: "presign_test.bin", FileSize: 10,
SavePath: savePath, Mode: "direct",
ExpireValue: 1, ExpireStyle: "day",
CreatedAt: time.Now(), ExpiresAt: time.Now().Add(time.Hour), Engine: "local",
}
if err := d.DB.WithContext(ctx).Create(&sess).Error; err != nil {
t.Fatalf("create session: %v", err)
}
res := model.StorageReservation{Token: "presign:" + uploadID, Size: 10, ExpiresAt: time.Now().Add(time.Hour)}
if err := d.DB.WithContext(ctx).Create(&res).Error; err != nil {
t.Fatalf("create reservation: %v", err)
}
w := postJSON(d, "/presign/upload/confirm/"+uploadID, nil)
if w.Code != http.StatusForbidden {
t.Fatalf("超限对象 confirm 应 403: %d %s", w.Code, w.Body.String())
}
// 对象应被删除、预留应释放
if ok, _ := d.Store.FileExists(ctx, savePath); ok {
t.Fatal("超限对象应被服务端删除")
}
var cnt int64
_ = d.DB.WithContext(ctx).Model(&model.StorageReservation{}).Where("token = ?", res.Token).Count(&cnt).Error
if cnt != 0 {
t.Fatal("预留应被释放")
}
}
// TestPresignConfirmRejectsSizeMismatch M2:实际大小与声明差超过 ±1KB 时 400。
func TestPresignConfirmRejectsSizeMismatch(t *testing.T) {
d := newPolicyTestDeps(t)
ctx := context.Background()
uploadID := "test-mismatch-confirm"
savePath := "share/data/presign_mismatch.bin"
if _, err := d.Store.SaveFile(ctx, bytes.NewReader(make([]byte, 2048)), savePath); err != nil {
t.Fatalf("SaveFile: %v", err)
}
sess := model.PresignUploadSession{
UploadID: uploadID, FileName: "presign_mismatch.bin", FileSize: 10,
SavePath: savePath, Mode: "proxy", // proxy 模式同样走大小核对(多引擎一致)
ExpireValue: 1, ExpireStyle: "day",
CreatedAt: time.Now(), ExpiresAt: time.Now().Add(time.Hour), Engine: "local",
}
if err := d.DB.WithContext(ctx).Create(&sess).Error; err != nil {
t.Fatalf("create session: %v", err)
}
res := model.StorageReservation{Token: "presign:" + uploadID, Size: 10, ExpiresAt: time.Now().Add(time.Hour)}
if err := d.DB.WithContext(ctx).Create(&res).Error; err != nil {
t.Fatalf("create reservation: %v", err)
}
w := postJSON(d, "/presign/upload/confirm/"+uploadID, nil)
if w.Code != http.StatusBadRequest {
t.Fatalf("大小不符 confirm 应 400: %d %s", w.Code, w.Body.String())
}
}
// TestAdminPasswordAutoUpgrade M1:明文/旧哈希经 VerifyPassword 后 NeedsRehash 为真,
// bcrypt 哈希不再需要升级。
func TestAdminPasswordAutoUpgrade(t *testing.T) {
if !settings.NeedsRehash("FileCodeBox2023") {
t.Fatal("明文哈希需要升级")
}
legacy := sha256LegacyHash("pwd12345")
if !settings.NeedsRehash(legacy) {
t.Fatal("sha256 哈希需要升级")
}
if !settings.VerifyPassword("pwd12345", legacy) {
t.Fatal("旧 sha256 哈希兼容校验失败")
}
b := settings.HashPassword("pwd12345")
if settings.NeedsRehash(b) {
t.Fatal("bcrypt 哈希不需要升级")
}
if !settings.VerifyPassword("pwd12345", b) {
t.Fatal("bcrypt 校验失败")
}
if settings.VerifyPassword("wrong", b) {
t.Fatal("错误密码不应通过")
}
}
+327
View File
@@ -0,0 +1,327 @@
package api
import (
"net/http"
"strconv"
"strings"
"github.com/gin-gonic/gin"
"filecodebox/internal/response"
"filecodebox/internal/settings"
)
// fileSizeUnits 文件大小单位(对齐参考 FILE_SIZE_UNITS)。
var fileSizeUnits = map[string]int64{"KB": 1024, "MB": 1024 * 1024, "GB": 1024 * 1024 * 1024}
// saveTimeUnits 保存时间单位(秒)。
var saveTimeUnits = map[string]int64{"second": 1, "minute": 60, "hour": 3600, "day": 86400}
// expireStyleOptions 可用过期方式(用于 setup 表单校验)。
var expireStyleOptions = []string{"day", "hour", "minute", "forever", "count"}
// setupFormValue 取表单/JSON 字符串值。
func setupFormValue(data map[string]any, key, def string) string {
v, ok := data[key]
if !ok || v == nil {
return def
}
if s, ok := v.(string); ok {
return s
}
return strings.TrimSpace(strconv.FormatFloat(toAnyFloat(v), 'f', -1, 64))
}
func toAnyFloat(v any) float64 {
switch n := v.(type) {
case float64:
return n
case int:
return float64(n)
case int64:
return float64(n)
}
return 0
}
// registerSetup 注册初始化向导(未初始化时唯一可用入口,白名单 /setup)。
func registerSetup(r *gin.Engine, d *Deps) {
r.GET("/setup", func(c *gin.Context) {
if d.Mgr.IsInitialized() {
c.Redirect(http.StatusSeeOther, "/")
return
}
c.Data(http.StatusOK, "text/html; charset=utf-8", []byte(buildSetupPage("")))
})
r.POST("/setup", func(c *gin.Context) {
if d.Mgr.IsInitialized() {
c.Redirect(http.StatusSeeOther, "/")
return
}
d.setupSubmit(c)
})
}
// setupSubmit 处理初始化提交(对齐参考 setup_submit + parse_setup_options)。
func (d *Deps) setupSubmit(c *gin.Context) {
// 兼容 JSON 与表单
data := map[string]any{}
if strings.Contains(c.GetHeader("Content-Type"), "application/json") {
if err := c.ShouldBindJSON(&data); err != nil {
c.Data(http.StatusBadRequest, "text/html; charset=utf-8", []byte(buildSetupPage("请求体格式错误")))
return
}
} else if err := c.Request.ParseForm(); err != nil {
c.Data(http.StatusBadRequest, "text/html; charset=utf-8", []byte(buildSetupPage("请求体格式错误")))
return
} else {
// 多值字段(如多个 expireStyle 复选框)保留完整列表,单值取首项
for k, v := range c.Request.PostForm {
switch {
case len(v) == 1:
data[k] = v[0]
case len(v) > 1:
data[k] = v
}
}
}
adminPassword := setupFormValue(data, "admin_password", "")
confirmPassword := setupFormValue(data, "confirm_password", "")
siteName := setupFormValue(data, "site_name", "")
if adminPassword == "" || len(adminPassword) < 8 {
c.Data(http.StatusBadRequest, "text/html; charset=utf-8", []byte(buildSetupPage("管理员密码至少 8 位")))
return
}
if adminPassword != confirmPassword {
c.Data(http.StatusBadRequest, "text/html; charset=utf-8", []byte(buildSetupPage("两次输入的管理员密码不一致")))
return
}
patch, errMsg := parseSetupOptions(data)
if errMsg != "" {
c.Data(http.StatusBadRequest, "text/html; charset=utf-8", []byte(buildSetupPage(errMsg)))
return
}
patch["site_name"] = firstNonEmpty(siteName, "文件快传")
patch["admin_token"] = settings.HashPassword(adminPassword)
patch["jwt_secret"] = settings.GenerateJWTSecret()
ctx := c.Request.Context()
if err := d.Mgr.UpdateKV(ctx, patch); err != nil {
c.Data(http.StatusInternalServerError, "text/html; charset=utf-8", []byte(buildSetupPage("初始化失败: "+err.Error())))
return
}
if err := d.Mgr.Reload(ctx); err != nil {
c.Data(http.StatusInternalServerError, "text/html; charset=utf-8", []byte(buildSetupPage("配置重载失败: "+err.Error())))
return
}
d.syncRateRules()
// JSON 请求返回 JSON;表单返回成功页
if strings.Contains(c.GetHeader("Accept"), "application/json") ||
strings.Contains(c.GetHeader("Content-Type"), "application/json") {
response.OK(c, gin.H{"ok": true, "admin": "/#/admin"})
return
}
c.Data(http.StatusOK, "text/html; charset=utf-8", []byte(buildSetupSuccessPage()))
}
// parseSetupOptions 解析并校验初始化选项(对齐参考 parse_setup_options)。
func parseSetupOptions(data map[string]any) (map[string]any, string) {
out := map[string]any{}
// 文件大小限制
unit := strings.ToUpper(setupFormValue(data, "upload_size_unit", "MB"))
if _, ok := fileSizeUnits[unit]; !ok {
return nil, "文件大小单位不正确"
}
sizeVal, err := strconv.Atoi(setupFormValue(data, "upload_size_value", "10"))
if err != nil || sizeVal < 1 {
return nil, "文件大小限制必须是正整数"
}
out["uploadSize"] = int64(sizeVal) * fileSizeUnits[unit]
// 最长保存时间
saveUnit := strings.ToLower(setupFormValue(data, "save_time_unit", "day"))
if _, ok := saveTimeUnits[saveUnit]; !ok {
return nil, "最长保存时间单位不正确"
}
saveVal, err := strconv.Atoi(setupFormValue(data, "save_time_value", "0"))
if err != nil || saveVal < 0 {
return nil, "最长保存时间必须是非负整数"
}
out["max_save_seconds"] = int64(saveVal) * saveTimeUnits[saveUnit]
// 过期方式白名单
var styles []string
if raw, ok := data["expireStyle"]; ok {
switch v := raw.(type) {
case []any:
for _, item := range v {
if s, ok := item.(string); ok {
styles = append(styles, s)
}
}
case []string:
styles = v
case string:
for _, s := range strings.Split(v, ",") {
styles = append(styles, strings.TrimSpace(s))
}
}
}
valid := map[string]bool{}
var finalStyles []string
for _, s := range styles {
s = strings.TrimSpace(s)
if s == "" || valid[s] {
continue
}
for _, opt := range expireStyleOptions {
if opt == s {
valid[s] = true
finalStyles = append(finalStyles, s)
break
}
}
}
if len(finalStyles) == 0 {
return nil, "至少需要选择一种过期方式"
}
out["expireStyle"] = finalStyles
// 取件码类型
codeType := setupFormValue(data, "code_generate_type", "secret")
if codeType != "number" && codeType != "secret" {
return nil, "提取码类型不正确"
}
out["code_generate_type"] = codeType
// 频率限制
for _, item := range []struct{ key, def string }{
{"errorCount", "10"}, {"errorMinute", "1"},
{"loginCount", "5"}, {"loginMinute", "15"},
{"uploadCount", "10"}, {"uploadMinute", "1"},
} {
n, err := strconv.Atoi(setupFormValue(data, item.key, item.def))
if err != nil || n < 1 {
return nil, item.key + " 必须是正整数"
}
out[item.key] = n
}
// 布尔开关
out["openUpload"] = boolToInt(parseSetupBool(data, "openUpload", true))
out["enableChunk"] = boolToInt(parseSetupBool(data, "enableChunk", false))
// 允许文件类型
allowed := setupFormValue(data, "allowed_file_types", "*")
var types []string
for _, item := range strings.Split(allowed, ",") {
if item = strings.TrimSpace(item); item != "" {
types = append(types, item)
}
}
if len(types) == 0 {
types = []string{"*"}
}
out["allowed_file_types"] = types
return out, ""
}
// parseSetupBool 解析表单布尔(缺省 default"1"/"true"/"on"/"yes" 为真)。
func parseSetupBool(data map[string]any, key string, def bool) bool {
v, ok := data[key]
if !ok {
return def
}
switch s := v.(type) {
case string:
switch strings.ToLower(strings.TrimSpace(s)) {
case "1", "true", "on", "yes":
return true
case "0", "false", "off", "no", "":
return false
}
case float64:
return s != 0
case bool:
return s
}
return def
}
// buildSetupPage 初始化向导页面(简洁中文表单)。
func buildSetupPage(errMsg string) string {
errBlock := ""
if errMsg != "" {
errBlock = `<div style="margin-bottom:12px;padding:10px 12px;border-radius:10px;background:#fef2f2;color:#b91c1c;font-size:13px">` + htmlEscape(errMsg) + `</div>`
}
return `<!doctype html>
<html lang="zh-CN">
<head>
<meta charset="utf-8">
<meta name="viewport" content="width=device-width, initial-scale=1">
<title>初始化 文件快传</title>
<style>
body{margin:0;min-height:100vh;display:grid;place-items:center;padding:16px;font-family:-apple-system,"Segoe UI",sans-serif;background:#f5f5f7;color:#18181b}
main{width:min(100%,640px);padding:24px;border-radius:16px;background:#fff;box-shadow:0 18px 50px rgba(23,32,51,.08)}
h1{margin:0 0 6px;font-size:20px} p{margin:0 0 16px;color:#71717a;font-size:13px}
label{display:block;margin:10px 0 4px;font-size:12px;color:#3f3f46;font-weight:600}
input{width:100%;height:36px;border:1px solid #e4e4e7;border-radius:8px;padding:0 10px;box-sizing:border-box;font:inherit}
.grid{display:grid;grid-template-columns:1fr 1fr;gap:0 12px}
button{width:100%;height:40px;margin-top:16px;border:0;border-radius:10px;background:#18181b;color:#fff;font:inherit;font-weight:700;cursor:pointer}
</style>
</head>
<body><main>
<h1>初始化 文件快传</h1>
<p>首次配置管理员密码、上传限制和取件策略,后续可在后台调整。</p>
` + errBlock + `
<form method="post" action="/setup" autocomplete="off">
<label>站点名称</label>
<input name="site_name" maxlength="80" placeholder="文件快传">
<div class="grid">
<div><label>管理员密码</label><input name="admin_password" type="password" minlength="8" required></div>
<div><label>确认管理员密码</label><input name="confirm_password" type="password" minlength="8" required></div>
<div><label>单文件大小限制</label><input name="upload_size_value" type="number" min="1" value="10" required></div>
<div><label>大小单位</label><input name="upload_size_unit" value="MB" required></div>
<div><label>上传频率(次/分钟)</label><input name="uploadCount" type="number" min="1" value="10" required></div>
<div><label>上传检测窗口(分钟)</label><input name="uploadMinute" type="number" min="1" value="1" required></div>
<div><label>取件错误频率(次/分钟)</label><input name="errorCount" type="number" min="1" value="10" required></div>
<div><label>取件错误窗口(分钟)</label><input name="errorMinute" type="number" min="1" value="1" required></div>
<div><label>登录失败频率(次/分钟)</label><input name="loginCount" type="number" min="1" value="5" required></div>
<div><label>登录失败窗口(分钟)</label><input name="loginMinute" type="number" min="1" value="15" required></div>
<div><label>最长保存时间</label><input name="save_time_value" type="number" min="0" value="0" required></div>
<div><label>保存时间单位</label><input name="save_time_unit" value="day" required></div>
</div>
<label>允许文件类型(逗号分隔,* 不限制)</label>
<input name="allowed_file_types" value="*">
<label>提取码类型(number=数字 / secret=随机字符)</label>
<input name="code_generate_type" value="secret">
<label><input type="checkbox" name="openUpload" value="1" checked style="width:auto"> 允许游客上传</label>
<label><input type="checkbox" name="enableChunk" value="1" style="width:auto"> 启用切片上传</label>
<input type="hidden" name="expireStyle" value="day">
<input type="hidden" name="expireStyle" value="hour">
<input type="hidden" name="expireStyle" value="minute">
<input type="hidden" name="expireStyle" value="forever">
<input type="hidden" name="expireStyle" value="count">
<button type="submit">完成初始化</button>
</form>
</main></body></html>`
}
// buildSetupSuccessPage 初始化完成页。
func buildSetupSuccessPage() string {
return `<!doctype html>
<html lang="zh-CN"><head><meta charset="utf-8"><meta http-equiv="refresh" content="2;url=/#/admin"><title>初始化完成</title></head>
<body style="display:grid;place-items:center;min-height:100vh;font-family:-apple-system,sans-serif;background:#f6f8fb;color:#172033">
<main style="text-align:center;padding:32px;background:#fff;border-radius:12px;box-shadow:0 18px 50px rgba(23,32,51,.08)">
<h1>初始化完成</h1><p>管理员密码已设置,请使用刚才的密码登录后台。</p><a href="/#/admin">进入后台</a>
</main></body></html>`
}
// htmlEscape HTML 转义(错误信息拼接用)。
func htmlEscape(s string) string {
r := strings.NewReplacer("&", "&amp;", "<", "&lt;", ">", "&gt;", `"`, "&#39;")
return r.Replace(s)
}
+626
View File
@@ -0,0 +1,626 @@
package api
import (
"errors"
"fmt"
"io"
"net/http"
"strconv"
"strings"
"time"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"filecodebox/internal/audit"
"filecodebox/internal/middleware"
"filecodebox/internal/model"
"filecodebox/internal/response"
"filecodebox/internal/storage"
)
// nowRFC3339 当前时间的 RFC3339 表示。
func nowRFC3339() string { return time.Now().Format(time.RFC3339) }
// lookupByCode 按取件码查询分享记录(对齐参考 get_code_file_by_code):
// 不存在返回 "文件不存在"expired=true 时过期返回 "文件已过期"。
// 服务端兜底:历史「链接+提取码」复制格式会把「CODE CODE」整串当码传入
// (空格经 URL 编码进 query/path),取第一段有效码避免误报不存在。
func (d *Deps) lookupByCode(c *gin.Context, code string, checkExpired bool) (*model.FileCodes, error) {
code = strings.TrimSpace(code)
if fields := strings.Fields(code); len(fields) > 1 {
code = fields[0]
}
if code == "" {
return nil, errNotFound("文件不存在")
}
var fc model.FileCodes
err := d.DB.WithContext(c.Request.Context()).
Where("code = ?", code).First(&fc).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, errNotFound("文件不存在")
}
if err != nil {
return nil, errInternal("查询失败: " + err.Error())
}
if checkExpired && fc.Expired(time.Now()) {
return nil, errNotFound("文件已过期")
}
return &fc, nil
}
// consumeUsage 原子校验分享状态并记录一次实际领取(对齐参考 consume_file_usage):
// 仅当 expired_count>0(次数剩余)或 expired_count<0 且未到过期时间时扣减成功。
func (d *Deps) consumeUsage(c *gin.Context, fc *model.FileCodes) bool {
now := time.Now()
res := d.DB.WithContext(c.Request.Context()).
Model(&model.FileCodes{}).
Where("id = ?", fc.ID).
Where("expired_count > 0 OR (expired_count < 0 AND (expired_at IS NULL OR expired_at > ?))", now).
Updates(map[string]any{
"expired_count": gorm.Expr("CASE WHEN expired_count > 0 THEN expired_count - 1 ELSE expired_count END"),
"used_count": gorm.Expr("used_count + 1"),
})
return res.Error == nil && res.RowsAffected > 0
}
// fileSavePath 拼接分享记录的存储相对路径(file_path/uuid_file_name)。
func fileSavePath(fc *model.FileCodes) string {
dir := ""
if fc.FilePath != nil {
dir = strings.Trim(*fc.FilePath, "/")
}
name := ""
if fc.UUIDFileName != nil {
name = *fc.UUIDFileName
}
if dir == "" {
return name
}
return dir + "/" + name
}
// ============ POST /share/text 文本分享 ============
// shareText 创建文本分享(对齐参考 share_text)。
func (d *Deps) shareText(c *gin.Context) {
if !d.requireShareLogin(c) {
return
}
if !requireUploadLimit(c, d.Limiter) {
return
}
// v3.1 修复:JSON/表单/ultipart 统一绑定(form+json 双标签——此前仅 PostForm 时,
// JSON 提交会静默存成空文本并 200,取件页空白)。
var body struct {
Text string `json:"text" form:"text"`
ExpireValue int `json:"expire_value" form:"expire_value"`
ExpireStyle string `json:"expire_style" form:"expire_style"`
Code string `json:"code" form:"code"`
}
if err := bindJSONOrForm(c, &body); err != nil {
respondError(c, err)
return
}
text := body.Text
if strings.TrimSpace(text) == "" {
response.Fail(c, http.StatusBadRequest, "分享内容不能为空")
return
}
// M3:前置拒绝超大 body(配合全局 BodyLimit441KB 为 222KB 内容 + 表单/JSON 编码余量)
if c.Request.ContentLength > 441*1024 {
response.Fail(c, http.StatusForbidden, "内容过多,建议采用文件形式")
return
}
expireValue := body.ExpireValue
if expireValue == 0 {
expireValue = 1
}
expireStyle := body.ExpireStyle
if expireStyle == "" {
expireStyle = "day"
}
// v3.1:自定义提取码格式校验(4-8 位字母数字,空=随机)
if err := validatePickupCode(body.Code); err != nil {
respondError(c, err)
return
}
exp, err := resolveExpire(d.Cfg, expireValue, expireStyle)
if err != nil {
respondError(c, err)
return
}
textSize := int64(len([]byte(text)))
if textSize > 222*1024 {
response.Fail(c, http.StatusForbidden, "内容过多,建议采用文件形式")
return
}
ctx := c.Request.Context()
token := "text:" + uuidHex()
if err := reserveStorage(ctx, d.DB, d.Cfg, token, textSize, 300*time.Second); err != nil {
respondError(c, err)
return
}
code, err := pickCustomCode(ctx, d.DB, d.Cfg, body.Code)
if err == nil {
fc := model.FileCodes{
Code: code,
Text: &text,
Size: textSize,
Prefix: "Text",
ExpiredAt: exp.ExpiredAt,
ExpiredCount: exp.ExpiredCount,
UsedCount: exp.UsedCount,
Engine: d.Store.CurrentName(), // v3:归属引擎戳(文本也记录,保持一致性)
}
err = d.DB.WithContext(ctx).Create(&fc).Error
}
err = mapCodeConflict(err) // v3.1:自定义码唯一索引冲突 → 友好 400
releaseStorage(ctx, d.DB, token)
if err != nil {
auditRecordFailed(c, d.AuditSvc, "文本分享创建失败")
respondError(c, err)
return
}
d.Limiter.Add(c, middleware.LimitUpload) // 上传成功才计数
auditUploadEntry(c, code, "Text", textSize, textSize)
auditRecordSuccess(c, d.AuditSvc)
response.OK(c, gin.H{"code": code})
}
// ============ POST /share/file 文件分享 ============
// shareFile 上传文件并创建分享(对齐参考 share_file)。
func (d *Deps) shareFile(c *gin.Context) {
if !d.requireShareLogin(c) {
return
}
if !requireUploadLimit(c, d.Limiter) {
return
}
fh, err := c.FormFile("file")
if err != nil {
auditRecordFailed(c, d.AuditSvc, "缺少 file 字段")
response.Fail(c, http.StatusBadRequest, "缺少上传文件 file 字段")
return
}
origName := fh.Filename
// v2 需求 ④⑩:动态策略校验(max_file_size0=回落 uploadSize
if err := d.CurrentUploadPolicy().CheckSize(fh.Size); err != nil {
auditUploadEntry(c, "", origName, fh.Size, 0)
auditRecordFailed(c, d.AuditSvc, "大小超过限制")
respondError(c, err)
return
}
expireValue := formInt(c, "expire_value", 1)
expireStyle := c.DefaultPostForm("expire_style", "day")
exp, err := resolveExpire(d.Cfg, expireValue, expireStyle)
if err != nil {
auditUploadEntry(c, "", origName, fh.Size, 0)
auditRecordFailed(c, d.AuditSvc, "过期策略非法")
respondError(c, err)
return
}
// v3.1:自定义提取码(落盘前校验,失败快速返回不占容量预留)
if err := validatePickupCode(c.PostForm("code")); err != nil {
auditUploadEntry(c, "", origName, fh.Size, 0)
auditRecordFailed(c, d.AuditSvc, "提取码非法")
respondError(c, err)
return
}
// magic bytes 防伪造(读前 64 字节)
f, err := fh.Open()
if err != nil {
auditUploadEntry(c, "", origName, fh.Size, 0)
auditRecordFailed(c, d.AuditSvc, "文件读取失败")
response.Fail(c, http.StatusBadRequest, "文件读取失败")
return
}
defer func() { _ = f.Close() }()
if err := validateFileMagic(d.Cfg, origName, fh.Header.Get("Content-Type"), readMultipartHeader(f, 64)); err != nil {
auditUploadEntry(c, "", origName, fh.Size, 0)
auditRecordFailed(c, d.AuditSvc, "文件类型被拒绝")
respondError(c, err)
return
}
dirPath, prefix, suffix, cleanName, savePath := buildSavePath(d.Cfg, origName, uuidHex())
ctx := c.Request.Context()
resToken := "file:" + uuidHex()
if err := reserveStorage(ctx, d.DB, d.Cfg, resToken, fh.Size, time.Hour); err != nil {
auditUploadEntry(c, "", origName, fh.Size, 0)
auditRecordFailed(c, d.AuditSvc, "容量预留失败")
respondError(c, err)
return
}
code, err := pickCustomCode(ctx, d.DB, d.Cfg, c.PostForm("code"))
if err == nil {
if _, err = d.Store.SaveFile(ctx, f, savePath); err != nil {
// 保存失败:清理半写文件
_ = d.Store.DeleteFile(ctx, savePath)
}
}
if err == nil {
fc := model.FileCodes{
Code: code,
Prefix: prefix,
Suffix: suffix,
UUIDFileName: &cleanName,
FilePath: &dirPath,
Size: fh.Size,
ExpiredAt: exp.ExpiredAt,
ExpiredCount: exp.ExpiredCount,
UsedCount: exp.UsedCount,
Engine: d.Store.CurrentName(), // v3:归属引擎戳(下载按此取回)
}
if err = d.DB.WithContext(ctx).Create(&fc).Error; err != nil {
err = mapCodeConflict(err) // v3.1
// 记录创建失败:清理已落盘文件
_ = d.Store.DeleteFile(ctx, savePath)
}
} else {
// 保存失败:尽力清理半写文件
_ = d.Store.DeleteFile(ctx, savePath)
}
releaseStorage(ctx, d.DB, resToken)
if err != nil {
auditUploadEntry(c, "", origName, fh.Size, 0)
auditRecordFailed(c, d.AuditSvc, "文件保存失败")
if be, ok := err.(*apiError); ok && be.Status == http.StatusBadRequest {
respondError(c, err) // v3.1:提取码冲突等业务 400 原样透出
} else {
respondError(c, mapStorageError(err))
}
return
}
d.Limiter.Add(c, middleware.LimitUpload)
auditUploadEntry(c, code, origName, fh.Size, fh.Size)
auditRecordSuccess(c, d.AuditSvc)
response.OK(c, gin.H{"code": code, "name": origName})
}
// ============ GET/POST /share/metadata 分享元信息(不消耗次数)============
// shareMetadataGET 查询分享元信息(对齐参考 get_file_metadata / build_file_metadata)。
func (d *Deps) shareMetadata(c *gin.Context) {
d.metadataCommon(c, c.Query("code"))
}
// shareMetadataPost JSON 体查询分享元信息(对齐参考 post_file_metadata)。
func (d *Deps) shareMetadataPost(c *gin.Context) {
var body struct {
Code string `json:"code"`
}
if err := bindJSONOrForm(c, &body); err != nil {
auditRecordFailed(c, d.AuditSvc, "请求体格式错误")
respondError(c, err)
return
}
d.metadataCommon(c, body.Code)
}
// metadataCommon 元信息查询公共实现。
func (d *Deps) metadataCommon(c *gin.Context, code string) {
fc, err := d.lookupByCode(c, code, true)
if err != nil {
auditUploadEntry(c, strings.TrimSpace(code), "", 0, 0)
auditRecordFailed(c, d.AuditSvc, err.Error())
respondError(c, err)
return
}
auditUploadEntry(c, fc.Code, fc.Prefix+fc.Suffix, fc.Size, 0)
auditRecordSuccess(c, d.AuditSvc)
response.OK(c, buildFileMetadata(fc))
}
// buildFileMetadata 构造分享元信息(对齐参考 build_file_metadata,不暴露存储路径)。
func buildFileMetadata(fc *model.FileCodes) gin.H {
isText := fc.Text != nil
var remaining any
if fc.ExpiredCount > 0 {
remaining = fc.ExpiredCount
}
var expiredAt any
if fc.ExpiredAt != nil {
expiredAt = fc.ExpiredAt.Format(time.RFC3339)
}
return gin.H{
"code": fc.Code,
"name": fc.Prefix + fc.Suffix,
"size": fc.Size,
"type": map[bool]string{true: "text", false: "file"}[isText],
"is_text": isText,
"created_at": fc.CreatedAt.Format(time.RFC3339),
"expired_at": expiredAt,
"expires_at": expiredAt,
"expired_count": fc.ExpiredCount,
"used_count": fc.UsedCount,
"remaining_downloads": remaining,
}
}
// ============ GET /share/select 取件(消耗次数,流式下载)============
// shareSelect 取件:文本直接返回纯文本;文件流式返回(支持 Range,对齐参考 get_code_file)。
func (d *Deps) shareSelect(c *gin.Context) {
// error 类限流:进入即检查(对齐参考 Depends(ip_limit["error"])
if allowed, _ := d.Limiter.Check(c, middleware.LimitError); !allowed {
response.Fail(c, http.StatusLocked, "请求次数过多,请稍后再试")
return
}
code := strings.TrimSpace(c.Query("code"))
fc, err := d.lookupByCode(c, code, true)
if err != nil {
d.Limiter.Add(c, middleware.LimitError) // 取件失败计数
auditUploadEntry(c, code, "", 0, 0)
auditRecordFailed(c, d.AuditSvc, err.Error())
respondError(c, err)
return
}
if !d.consumeUsage(c, fc) {
auditUploadEntry(c, fc.Code, fc.Prefix+fc.Suffix, fc.Size, 0)
auditRecordFailed(c, d.AuditSvc, "文件已过期")
response.Fail(c, http.StatusNotFound, "文件已过期")
return
}
if fc.Text != nil {
// 文本分享:text/plain 响应(对齐参考 Response(content=text, media_type=text/plain)
name := fc.Prefix + suffixOrTxt(fc)
auditUploadEntry(c, fc.Code, name, fc.Size, int64(len(*fc.Text)))
auditRecordSuccess(c, d.AuditSvc)
c.Header("Content-Disposition", contentDisposition(name))
c.Data(http.StatusOK, "text/plain; charset=utf-8", []byte(*fc.Text))
return
}
d.serveFile(c, fc)
}
func suffixOrTxt(fc *model.FileCodes) string {
if fc.Suffix != "" {
return fc.Suffix
}
return ".txt"
}
// ============ POST /share/select 取件详情(JSON,对齐参考 select_file============
// shareSelectPost 返回分享详情 JSON(元信息+内容/下载地址,对齐参考 select_file)。
// 有次数限制的文件返回代理下载地址(消耗发生在 download 时),其余在本次消耗。
func (d *Deps) shareSelectPost(c *gin.Context) {
if allowed, _ := d.Limiter.Check(c, middleware.LimitError); !allowed {
response.Fail(c, http.StatusLocked, "请求次数过多,请稍后再试")
return
}
var body struct {
Code string `json:"code"`
}
if err := bindJSONOrForm(c, &body); err != nil {
respondError(c, err)
return
}
fc, err := d.lookupByCode(c, body.Code, true)
if err != nil {
d.Limiter.Add(c, middleware.LimitError)
respondError(c, err)
return
}
detail := buildFileMetadata(fc)
var downloadURL string
if fc.Text != nil {
detail["text"] = *fc.Text
detail["content"] = *fc.Text
} else if fc.ExpiredCount >= 0 {
// 有次数限制:必须经代理下载接口扣次数
downloadURL = d.proxyDownloadURL(fc.Code)
detail["text"] = downloadURL
} else {
// 时间型/永久:优先引擎直链(如 S3 预签名),不支持则代理
if url, err := d.Store.PresignGetURL(c.Request.Context(), fileSavePath(fc), 3600); err == nil {
downloadURL = url
} else {
downloadURL = d.proxyDownloadURL(fc.Code)
}
detail["text"] = downloadURL
}
detail["download_url"] = nil
if downloadURL != "" {
detail["download_url"] = downloadURL
}
// 仅当下载地址不是代理地址时在本次消耗次数(对齐参考 consumes_on_download 判定)
consumesOnDownload := strings.HasPrefix(downloadURL, "/share/download?")
if !consumesOnDownload {
if !d.consumeUsage(c, fc) {
response.Fail(c, http.StatusNotFound, "文件已过期")
return
}
for k, v := range buildFileMetadata(fc) {
detail[k] = v
}
}
response.OK(c, detail)
}
// proxyDownloadURL 生成代理下载地址(对齐参考 get_file_url)。
func (d *Deps) proxyDownloadURL(code string) string {
secret := d.jwtSecret()
return "/share/download?key=" + GetSelectToken(code, secret, 0) + "&code=" + code
}
// ============ GET /share/download 代理下载(token 鉴权,消耗次数)============
// shareDownload 代理下载(对齐参考 download_file):
// key 为 GetSelectToken 生成的窗口令牌,同时接受当前与上一窗口。
func (d *Deps) shareDownload(c *gin.Context) {
key := c.Query("key")
code := strings.TrimSpace(c.Query("code"))
secret := d.jwtSecret()
// L2:HMAC 令牌 + 常量时间比较
if key == "" || secret == "" || !VerifySelectToken(code, secret, key) {
d.Limiter.Add(c, middleware.LimitError)
auditUploadEntry(c, code, "", 0, 0)
auditRecordFailed(c, d.AuditSvc, "下载鉴权失败")
response.Fail(c, http.StatusForbidden, "下载鉴权失败")
return
}
fc, err := d.lookupByCode(c, code, true)
if err != nil {
auditUploadEntry(c, code, "", 0, 0)
auditRecordFailed(c, d.AuditSvc, err.Error())
respondError(c, err)
return
}
if !d.consumeUsage(c, fc) {
auditUploadEntry(c, fc.Code, fc.Prefix+fc.Suffix, fc.Size, 0)
auditRecordFailed(c, d.AuditSvc, "文件已过期")
response.Fail(c, http.StatusNotFound, "文件已过期")
return
}
if fc.Text != nil {
auditUploadEntry(c, fc.Code, fc.Prefix+fc.Suffix, fc.Size, int64(len(*fc.Text)))
auditRecordSuccess(c, d.AuditSvc)
response.OK(c, *fc.Text)
return
}
d.serveFile(c, fc)
}
// ============ 文件流式下载(含 Range============
// parseRangeHeader 解析 Range 头(仅 bytes 单区间,多区间/单位错误返回 nil 交由全量处理;
// 合法但越界由引擎返回 ErrRangeNotSatisfiable→416)。
func parseRangeHeader(header string, totalSize int64) *storage.Range {
header = strings.TrimSpace(header)
if header == "" || !strings.HasPrefix(header, "bytes=") {
return nil
}
spec := strings.TrimPrefix(header, "bytes=")
if strings.Contains(spec, ",") { // 多区间不支持,回退全量
return nil
}
dash := strings.Index(spec, "-")
if dash < 0 {
return nil
}
startStr, endStr := strings.TrimSpace(spec[:dash]), strings.TrimSpace(spec[dash+1:])
if startStr == "" {
// bytes=-N:后缀区间
n, err := strconv.ParseInt(endStr, 10, 64)
if err != nil || n <= 0 {
return nil
}
if totalSize > 0 && n >= totalSize {
n = totalSize
}
return &storage.Range{Start: totalSize - n, End: -1}
}
start, err := strconv.ParseInt(startStr, 10, 64)
if err != nil || start < 0 {
return nil
}
if endStr == "" {
return &storage.Range{Start: start, End: -1}
}
end, err := strconv.ParseInt(endStr, 10, 64)
if err != nil || end < start {
return nil
}
return &storage.Range{Start: start, End: end}
}
// serveFile 打开存储流并写响应(200 全量 / 206 区间,自动 Accept-Ranges/Content-Range)。
// 响应字节数由审计中间件包装的 Writer 自动统计。
func (d *Deps) serveFile(c *gin.Context, fc *model.FileCodes) {
ctx := c.Request.Context()
savePath := fileSavePath(fc)
name := fc.Prefix + fc.Suffix
// v3:按文件归属引擎取回(切换引擎后旧文件仍可下载);空戳=历史数据回落当前引擎
store, err := d.storeFor(fc.Engine)
if err != nil {
auditUploadEntry(c, fc.Code, name, fc.Size, 0)
auditRecordFailed(c, d.AuditSvc, "存储引擎不可用: "+err.Error())
respondError(c, mapStorageError(err))
return
}
// 先 Stat 拿总大小(用于审计与 Range 后缀解析)
var total int64 = -1
if meta, err := store.Stat(ctx, savePath); err == nil && meta != nil {
total = meta.Size
}
rng := parseRangeHeader(c.GetHeader("Range"), total)
dl, err := store.Open(ctx, savePath, rng)
if err != nil {
auditUploadEntry(c, fc.Code, name, fc.Size, 0)
auditRecordFailed(c, d.AuditSvc, "文件读取失败")
respondError(c, mapStorageError(err))
return
}
defer func() { _ = dl.Close() }()
if dl.Total >= 0 {
total = dl.Total
}
c.Header("Accept-Ranges", "bytes")
c.Header("Content-Disposition", contentDisposition(name))
c.Header("Content-Type", "application/octet-stream")
status := http.StatusOK
if rng != nil {
end := dl.End
if end < 0 && total >= 0 {
end = total - 1
}
if end < dl.Start {
auditUploadEntry(c, fc.Code, name, total, 0)
auditRecordFailed(c, d.AuditSvc, "请求范围超出文件大小")
c.Header("Content-Range", fmt.Sprintf("bytes */%d", total))
response.Fail(c, http.StatusRequestedRangeNotSatisfiable, "请求范围超出文件大小")
return
}
status = http.StatusPartialContent
c.Header("Content-Range", fmt.Sprintf("bytes %d-%d/%d", dl.Start, end, total))
c.Header("Content-Length", strconv.FormatInt(end-dl.Start+1, 10))
} else if total >= 0 {
c.Header("Content-Length", strconv.FormatInt(total, 10))
}
auditUploadEntry(c, fc.Code, name, total, 0)
c.Status(status)
n, _ := io.Copy(c.Writer, dl)
middleware.AuditSet(c, func(e *audit.Entry) {
if e.TransferredBytes == 0 {
e.TransferredBytes = n
}
if e.SizeBytes == 0 {
e.SizeBytes = total
}
})
auditRecordSuccess(c, d.AuditSvc)
}
// formInt 读取表单整数(缺省 default 值,非法值亦回退 default)。
func formInt(c *gin.Context, key string, def int) int {
raw := c.PostForm(key)
if raw == "" {
raw = c.Query(key)
}
if raw == "" {
return def
}
n, err := strconv.Atoi(raw)
if err != nil {
return def
}
return n
}
+170
View File
@@ -0,0 +1,170 @@
package api
// v3.1:自定义提取码与站点域名单元测试。
import (
"encoding/json"
"net/http"
"net/http/httptest"
"net/url"
"strings"
"testing"
"github.com/gin-gonic/gin"
)
// postForm 以 urlencoded 表单调用 handlerv3.1 测试辅助)。
func postForm(d *Deps, path string, fields map[string]string) *httptest.ResponseRecorder {
form := url.Values{}
for k, v := range fields {
form.Set(k, v)
}
req := httptest.NewRequest(http.MethodPost, path, strings.NewReader(form.Encode()))
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
var handler gin.HandlerFunc
switch path {
case "/share/text":
handler = d.shareText
default:
handler = func(c *gin.Context) { c.AbortWithStatus(http.StatusNotFound) }
}
return invoke(handler, req)
}
func TestValidatePickupCode(t *testing.T) {
// 合法:空(用随机码)
if err := validatePickupCode(""); err != nil {
t.Fatalf("空码应合法: %v", err)
}
// 合法:5-8 位字母数字(L3:最小长度由 4 提升至 5)
for _, c := range []string{"abcde", "AB123", "12345678", "a1B2c"} {
if err := validatePickupCode(c); err != nil {
t.Fatalf("合法码 %s 不应报错: %v", c, err)
}
}
// 非法:长度(4 位及以下不再允许)
for _, c := range []string{"abcd", "a1B2", "abc", "123456789"} {
if err := validatePickupCode(c); err == nil {
t.Fatalf("非法长度 %s 应报错", c)
}
}
// 非法:字符
for _, c := range []string{"ab c1", "提码", "ab-cd", "ab.cd", "ab+cd"} {
if err := validatePickupCode(c); err == nil {
t.Fatalf("非法字符 %s 应报错", c)
}
}
}
func TestNormalizeSiteDomain(t *testing.T) {
// 空 = 当前地址
d, err := normalizeSiteDomain("")
if err != nil {
t.Fatalf("unexpected err: %v", err)
}
if d != "" {
t.Fatalf("want empty, got %q", d)
}
// 完整 URL
d, err = normalizeSiteDomain("https://share.example.com")
if err != nil {
t.Fatalf("unexpected err: %v", err)
}
if d != "https://share.example.com" {
t.Fatalf("want https://share.example.com, got %q", d)
}
// 带端口 + 去尾斜杠
d, err = normalizeSiteDomain("http://192.168.1.5:8466/")
if err != nil {
t.Fatalf("unexpected err: %v", err)
}
if d != "http://192.168.1.5:8466" {
t.Fatalf("want http://192.168.1.5:8466, got %q", d)
}
// 裸主机自动补 http
d, err = normalizeSiteDomain("share.example.com")
if err != nil {
t.Fatalf("unexpected err: %v", err)
}
if d != "http://share.example.com" {
t.Fatalf("want http://share.example.com, got %q", d)
}
// 非法:路径 / 协议
for _, bad := range []string{"https://a.com/path", "ftp://a.com", "javascript:alert(1)"} {
if _, err := normalizeSiteDomain(bad); err == nil {
t.Fatalf("非法域名 %s 应报错", bad)
}
}
}
// TestShareTextTextPlainCompat 复刻真实浏览器请求形态:
// 旧前端 bundle 发 text/plain Content-Type + urlencoded bodyfetch 字符串 body 默认头)。
// 修复前:该形态被静默存成空文本(bug 1)或 400「分享内容不能为空」(bug 2)。
func TestShareTextTextPlainCompat(t *testing.T) {
d := newPolicyTestDeps(t)
body := strings.NewReader("text=111&expire_value=1&expire_style=day&code=")
req := httptest.NewRequest(http.MethodPost, "/share/text", body)
req.Header.Set("Content-Type", "text/plain;charset=UTF-8")
w := invoke(d.shareText, req)
if w.Code != http.StatusOK {
t.Fatalf("text/plain+urlencoded 应 200: %d %s", w.Code, w.Body.String())
}
// JSON 体但 Content-Type 缺失/为 text/plain 也应可解析
req2 := httptest.NewRequest(http.MethodPost, "/share/text",
strings.NewReader(`{"text":"无头JSON","expire_value":1,"expire_style":"day"}`))
req2.Header.Set("Content-Type", "text/plain;charset=UTF-8")
w2 := invoke(d.shareText, req2)
if w2.Code != http.StatusOK {
t.Fatalf("text/plain+JSON体 应 200: %d %s", w2.Code, w2.Body.String())
}
// 取件确认内容真实落库
req3 := httptest.NewRequest(http.MethodPost, "/share/select",
strings.NewReader(`{"code":"`+codeOf(w)+`"}`))
req3.Header.Set("Content-Type", "application/json")
w3 := invoke(d.shareSelectPost, req3)
if !strings.Contains(w3.Body.String(), "111") {
t.Fatalf("落库内容应为 111: %s", w3.Body.String())
}
}
// codeOf 从创建响应提取取件码。
func codeOf(w *httptest.ResponseRecorder) string {
var env struct {
Data struct {
Code string `json:"code"`
} `json:"data"`
}
_ = json.Unmarshal(w.Body.Bytes(), &env)
return env.Data.Code
}
func TestShareTextCustomCode(t *testing.T) {
d := newPolicyTestDeps(t)
// 自定义码成功创建
w := postForm(d, "/share/text", map[string]string{"text": "自定义码测试", "code": "MYCODE1"})
if w.Code != http.StatusOK {
t.Fatalf("自定义码创建失败: %d %s", w.Code, w.Body.String())
}
// 重复占用 → 400
w = postForm(d, "/share/text", map[string]string{"text": "第二条", "code": "MYCODE1"})
if w.Code != http.StatusBadRequest {
t.Fatalf("占用码应 400: %d %s", w.Code, w.Body.String())
}
// 非法码 → 400
w = postForm(d, "/share/text", map[string]string{"text": "第三条", "code": "abc"})
if w.Code != http.StatusBadRequest {
t.Fatalf("过短码应 400: %d %s", w.Code, w.Body.String())
}
// 空码 → 随机码仍正常
w = postForm(d, "/share/text", map[string]string{"text": "第四条", "code": ""})
if w.Code != http.StatusOK {
t.Fatalf("空码应回退随机: %d %s", w.Code, w.Body.String())
}
}
+75
View File
@@ -0,0 +1,75 @@
package api
import (
"io/fs"
"net/http"
"path"
"strings"
"github.com/gin-gonic/gin"
"filecodebox/internal/response"
web "filecodebox/web"
)
// registerWeb 注册前端静态资源与 SPA 回退(必须最后注册):
// - 静态资源命中 web/dist 内文件则直接服务(带 Immutable 缓存,html 不缓存);
// - 未命中且为 GET/HEAD 且非 /api 前缀:回退 index.html(前端 history 路由
// /s/:code、/admin/*、/docs、/openapi 由 SPA 接管);
// - /api/* 未命中路由:JSON 404(避免调试时拿到 HTML 掩盖真实错误)。
func registerWeb(r *gin.Engine, d *Deps) {
dist, err := web.Dist()
if err != nil {
return // 嵌入异常时跳过(API 仍可用)
}
fileServer := http.StripPrefix("/", http.FileServer(http.FS(dist)))
indexHTML := readIndexHTML(dist)
r.NoRoute(func(c *gin.Context) {
p := c.Request.URL.Path
// 1. API 未命中:JSON 404
if strings.HasPrefix(p, "/api/") || p == "/api" {
response.Fail(c, http.StatusNotFound, "接口不存在")
return
}
// 2. 静态资源命中:直接服务
if p != "/" {
clean := strings.TrimPrefix(path.Clean(p), "/")
if clean != "" {
if f, err := dist.Open(clean); err == nil {
_ = f.Close()
fileServer.ServeHTTP(c.Writer, c.Request)
return
}
}
}
// 3. SPA 回退:仅 GET/HEAD 且接受 HTML
if c.Request.Method == http.MethodGet || c.Request.Method == http.MethodHead {
accept := c.GetHeader("Accept")
if accept == "" || strings.Contains(accept, "text/html") || strings.Contains(accept, "*/*") {
if indexHTML != nil {
c.Data(http.StatusOK, "text/html; charset=utf-8", indexHTML)
return
}
}
// 非 HTML 请求未命中:普通 404
c.Status(http.StatusNotFound)
return
}
c.Status(http.StatusNotFound)
})
}
// readIndexHTML 读取嵌入的 index.htmlSPA 回退用)。
func readIndexHTML(dist fs.FS) []byte {
f, err := dist.Open("index.html")
if err != nil {
return nil
}
defer func() { _ = f.Close() }()
data, err := fs.ReadFile(dist, "index.html")
if err != nil {
return nil
}
return data
}
+212
View File
@@ -0,0 +1,212 @@
// Package audit 提供上传/下载审计日志服务(需求 ③):
// 记录操作时间/IP/UA/设备解析/动作/结果/字节数/耗时,落库 Postgres。
package audit
import (
"context"
"errors"
"log"
"strings"
"sync"
"time"
"gorm.io/gorm"
"filecodebox/internal/model"
)
// Service 审计日志服务。
type Service struct {
sink Sink
}
// Sink 审计落库抽象(生产为 Postgres,测试为内存实现)。
type Sink interface {
// Save 批量落库。
Save(ctx context.Context, logs []model.AuditLog) error
}
// DBSink 基于 GORM 的落库实现。
type DBSink struct{ db *gorm.DB }
// NewDBSink 构造数据库落库实现。
func NewDBSink(db *gorm.DB) *DBSink { return &DBSink{db: db} }
// Save 批量插入审计记录。
func (s *DBSink) Save(ctx context.Context, logs []model.AuditLog) error {
if len(logs) == 0 {
return nil
}
return s.db.WithContext(ctx).CreateInBatches(&logs, 200).Error
}
// NewService 构造审计服务。
func NewService(sink Sink) *Service {
return &Service{sink: sink}
}
// Entry 一次待落库的审计事件。
type Entry struct {
Action string // upload | download
FileCode string // 取件码
FileName string // 原始文件名
SizeBytes int64 // 文件总字节数
TransferredBytes int64 // 实际传输字节数
IP string // 客户端 IP
UserAgent string // User-Agent
DeviceOS string // 操作系统
DeviceBrowser string // 浏览器
DeviceType string // desktop/mobile/tablet/bot/other
Actor string // admin | guest
Result string // success | denied | failed
ErrorMsg string // 失败原因
Duration time.Duration // 耗时
}
// Record 异步写入一条审计日志:先尝试同步落库,失败时进入内存缓冲等待重试,
// 避免审计失败影响主请求,也避免高峰期阻塞。
func (s *Service) Record(entry Entry) {
record := model.AuditLog{
Action: entry.Action,
FileCode: truncate(entry.FileCode, 64),
FileName: truncate(entry.FileName, 255),
SizeBytes: entry.SizeBytes,
TransferredBytes: entry.TransferredBytes,
IP: truncate(entry.IP, 64),
UserAgent: truncate(entry.UserAgent, 512),
DeviceOS: truncate(entry.DeviceOS, 64),
DeviceBrowser: truncate(entry.DeviceBrowser, 64),
DeviceType: truncate(entry.DeviceType, 32),
Actor: truncate(entry.Actor, 64),
Result: normalizeResult(entry.Result),
ErrorMsg: truncate(entry.ErrorMsg, 512),
DurationMs: entry.Duration.Milliseconds(),
}
go func() {
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
if err := s.sink.Save(ctx, []model.AuditLog{record}); err != nil {
log.Printf("[audit] 审计日志落库失败,进入重试队列: %v", err)
s.enqueue(record)
}
}()
}
// retryBuf 落库失败时的内存重试缓冲。
var retryBuf struct {
sync.Mutex
items []model.AuditLog
}
const maxRetryBuffer = 10000
// enqueue 入队;超出上限时丢弃最旧的,防止内存无限增长。
func (s *Service) enqueue(record model.AuditLog) {
retryBuf.Lock()
if len(retryBuf.items) >= maxRetryBuffer {
retryBuf.items = retryBuf.items[1:]
}
retryBuf.items = append(retryBuf.items, record)
retryBuf.Unlock()
}
// FlushRetry 将缓冲中的审计日志重新落库;由后台定时任务调用。
func (s *Service) FlushRetry() {
retryBuf.Lock()
if len(retryBuf.items) == 0 {
retryBuf.Unlock()
return
}
items := retryBuf.items
retryBuf.items = nil
retryBuf.Unlock()
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
if err := s.sink.Save(ctx, items); err != nil {
log.Printf("[audit] 重试队列落库失败: %v", err)
// 失败则放回队首
retryBuf.Lock()
retryBuf.items = append(items, retryBuf.items...)
if len(retryBuf.items) > maxRetryBuffer {
retryBuf.items = retryBuf.items[:maxRetryBuffer]
}
retryBuf.Unlock()
}
}
// StartRetryLoop 启动后台重试循环。
func (s *Service) StartRetryLoop(stop <-chan struct{}) {
go func() {
ticker := time.NewTicker(30 * time.Second)
defer ticker.Stop()
for {
select {
case <-stop:
s.FlushRetry()
return
case <-ticker.C:
s.FlushRetry()
}
}
}()
}
// Query 按条件分页查询审计日志(管理端使用)。
// action/ip/result 为可选过滤;begin/end 为创建时间范围(可选)。
func (s *Service) Query(page, pageSize int, action, ip, result string, begin, end *time.Time) ([]model.AuditLog, int64, error) {
dbSink, ok := s.sink.(*DBSink)
if !ok {
return nil, 0, errors.New("audit: 当前 sink 不支持查询")
}
if page < 1 {
page = 1
}
if pageSize < 1 || pageSize > 200 {
pageSize = 20
}
q := dbSink.db.Model(&model.AuditLog{})
if action != "" {
q = q.Where("action = ?", action)
}
if ip != "" {
q = q.Where("ip = ?", ip)
}
if result != "" {
q = q.Where("result = ?", result)
}
if begin != nil {
q = q.Where("created_at >= ?", *begin)
}
if end != nil {
q = q.Where("created_at <= ?", *end)
}
var total int64
if err := q.Count(&total).Error; err != nil {
return nil, 0, err
}
var logs []model.AuditLog
err := q.Order("id DESC").
Offset((page - 1) * pageSize).
Limit(pageSize).
Find(&logs).Error
return logs, total, err
}
func truncate(s string, n int) string {
if len(s) <= n {
return s
}
return s[:n]
}
func normalizeResult(r string) string {
switch strings.TrimSpace(r) {
case model.AuditResultSuccess, model.AuditResultDenied, model.AuditResultFailed:
return strings.TrimSpace(r)
case "":
return model.AuditResultFailed
default:
return model.AuditResultFailed
}
}
+84
View File
@@ -0,0 +1,84 @@
package audit
import (
"strings"
)
// DeviceInfo 从 User-Agent 解析出的设备信息。
type DeviceInfo struct {
OS string // Windows/macOS/Android/iOS/Linux/Unknown
Browser string // Chrome/Firefox/Safari/Edge/Other
Type string // desktop/mobile/tablet/bot/other
}
// 动作常量。
const (
ActionUpload = "upload"
ActionDownload = "download"
// ActionAdmin 管理端敏感操作(登录/登出/配置/密码/引擎切换/文件删除等),
// L5:纳入审计以便追溯登录失败与配置变更。
ActionAdmin = "admin"
)
// 角色常量。
const (
ActorAdmin = "admin"
ActorGuest = "guest"
)
// botKeywords 常见爬虫/机器人标识。
var botKeywords = []string{"bot", "spider", "crawl", "slurp", "curl/", "wget", "python-requests", "go-http-client"}
// ParseUserAgent 解析 User-Agent 为设备信息(轻量规则,避免引入重依赖)。
func ParseUserAgent(ua string) DeviceInfo {
ua = strings.TrimSpace(ua)
lower := strings.ToLower(ua)
if ua == "" {
return DeviceInfo{OS: "Unknown", Browser: "Other", Type: "other"}
}
for _, kw := range botKeywords {
if strings.Contains(lower, kw) {
return DeviceInfo{OS: "Unknown", Browser: "Other", Type: "bot"}
}
}
info := DeviceInfo{OS: "Unknown", Browser: "Other", Type: "desktop"}
// 操作系统
switch {
case strings.Contains(lower, "windows"):
info.OS = "Windows"
case strings.Contains(lower, "iphone"), strings.Contains(lower, "ipod"):
info.OS = "iOS"
info.Type = "mobile"
case strings.Contains(lower, "ipad"):
info.OS = "iOS"
info.Type = "tablet"
case strings.Contains(lower, "mac os x"), strings.Contains(lower, "macintosh"):
info.OS = "macOS"
case strings.Contains(lower, "android"):
info.OS = "Android"
info.Type = "mobile"
if strings.Contains(lower, "tablet") || !strings.Contains(lower, "mobile") {
info.Type = "tablet"
}
case strings.Contains(lower, "linux"), strings.Contains(lower, "ubuntu"), strings.Contains(lower, "fedora"):
info.OS = "Linux"
}
// 浏览器(顺序重要:Edge/OPR 必须在 Chrome 之前判断)
switch {
case strings.Contains(lower, "edg/"), strings.Contains(lower, "edge/"):
info.Browser = "Edge"
case strings.Contains(lower, "opr/"), strings.Contains(lower, "opera"):
info.Browser = "Opera"
case strings.Contains(lower, "chrome/"), strings.Contains(lower, "crios/"):
info.Browser = "Chrome"
case strings.Contains(lower, "firefox/"), strings.Contains(lower, "fxios/"):
info.Browser = "Firefox"
case strings.Contains(lower, "safari/"):
info.Browser = "Safari"
}
return info
}
+45
View File
@@ -0,0 +1,45 @@
package audit
import "testing"
func TestParseUserAgent(t *testing.T) {
cases := []struct {
ua string
os string
browser string
typ string
}{
{
ua: "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36",
os: "Windows", browser: "Chrome", typ: "desktop",
},
{
ua: "Mozilla/5.0 (iPhone; CPU iPhone OS 17_0 like Mac OS X) AppleWebKit/605.1.15 (KHTML, like Gecko) Version/17.0 Mobile/15E148 Safari/604.1",
os: "iOS", browser: "Safari", typ: "mobile",
},
{
ua: "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/119.0.0.0 Safari/537.36 Edg/119.0.0.0",
os: "macOS", browser: "Edge", typ: "desktop",
},
{
ua: "Mozilla/5.0 (Linux; Android 13; Pixel 7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Mobile Safari/537.36",
os: "Android", browser: "Chrome", typ: "mobile",
},
{
ua: "curl/8.4.0",
os: "Unknown", browser: "Other", typ: "bot",
},
{
ua: "Mozilla/5.0 (X11; Linux x86_64; rv:121.0) Gecko/20100101 Firefox/121.0",
os: "Linux", browser: "Firefox", typ: "desktop",
},
{ua: "", os: "Unknown", browser: "Other", typ: "other"},
}
for i, tc := range cases {
got := ParseUserAgent(tc.ua)
if got.OS != tc.os || got.Browser != tc.browser || got.Type != tc.typ {
t.Errorf("case %d: ParseUserAgent(%q) = %+v, want os=%s browser=%s type=%s",
i, tc.ua, got, tc.os, tc.browser, tc.typ)
}
}
}
+42
View File
@@ -0,0 +1,42 @@
// Package cache 提供统一缓存接口:FCB_REDIS_ADDR 未配置时自动降级为进程内存实现,
// 用于 IP 限流计数与热点配置缓存(需求 ② 的可选 Redis 增强)。
package cache
import (
"context"
"errors"
"time"
)
// ErrNotFound 表示键不存在。
var ErrNotFound = errors.New("cache: key 不存在")
// Cache 缓存统一接口。
type Cache interface {
// Get 读取字符串值;键不存在返回 ErrNotFound。
Get(ctx context.Context, key string) (string, error)
// Set 写入字符串值,ttl<=0 表示不过期。
Set(ctx context.Context, key, value string, ttl time.Duration) error
// Delete 删除键。
Delete(ctx context.Context, keys ...string) error
// Exists 判断键是否存在。
Exists(ctx context.Context, key string) (bool, error)
// Incr 原子自增;键不存在时从 0 开始并设置 ttl 窗口(限流固定窗口用)。
Incr(ctx context.Context, key string, ttl time.Duration) (int64, error)
// Close 释放底层资源(Redis 连接;内存实现为空操作)。
Close() error
}
// RedisOptions Redis 连接参数(addr 为空 → 内存实现;db 为 FCB_REDIS_DB 库号)。
type RedisOptions struct {
Addr string
DB int // 逻辑库号 0-15cluster 模式忽略)
}
// New 按配置构造缓存实现:redisAddr 为空 → 内存实现。
func New(ctx context.Context, opt RedisOptions) (Cache, error) {
if opt.Addr == "" {
return NewMemory(), nil
}
return NewRedis(ctx, opt.Addr, opt.DB)
}
+81
View File
@@ -0,0 +1,81 @@
package cache
import (
"context"
"sync"
"testing"
"time"
)
func TestMemoryCacheSetGet(t *testing.T) {
c := NewMemory()
defer c.Close()
ctx := context.Background()
if err := c.Set(ctx, "k1", "v1", 0); err != nil {
t.Fatalf("Set 失败: %v", err)
}
v, err := c.Get(ctx, "k1")
if err != nil || v != "v1" {
t.Fatalf("Get = (%q, %v)", v, err)
}
if _, err := c.Get(ctx, "missing"); err != ErrNotFound {
t.Fatalf("缺失键应返回 ErrNotFound: %v", err)
}
_ = c.Delete(ctx, "k1")
if _, err := c.Get(ctx, "k1"); err != ErrNotFound {
t.Fatal("删除后应不存在")
}
}
func TestMemoryCacheTTL(t *testing.T) {
c := NewMemory()
defer c.Close()
ctx := context.Background()
_ = c.Set(ctx, "ttl", "x", 50*time.Millisecond)
if ok, _ := c.Exists(ctx, "ttl"); !ok {
t.Fatal("TTL 内应存在")
}
time.Sleep(80 * time.Millisecond)
if _, err := c.Get(ctx, "ttl"); err != ErrNotFound {
t.Fatal("过期后应不存在")
}
}
func TestMemoryCacheIncrWindow(t *testing.T) {
c := NewMemory()
defer c.Close()
ctx := context.Background()
for i := int64(1); i <= 3; i++ {
n, err := c.Incr(ctx, "rl", time.Minute)
if err != nil || n != i {
t.Fatalf("Incr = (%d, %v), want (%d, nil)", n, err, i)
}
}
// 窗口过期后重新计数
_ = c.Set(ctx, "short", "seed", time.Millisecond)
time.Sleep(5 * time.Millisecond)
n, err := c.Incr(ctx, "short", time.Millisecond)
if err != nil || n != 1 {
t.Fatalf("过期窗口重置失败: (%d, %v)", n, err)
}
}
func TestMemoryCacheConcurrentIncr(t *testing.T) {
c := NewMemory()
defer c.Close()
ctx := context.Background()
var wg sync.WaitGroup
for i := 0; i < 50; i++ {
wg.Add(1)
go func() {
defer wg.Done()
_, _ = c.Incr(ctx, "cnt", time.Minute)
}()
}
wg.Wait()
n, _ := c.Incr(ctx, "cnt", time.Minute)
if n != 51 {
t.Fatalf("并发计数丢失: %d != 51", n)
}
}
+159
View File
@@ -0,0 +1,159 @@
package cache
import (
"context"
"sync"
"time"
)
// memoryItem 内存缓存条目。
type memoryItem struct {
value string
expiresAt time.Time // 零值表示不过期
}
// MemoryCache 进程内存缓存实现(单机、无持久化)。
type MemoryCache struct {
mu sync.RWMutex
items map[string]memoryItem
done chan struct{}
}
// NewMemory 构造内存缓存,并启动后台过期清理。
func NewMemory() *MemoryCache {
m := &MemoryCache{
items: make(map[string]memoryItem),
done: make(chan struct{}),
}
go m.gcLoop()
return m
}
// gcLoop 每分钟清理一次过期键,避免长期运行内存膨胀。
func (m *MemoryCache) gcLoop() {
ticker := time.NewTicker(time.Minute)
defer ticker.Stop()
for {
select {
case <-m.done:
return
case now := <-ticker.C:
m.mu.Lock()
for k, item := range m.items {
if !item.expiresAt.IsZero() && now.After(item.expiresAt) {
delete(m.items, k)
}
}
m.mu.Unlock()
}
}
}
// Get 读取键值。
func (m *MemoryCache) Get(_ context.Context, key string) (string, error) {
m.mu.RLock()
item, ok := m.items[key]
m.mu.RUnlock()
if !ok {
return "", ErrNotFound
}
if !item.expiresAt.IsZero() && time.Now().After(item.expiresAt) {
return "", ErrNotFound
}
return item.value, nil
}
// Set 写入键值。
func (m *MemoryCache) Set(_ context.Context, key, value string, ttl time.Duration) error {
item := memoryItem{value: value}
if ttl > 0 {
item.expiresAt = time.Now().Add(ttl)
}
m.mu.Lock()
m.items[key] = item
m.mu.Unlock()
return nil
}
// Delete 删除键。
func (m *MemoryCache) Delete(_ context.Context, keys ...string) error {
m.mu.Lock()
for _, k := range keys {
delete(m.items, k)
}
m.mu.Unlock()
return nil
}
// Exists 判断键是否存在。
func (m *MemoryCache) Exists(_ context.Context, key string) (bool, error) {
_, err := m.Get(context.Background(), key)
return err == nil, nil
}
// Incr 原子自增;首次创建时记录窗口起点(以过期时间体现)。
func (m *MemoryCache) Incr(_ context.Context, key string, ttl time.Duration) (int64, error) {
m.mu.Lock()
defer m.mu.Unlock()
now := time.Now()
item, ok := m.items[key]
if ok && !item.expiresAt.IsZero() && now.After(item.expiresAt) {
// 窗口已过期,重新计数
ok = false
}
var n int64
if !ok {
n = 1
newItem := memoryItem{value: "1"}
if ttl > 0 {
newItem.expiresAt = now.Add(ttl)
}
m.items[key] = newItem
return n, nil
}
// 解析现有值
for _, c := range item.value {
if c < '0' || c > '9' {
n = 0
break
}
n = n*10 + int64(c-'0')
}
n++
newItem := memoryItem{value: itoa(n), expiresAt: item.expiresAt}
m.items[key] = newItem
return n, nil
}
// Close 停止清理协程。
func (m *MemoryCache) Close() error {
select {
case <-m.done:
default:
close(m.done)
}
return nil
}
// itoa 简单整数转字符串,避免在锁内依赖 strconv 的额外开销(数值都很小)。
func itoa(n int64) string {
if n == 0 {
return "0"
}
var buf [20]byte
i := len(buf)
neg := n < 0
if neg {
n = -n
}
for n > 0 {
i--
buf[i] = byte('0' + n%10)
n /= 10
}
if neg {
i--
buf[i] = '-'
}
return string(buf[i:])
}
+117
View File
@@ -0,0 +1,117 @@
package cache
import (
"context"
"fmt"
"net/url"
"strings"
"time"
"github.com/redis/go-redis/v9"
)
// RedisCache 基于 Redis 的缓存实现(可选增强)。
type RedisCache struct {
client *redis.Client
}
// NewRedis 连接 Redis 并校验可用性。addr 支持两种形式:
// - host:port(纯地址,库号由 db 参数指定)
// - redis://[:password@]host:port[/db]URL 形式,URL 中的库号优先于 db 参数)
func NewRedis(ctx context.Context, addr string, db int) (*RedisCache, error) {
opts, err := buildRedisOptions(addr, db)
if err != nil {
return nil, err
}
client := redis.NewClient(opts)
if err := client.Ping(ctx).Err(); err != nil {
_ = client.Close()
return nil, fmt.Errorf("cache: Redis 连接失败 %s: %w", addr, err)
}
return &RedisCache{client: client}, nil
}
// buildRedisOptions 构造 go-redis 连接选项(纯地址 / URL 形式统一入口)。
func buildRedisOptions(addr string, db int) (*redis.Options, error) {
opts := &redis.Options{
Addr: addr,
DB: db,
DialTimeout: 5 * time.Second,
ReadTimeout: 3 * time.Second,
WriteTimeout: 3 * time.Second,
PoolSize: 32,
}
if strings.HasPrefix(addr, "redis://") || strings.HasPrefix(addr, "rediss://") {
u, err := redis.ParseURL(addr)
if err != nil {
return nil, fmt.Errorf("cache: Redis 地址解析失败 %s: %w", addr, err)
}
// URL 未显式携带库号(路径为空或 /)时用 db 参数;显式 /N 优先
if u.DB == 0 && !urlHasDBPath(addr) {
u.DB = db
}
u.DialTimeout = opts.DialTimeout
u.ReadTimeout = opts.ReadTimeout
u.WriteTimeout = opts.WriteTimeout
u.PoolSize = opts.PoolSize
opts = u
}
return opts, nil
}
// urlHasDBPath 判断 redis:// URL 是否显式携带了库号路径(如 /5)。
func urlHasDBPath(raw string) bool {
u, err := url.Parse(raw)
if err != nil {
return false
}
return strings.Trim(u.Path, "/") != ""
}
// Get 读取键值。
func (r *RedisCache) Get(ctx context.Context, key string) (string, error) {
val, err := r.client.Get(ctx, key).Result()
if err == redis.Nil {
return "", ErrNotFound
}
return val, err
}
// Set 写入键值。
func (r *RedisCache) Set(ctx context.Context, key, value string, ttl time.Duration) error {
return r.client.Set(ctx, key, value, ttl).Err()
}
// Delete 删除键。
func (r *RedisCache) Delete(ctx context.Context, keys ...string) error {
if len(keys) == 0 {
return nil
}
return r.client.Del(ctx, keys...).Err()
}
// Exists 判断键是否存在。
func (r *RedisCache) Exists(ctx context.Context, key string) (bool, error) {
n, err := r.client.Exists(ctx, key).Result()
return n > 0, err
}
// Incr 原子自增;首次创建时设置窗口 TTL。
// 使用 Lua 脚本保证 INCR+EXPIRE 原子性,避免多实例下窗口被反复重置。
func (r *RedisCache) Incr(ctx context.Context, key string, ttl time.Duration) (int64, error) {
var incrScript = redis.NewScript(`
local n = redis.call('INCR', KEYS[1])
if n == 1 and ARGV[1] ~= '0' then
redis.call('PEXPIRE', KEYS[1], ARGV[1])
end
return n
`)
ttlMs := int64(0)
if ttl > 0 {
ttlMs = ttl.Milliseconds()
}
return incrScript.Run(ctx, r.client, []string{key}, ttlMs).Int64()
}
// Close 关闭 Redis 连接。
func (r *RedisCache) Close() error { return r.client.Close() }
+70
View File
@@ -0,0 +1,70 @@
// redis_options_test.go — FCB_REDIS_DB / URL 库号解析单测。
package cache
import "testing"
func TestBuildRedisOptionsPlainAddr(t *testing.T) {
opts, err := buildRedisOptions("127.0.0.1:6379", 0)
if err != nil {
t.Fatalf("plain addr: %v", err)
}
if opts.DB != 0 {
t.Fatalf("默认库号应为 0, got %d", opts.DB)
}
opts, err = buildRedisOptions("127.0.0.1:6379", 5)
if err != nil {
t.Fatalf("plain addr db=5: %v", err)
}
if opts.Addr != "127.0.0.1:6379" || opts.DB != 5 {
t.Fatalf("host:port + db: got addr=%s db=%d", opts.Addr, opts.DB)
}
}
func TestBuildRedisOptionsURL(t *testing.T) {
cases := []struct {
name string
url string
dbParam int
wantDB int
wantPw string
}{
{"URL 无库号用参数", "redis://127.0.0.1:6379", 3, 3, ""},
{"URL 显式库号优先", "redis://127.0.0.1:6379/7", 3, 7, ""},
{"URL 带密码", "redis://:secretpw@127.0.0.1:6379/2", 0, 2, "secretpw"},
{"rediss 无库号用参数", "rediss://127.0.0.1:6379", 9, 9, ""},
{"URL 根路径视为无库号", "redis://127.0.0.1:6379/", 4, 4, ""},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
opts, err := buildRedisOptions(tc.url, tc.dbParam)
if err != nil {
t.Fatalf("buildRedisOptions(%q): %v", tc.url, err)
}
if opts.DB != tc.wantDB {
t.Fatalf("db = %d, want %d", opts.DB, tc.wantDB)
}
if opts.Password != tc.wantPw {
t.Fatalf("password = %q, want %q", opts.Password, tc.wantPw)
}
if opts.Addr != "127.0.0.1:6379" {
t.Fatalf("addr = %q", opts.Addr)
}
})
}
}
func TestBuildRedisOptionsInvalidURL(t *testing.T) {
if _, err := buildRedisOptions("redis://[bad", 0); err == nil {
t.Fatal("非法 URL 应报错")
}
}
func TestURLHasDBPath(t *testing.T) {
if urlHasDBPath("redis://h:6379") || urlHasDBPath("redis://h:6379/") {
t.Fatal("无路径或根路径应视为 false")
}
if !urlHasDBPath("redis://h:6379/5") {
t.Fatal("/5 应视为 true")
}
}
+417
View File
@@ -0,0 +1,417 @@
// Package config 提供全局配置:默认值对齐参考实现 core/settings.py
// 支持 FCB_* 环境变量覆盖默认值,再由数据库 settings KV 做运行时覆盖。
package config
import (
"fmt"
"os"
"strconv"
"strings"
)
// 会话有效期边界(与参考实现保持一致:天级、可配 1~365 天)。
const (
// AdminSessionExpireDefault 默认 7 天(L8:由 30 天缩短,降低 localStorage
// token 泄露后的暴露窗口;管理员可在 1~365 天内自行调整)
AdminSessionExpireDefault = 7 * 24 * 60 * 60
AdminSessionExpireMin = 24 * 60 * 60 // 最小 1 天
AdminSessionExpireMax = 365 * 24 * 60 * 60 // 最大 365 天
// DefaultSQLitePath SQLite 模式默认数据库文件路径(相对运行目录,自动创建 data/)。
DefaultSQLitePath = "./data/filecodebox.db"
)
// 数据库驱动常量(需求 ⑧:SQLite 默认、Postgres 可选)。
const (
DBDriverSQLite = "sqlite"
DBDriverPostgres = "postgres"
)
// DefaultLogoURL / DefaultFaviconURL 默认 Logo 与 favicon(需求 ⑤):
// v2 起默认改用前端打包的本地资源(web/src/assets/brand/logo.svg + favicon.png
// 经 Vite 产出 /assets/logo-*.svg 与 /assets/favicon-*.png)。此处留空,
// GET /api/v1/config 下发空值时前端 displayLogoUrl/displayFaviconUrl 回落到本地打包资源;
// 管理端仍可设置任意 URL 全站替换。
const DefaultLogoURL = ""
// DefaultFaviconURL favicon/备用 Logo 默认空串(语义见 DefaultLogoURL 注释)。
const DefaultFaviconURL = ""
// Config 运行时配置。Env 为 FCB_* 环境变量解析结果(进程级),
// KV 为数据库 settings 键值覆盖(可被管理端动态修改)。
type Config struct {
Env *EnvConfig
KV map[string]any
}
// EnvConfig 进程级环境变量配置,仅能通过环境变量修改。
type EnvConfig struct {
DBDriver string // FCB_DB_DRIVERsqlite|postgres,默认 sqlite(需求 ⑧)
DBDSN string // FCB_DB_DSNpostgres 必需;sqlite 为空时用 DefaultSQLitePath
RedisAddr string // FCB_REDIS_ADDR,可选;为空时缓存降级为内存实现
RedisDB int // FCB_REDIS_DBRedis 逻辑库号 0-15,默认 0(URL 形式地址以 URL 内库号优先)
Listen string // FCB_LISTEN,监听地址,默认 :8466
StorageEngine string // FCB_STORAGE_ENGINElocal|s3|webdav,默认 local
TrustedProxies []string // FCB_TRUSTED_PROXIES,逗号分隔的可信代理 CIDR
}
// defaults 返回与参考实现 core/settings.py DEFAULT_CONFIG 对齐的默认配置。
func defaults() map[string]any {
return map[string]any{
// 存储引擎与路径
"file_storage": "local",
"storage_path": "",
"storageLimit": 0,
// v3:存储引擎运行时可配(热切换);空=沿用 Env.StorageEngine 启动值
"storage_engine": "",
"site_domain": "",
// 站点信息
"name": "文件快传",
"site_name": "文件快传", // 新增:管理端可自定义
"description": "开箱即用的文件快传系统",
"notify_title": "系统通知",
"notify_content": "欢迎使用文件快传,拖拽或粘贴即可分享文本与文件。",
"page_explain": "请勿上传或分享违法内容。根据《中华人民共和国网络安全法》、《中华人民共和国刑法》、《中华人民共和国治安管理处罚法》等相关规定。 传播或存储违法、违规内容,会受到相关处罚,严重者将承担刑事责任。本站坚决配合相关部门,确保网络内容的安全,和谐,打造绿色网络环境。",
"keywords": "文件快传, 文件分享, 匿名口令分享文本, 文件",
// 需求 ⑤:默认 Logo 与 favicon(空 = 前端使用打包的本地资源)
"logo_url": DefaultLogoURL,
"favicon_url": DefaultFaviconURL,
// 需求 ①:背景图(v2 新增 background_urlbackground 为参考实现既有键,保留兼容)
"background": "",
"background_url": "",
// 需求 ②:页脚自定义内容与备案号
"footer_text": "",
"footer_beian": "",
// 需求 ③:系统通知(notify_enabled 新增开关,title/content 沿用参考语义)
"notify_enabled": 1,
// 需求 ④:保存策略(次数上限新增;时间上限沿用 max_save_seconds
"max_save_count": 0,
// 需求 ⑩:存储策略-单文件上限(0=回落 uploadSize,避免与参考键冲突)
"max_file_size": 0,
// 本地存储
"local_storage_path": "",
// S3 引擎
"s3_access_key_id": "",
"s3_secret_access_key": "",
"s3_bucket_name": "",
"s3_endpoint_url": "",
"s3_region_name": "auto",
"s3_signature_version": "s3v4",
"s3_hostname": "",
"s3_addressing_style": "auto",
"s3_proxy": 0,
"aws_session_token": "",
// WebDAV 引擎
"webdav_url": "",
"webdav_username": "",
"webdav_password": "",
"webdav_root_path": "filebox_storage",
"webdav_proxy": 0,
// 安全
"admin_token": "", // 管理员密码哈希;为空表示未初始化
"jwt_secret": "",
"adminSessionExpire": AdminSessionExpireDefault,
// 上传与分享策略
"openUpload": 1,
"uploadSize": 1024 * 1024 * 10,
"allowed_file_types": []string{"*"},
"expireStyle": []string{"day", "hour", "minute", "forever", "count"},
"code_generate_type": "secret",
"uploadMinute": 1,
"uploadCount": 10,
"errorMinute": 1,
"errorCount": 10,
"loginCount": 5,
"loginMinute": 15,
"max_save_seconds": 0,
"enableChunk": 0,
// 界面
"opacity": 0.9,
"showAdminAddr": 0,
"robotsText": "User-agent: *\nDisallow: /",
"serverWorkers": 1,
"serverHost": "0.0.0.0",
"serverPort": 8466,
}
}
// loadEnv 解析 FCB_* 环境变量;返回 nil 表示未设置任何必需项。
func loadEnv() (*EnvConfig, error) {
env := &EnvConfig{
DBDriver: strings.ToLower(strings.TrimSpace(os.Getenv("FCB_DB_DRIVER"))),
DBDSN: strings.TrimSpace(os.Getenv("FCB_DB_DSN")),
RedisAddr: strings.TrimSpace(os.Getenv("FCB_REDIS_ADDR")),
Listen: strings.TrimSpace(os.Getenv("FCB_LISTEN")),
StorageEngine: strings.TrimSpace(os.Getenv("FCB_STORAGE_ENGINE")),
}
// Redis 库号(FCB_REDIS_DB0-15;非法值忽略用默认 0)
if v := strings.TrimSpace(os.Getenv("FCB_REDIS_DB")); v != "" {
if n, err := strconv.Atoi(v); err == nil && n >= 0 && n <= 15 {
env.RedisDB = n
}
}
if env.Listen == "" {
env.Listen = ":8466"
}
if env.StorageEngine == "" {
env.StorageEngine = "local"
}
switch env.StorageEngine {
case "local", "s3", "webdav":
default:
return nil, fmt.Errorf("FCB_STORAGE_ENGINE 无效值 %q,仅支持 local|s3|webdav", env.StorageEngine)
}
if raw := strings.TrimSpace(os.Getenv("FCB_TRUSTED_PROXIES")); raw != "" {
for _, item := range strings.Split(raw, ",") {
if item = strings.TrimSpace(item); item != "" {
env.TrustedProxies = append(env.TrustedProxies, item)
}
}
}
return env, nil
}
// New 从环境变量构造配置;KV 覆盖先为空。
// 需求 ⑧:FCB_DB_DRIVER 默认 sqlite(零依赖);postgres 必须提供 FCB_DB_DSN。
func New() (*Config, error) {
env, err := loadEnv()
if err != nil {
return nil, err
}
switch env.DBDriver {
case "", DBDriverSQLite:
env.DBDriver = DBDriverSQLite
// sqlite 模式 DSN 可为空:数据库层回退到 DefaultSQLitePath
case DBDriverPostgres:
if env.DBDSN == "" {
return nil, fmt.Errorf("FCB_DB_DRIVER=postgres 时必须提供 FCB_DB_DSNPostgres 连接串)")
}
default:
return nil, fmt.Errorf("FCB_DB_DRIVER 无效值 %q,仅支持 sqlite|postgres", env.DBDriver)
}
return &Config{Env: env, KV: map[string]any{}}, nil
}
// ApplyKV 用数据库 settings KV 覆盖运行时配置(内部键以 _ 开头的不允许覆盖)。
func (c *Config) ApplyKV(kv map[string]any) {
for k, v := range kv {
if strings.HasPrefix(k, "_") {
continue
}
c.KV[k] = v
}
}
// Get 按 键读取:KV 覆盖 > 默认值;找不到返回零值与 false。
func (c *Config) Get(key string) (any, bool) {
if v, ok := c.KV[key]; ok {
return v, true
}
v, ok := defaults()[key]
return v, ok
}
// GetString 取字符串配置。
func (c *Config) GetString(key string) string {
v, ok := c.Get(key)
if !ok || v == nil {
return ""
}
if s, ok := v.(string); ok {
return s
}
return fmt.Sprintf("%v", v)
}
// GetInt 取整型配置,兼容 JSON 数字(float64)与字符串。
func (c *Config) GetInt(key string) int {
n, _ := c.getInt64(key)
return int(n)
}
// GetInt64 取长整型配置。
func (c *Config) GetInt64(key string) int64 {
n, _ := c.getInt64(key)
return n
}
func (c *Config) getInt64(key string) (int64, bool) {
v, ok := c.Get(key)
if !ok || v == nil {
return 0, false
}
switch n := v.(type) {
case int:
return int64(n), true
case int64:
return n, true
case float64:
return int64(n), true
case string:
if n, err := strconv.ParseInt(strings.TrimSpace(n), 10, 64); err == nil {
return n, true
}
}
return 0, false
}
// GetBool 取布尔配置,兼容 1/0、"true"/"false"/"on"/"yes"。
func (c *Config) GetBool(key string) bool {
v, ok := c.Get(key)
if !ok || v == nil {
return false
}
switch b := v.(type) {
case bool:
return b
case int:
return b != 0
case float64:
return b != 0
case string:
switch strings.ToLower(strings.TrimSpace(b)) {
case "1", "true", "on", "yes":
return true
}
}
return false
}
// GetStringSlice 取字符串切片配置。
// SiteDomain 站点对外域名(v3.1):空=分享链接用当前访问地址。
func (c *Config) SiteDomain() string {
return strings.TrimRight(strings.TrimSpace(c.GetString("site_domain")), "/")
}
func (c *Config) GetStringSlice(key string) []string {
v, ok := c.Get(key)
if !ok || v == nil {
return nil
}
switch s := v.(type) {
case []string:
return s
case []any:
out := make([]string, 0, len(s))
for _, item := range s {
if item == nil {
continue
}
out = append(out, fmt.Sprintf("%v", item))
}
return out
case string:
var out []string
for _, item := range strings.Split(s, ",") {
if item = strings.TrimSpace(item); item != "" {
out = append(out, item)
}
}
return out
}
return nil
}
// —— 常用字段的便捷访问(与参考 settings.xxx 对齐)——
// SiteName 站点名称。
func (c *Config) SiteName() string {
if v := c.GetString("site_name"); v != "" {
return v
}
return c.GetString("name")
}
// LogoURL 页面 Logo。
func (c *Config) LogoURL() string { return c.GetString("logo_url") }
// FaviconURL favicon 地址。
func (c *Config) FaviconURL() string { return c.GetString("favicon_url") }
// OpenUpload 是否允许游客上传。
func (c *Config) OpenUpload() bool { return c.GetBool("openUpload") }
// UploadSize 单文件大小上限(字节)。
func (c *Config) UploadSize() int64 { return c.GetInt64("uploadSize") }
// AllowedFileTypes 允许的文件类型列表("*" 表示不限制)。
func (c *Config) AllowedFileTypes() []string { return c.GetStringSlice("allowed_file_types") }
// ExpireStyle 允许的过期方式。
func (c *Config) ExpireStyle() []string { return c.GetStringSlice("expireStyle") }
// EnableChunk 是否启用分片上传。
func (c *Config) EnableChunk() bool { return c.GetBool("enableChunk") }
// MaxSaveSeconds 最长保存秒数,0 表示不限制。
func (c *Config) MaxSaveSeconds() int64 { return c.GetInt64("max_save_seconds") }
// MaxSaveCount 单次分享最大可取次数上限(需求 ④),0 表示不限制。
func (c *Config) MaxSaveCount() int { return c.GetInt("max_save_count") }
// MaxFileSize 存储策略-单文件上限(需求 ⑩);0 表示回落 uploadSize。
func (c *Config) MaxFileSize() int64 {
if n := c.GetInt64("max_file_size"); n > 0 {
return n
}
return c.UploadSize()
}
// FooterText 页脚自定义内容(需求 ②)。
func (c *Config) FooterText() string { return c.GetString("footer_text") }
// FooterBeian 备案号(需求 ②)。
func (c *Config) FooterBeian() string { return c.GetString("footer_beian") }
// BackgroundURL 背景图地址(需求 ①);空表示使用主题默认。
func (c *Config) BackgroundURL() string {
if v := c.GetString("background_url"); v != "" {
return v
}
return c.GetString("background")
}
// NotifyEnabled 系统通知开关(需求 ③):默认开启。
func (c *Config) NotifyEnabled() bool {
if v, ok := c.Get("notify_enabled"); ok && v != nil {
return c.GetBool("notify_enabled")
}
return true
}
// SQLitePath 数据库文件路径:sqlite 模式下 DSN 为空时回退默认路径(需求 ⑧)。
func (c *Config) SQLitePath() string {
if c.Env.DBDriver != DBDriverSQLite {
return ""
}
if c.Env.DBDSN != "" {
return c.Env.DBDSN
}
return DefaultSQLitePath
}
// AdminSessionExpireSeconds 管理员会话有效期(秒),
// 参考 apps/admin/dependencies.py 的 get_admin_session_expire_seconds。
func (c *Config) AdminSessionExpireSeconds() int {
n := c.GetInt("adminSessionExpire")
if n < AdminSessionExpireMin || n > AdminSessionExpireMax || n%AdminSessionExpireMin != 0 {
return AdminSessionExpireDefault
}
return n
}
// Engine 当前存储引擎。
// Engine 返回当前存储引擎名:KV storage_engine 优先(v3 运行时可改),
// 空(未设置/历史数据)回落启动值 Env.StorageEngineenv 校验过的 local|s3|webdav)。
// 枚举校验内联(避免 config→storage 反向依赖)。
func (c *Config) Engine() string {
if v, ok := c.Get(KeyStorageEngine); ok {
if s, isStr := v.(string); isStr {
switch s {
case "local", "s3", "webdav":
return s
}
}
}
return c.Env.StorageEngine
}
+160
View File
@@ -0,0 +1,160 @@
package config
import (
"testing"
)
func TestNewDefaultsToSQLite(t *testing.T) {
t.Setenv("FCB_DB_DRIVER", "")
t.Setenv("FCB_DB_DSN", "")
t.Setenv("FCB_REDIS_ADDR", "")
t.Setenv("FCB_LISTEN", "")
t.Setenv("FCB_STORAGE_ENGINE", "")
c, err := New()
if err != nil {
t.Fatalf("默认(无 DSN)应可构造: %v", err)
}
if c.Env.DBDriver != DBDriverSQLite {
t.Errorf("默认驱动应为 sqlite,实际 %s", c.Env.DBDriver)
}
if c.SQLitePath() != DefaultSQLitePath {
t.Errorf("SQLite 默认路径 = %s", c.SQLitePath())
}
if c.Env.Listen != ":8466" {
t.Errorf("默认监听地址错误: %s", c.Env.Listen)
}
if c.Env.StorageEngine != "local" {
t.Errorf("默认存储引擎错误: %s", c.Env.StorageEngine)
}
if c.Env.RedisAddr != "" {
t.Errorf("RedisAddr 应为空: %s", c.Env.RedisAddr)
}
}
func TestNewPostgresRequiresDSN(t *testing.T) {
t.Setenv("FCB_DB_DRIVER", "postgres")
t.Setenv("FCB_DB_DSN", "")
if _, err := New(); err == nil {
t.Fatal("postgres 模式缺少 FCB_DB_DSN 应报错")
}
t.Setenv("FCB_DB_DSN", "postgres://user:pass@localhost:5432/fcb")
c, err := New()
if err != nil {
t.Fatalf("postgres + DSN 应可构造: %v", err)
}
if c.Env.DBDriver != DBDriverPostgres {
t.Errorf("驱动应为 postgres,实际 %s", c.Env.DBDriver)
}
if c.SQLitePath() != "" {
t.Errorf("postgres 模式 SQLitePath 应为空: %s", c.SQLitePath())
}
}
func TestNewInvalidDriver(t *testing.T) {
t.Setenv("FCB_DB_DRIVER", "mysql")
t.Setenv("FCB_DB_DSN", "x")
if _, err := New(); err == nil {
t.Fatal("非法驱动应报错")
}
}
func TestNewInvalidEngine(t *testing.T) {
t.Setenv("FCB_DB_DSN", "postgres://x")
t.Setenv("FCB_STORAGE_ENGINE", "onedrive")
if _, err := New(); err == nil {
t.Fatal("非法引擎应报错")
}
}
func TestEnvOverridesAndDefaults(t *testing.T) {
t.Setenv("FCB_DB_DSN", "postgres://x")
t.Setenv("FCB_LISTEN", ":9999")
t.Setenv("FCB_STORAGE_ENGINE", "webdav")
c, err := New()
if err != nil {
t.Fatalf("New 失败: %v", err)
}
if c.Env.Listen != ":9999" || c.Env.StorageEngine != "webdav" {
t.Fatalf("env 覆盖失败: %+v", c.Env)
}
// 默认值对齐参考 DEFAULT_CONFIG
if got := c.GetInt("uploadSize"); got != 1024*1024*10 {
t.Errorf("uploadSize 默认值 = %d", got)
}
if got := c.GetInt("errorCount"); got != 10 {
t.Errorf("errorCount 默认值 = %d", got)
}
if got := c.GetInt("loginCount"); got != 5 {
t.Errorf("loginCount 默认值 = %d", got)
}
if got := c.GetInt("loginMinute"); got != 15 {
t.Errorf("loginMinute 默认值 = %d", got)
}
if got := c.GetBool("openUpload"); !got {
t.Error("openUpload 默认应为开启")
}
if c.EnableChunk() {
t.Error("enableChunk 默认应关闭")
}
// 新增字段(需求 ①)
if c.LogoURL() != DefaultLogoURL {
t.Errorf("logo_url 默认值 = %s", c.LogoURL())
}
if c.FaviconURL() != DefaultFaviconURL {
t.Errorf("favicon_url 默认值 = %s", c.FaviconURL())
}
if c.SiteName() == "" {
t.Error("site_name 默认值不应为空")
}
// 过期方式与文件类型
if len(c.ExpireStyle()) != 5 {
t.Errorf("expireStyle 默认值 = %v", c.ExpireStyle())
}
if len(c.AllowedFileTypes()) != 1 || c.AllowedFileTypes()[0] != "*" {
t.Errorf("allowed_file_types 默认值 = %v", c.AllowedFileTypes())
}
}
func TestKVOverridesEnvAndDefaults(t *testing.T) {
t.Setenv("FCB_DB_DSN", "postgres://x")
c, _ := New()
c.ApplyKV(map[string]any{
"uploadSize": 1024,
"openUpload": 0,
"site_name": "我的快递柜",
"logo_url": "https://example.com/logo.svg",
"internalKey": "x", // 非下划线开头允许;下划线开头被拒
"_secret": "no",
})
if got := c.GetInt("uploadSize"); got != 1024 {
t.Errorf("KV 覆盖 uploadSize 失败: %d", got)
}
if c.OpenUpload() {
t.Error("KV 覆盖 openUpload 失败")
}
if c.SiteName() != "我的快递柜" {
t.Errorf("site_name KV 覆盖失败: %s", c.SiteName())
}
if c.LogoURL() != "https://example.com/logo.svg" {
t.Errorf("logo_url KV 覆盖失败: %s", c.LogoURL())
}
if _, ok := c.Get("_secret"); ok {
t.Error("下划线内部键不应可通过 ApplyKV 覆盖")
}
}
func TestAdminSessionExpireClamp(t *testing.T) {
t.Setenv("FCB_DB_DSN", "postgres://x")
c, _ := New()
if got := c.AdminSessionExpireSeconds(); got != AdminSessionExpireDefault {
t.Errorf("默认会话有效期 = %d", got)
}
c.ApplyKV(map[string]any{"adminSessionExpire": 7 * 24 * 60 * 60})
if got := c.AdminSessionExpireSeconds(); got != 7*24*60*60 {
t.Errorf("7 天会话有效期 = %d", got)
}
c.ApplyKV(map[string]any{"adminSessionExpire": 3600}) // 非整天,回落默认
if got := c.AdminSessionExpireSeconds(); got != AdminSessionExpireDefault {
t.Errorf("非法值应回落默认 = %d", got)
}
}
+97
View File
@@ -0,0 +1,97 @@
// Package config — schema.go 定义 v2 新增配置键(KVschema
// 键名常量、类型、默认值与取值边界。管理与 API 层(t2)按下表读写与校验,
// 文档(t4)按本表生成说明。键名除参考实现既有 camelCase 键外,
// v2 新增键统一 snake_case。
package config
// —— v2 新增/沿用键名常量(单一事实来源;settings 包会 re-export)——
// 命名规则:v2 新增键 snake_case;与参考实现对齐的既有键保持原拼写。
const (
// 需求 ①:背景图
KeyBackground = "background" // 参考实现既有键(v1 兼容保留)
KeyBackgroundURL = "background_url" // v2 新增:背景图 URL 或上传后的访问地址(空=默认主题)
// 需求 ②:页脚
KeyFooterText = "footer_text" // v2 新增:页脚自定义内容(纯文本或受控 HTML 片段)
KeyFooterBeian = "footer_beian" // v2 新增:备案号(如 京ICP备2024xxxxxx号-1
// 需求 ③:系统通知
KeyNotifyEnabled = "notify_enabled" // v2 新增:通知开关,1 开启 / 0 关闭
KeyNotifyTitle = "notify_title" // 既有键:通知标题
KeyNotifyContent = "notify_content" // 既有键:通知内容(允许 <a> 等受控 HTML
// 需求 ④:保存策略(上传页动态读取并在范围内选择)
KeyMaxSaveSeconds = "max_save_seconds" // 既有键:最长保存秒数,0=不限制(仅受默认 7 天兜底)
KeyMaxSaveCount = "max_save_count" // v2 新增:单次分享最大可取(保存)次数上限,0=不限制
KeyExpireStyle = "expireStyle" // 既有键:允许的过期方式白名单
// 需求 ④:上传频率限制(既有键,对齐参考 ip_limit["upload"]
KeyUploadCount = "uploadCount" // 窗口内允许上传次数
KeyUploadMinute = "uploadMinute" // 频率窗口(分钟)
// 需求 ④⑩:存储策略(最大文件大小/允许类型/总容量)
KeyUploadSize = "uploadSize" // 既有键:单文件上限(字节),参考实现语义
KeyMaxFileSize = "max_file_size" // v2 新增:存储策略-单文件上限(字节),0=回落 uploadSize
KeyAllowedTypes = "allowed_file_types" // 既有键:允许类型白名单("*" 不限制)
KeyStorageLimit = "storageLimit" // 既有键:站点总容量(字节),0=不限制
KeyOpenUpload = "openUpload" // 既有键:游客上传开关
// v3:存储引擎运行时可配(热切换;file_storage 为参考既有键保留兼容)
KeyStorageEngine = "storage_engine" // 当前存储引擎:local|s3|webdav
KeySiteDomain = "site_domain" // 站点对外域名(空=分享链接用当前地址)
)
// —— 取值边界(管理端保存与 API 校验用)——
const (
// 保存时间上限:最长 365 天,0 表示不限制。
MaxSaveSecondsMax = 365 * 24 * 60 * 60
// 保存次数上限:最长 100000 次,0 表示不限制。
MaxSaveCountMax = 100000
// 单文件大小上限:最长 10 GiB0 表示回落 uploadSize。
MaxFileSizeMax = 10 * 1024 * 1024 * 1024
// 背景图 URL 最大长度(含 data: 之外的普通 http(s) URL)。
BackgroundURLMaxLen = 2048
// 页脚自定义内容最大长度。
FooterTextMaxLen = 2000
// 备案号最大长度。
FooterBeianMaxLen = 128
// 通知标题/内容最大长度。
NotifyTitleMaxLen = 128
NotifyContentMaxLen = 2000
)
// KVSchemaEntry 配置键元数据:类型 / 默认值 / 说明,供管理端 UI 与文档生成。
type KVSchemaEntry struct {
Key string // KV 键名
Type string // string | int | int64 | bool | []string
Default any // 默认值(与 defaults() 保持一致,测试保证同步)
Min int64 // 数值键最小值(字符串键为长度下界)
Max int64 // 数值键最大值(字符串键为长度上界;-1 不限制)
Description string // 中文说明
}
// KVSchema v2 全量配置键 schema 表(含既有策略键,供管理端/文档/AI 校验)。
// 注意:Default 与 config defaults() 逐一对应(schema_test 保证)。
func KVSchema() []KVSchemaEntry {
return []KVSchemaEntry{
// —— 需求 ① 背景图 ——
{KeyBackgroundURL, "string", "", 0, BackgroundURLMaxLen, "背景图 URL 或上传后地址(空=主题默认)"},
// —— 需求 ② 页脚 ——
{KeyFooterText, "string", "", 0, FooterTextMaxLen, "页脚自定义内容(纯文本或受控 HTML 片段)"},
{KeyFooterBeian, "string", "", 0, FooterBeianMaxLen, "备案号,展示于页脚"},
// —— 需求 ③ 系统通知 ——
{KeyNotifyEnabled, "int", 1, 0, 1, "系统通知开关:1 右上角悬浮窗展示 / 0 关闭"},
{KeyNotifyTitle, "string", "系统通知", 0, NotifyTitleMaxLen, "通知标题"},
{KeyNotifyContent, "string", "欢迎使用文件快传,拖拽或粘贴即可分享文本与文件。", 0, NotifyContentMaxLen, "通知内容(允许 <a> 等受控 HTML"},
// —— 需求 ④ 保存策略 ——
{KeyMaxSaveSeconds, "int64", int64(0), 0, MaxSaveSecondsMax, "最长保存秒数上限,0=不限制(默认 7 天兜底)"},
{KeyMaxSaveCount, "int", 0, 0, MaxSaveCountMax, "单次分享最大可取次数上限,0=不限制"},
{KeyExpireStyle, "[]string", []string{"day", "hour", "minute", "forever", "count"}, -1, -1, "上传页可选过期方式白名单"},
// —— 需求 ④ 上传频率限制(既有键对齐参考)——
{KeyUploadCount, "int", 10, 1, 10000, "频率窗口内允许的上传次数"},
{KeyUploadMinute, "int", 1, 1, 1440, "上传频率窗口(分钟)"},
// —— 需求 ④⑩ 存储策略 ——
{KeyMaxFileSize, "int64", int64(0), 0, MaxFileSizeMax, "存储策略-单文件上限(字节),0=回落 uploadSize"},
{KeyUploadSize, "int64", int64(1024 * 1024 * 10), 1024, MaxFileSizeMax, "单文件上限(字节),参考实现语义"},
{KeyAllowedTypes, "[]string", []string{"*"}, -1, -1, "允许上传类型白名单(\"*\" 不限制)"},
{KeyStorageLimit, "int64", int64(0), 0, -1, "站点总容量(字节),0=不限制"},
{KeyOpenUpload, "int", 1, 0, 1, "游客上传开关:1 开 / 0 需管理员登录"},
// —— v3 存储引擎(热切换;引擎参数键沿用 defaults() 既有键,管理端经 config get/update 读写)——
{KeyStorageEngine, "string", "", 0, 16, "当前存储引擎:local|s3|webdav(热切换,健康检查通过才生效;空=回落启动值 FCB_STORAGE_ENGINE"},
{KeySiteDomain, "string", "", 0, 256, "站点对外域名(http(s)://host[:port],不带路径;空=分享链接用当前访问地址)"},
}
}
+91
View File
@@ -0,0 +1,91 @@
// schema 同步测试:保证 config.KVSchema() 的默认值/键集合与 defaults() 完全一致,
// 与 settings 包 re-export 的键名常量同源。新增键时任何一处漏改都会在此失败。
package config
import (
"encoding/json"
"testing"
)
// TestKVSchemaDefaultsMatchDefaults KVSchema 的 Default 必须 === defaults() 中同名键。
func TestKVSchemaDefaultsMatchDefaults(t *testing.T) {
def := defaults()
for _, e := range KVSchema() {
want, ok := def[e.Key]
if !ok {
t.Fatalf("schema 键 %q 缺少 defaults() 默认值", e.Key)
}
// 类型规范化比较(JSON 序列化可比较 []string / int / float
a, _ := json.Marshal(e.Default)
b, _ := json.Marshal(want)
if string(a) != string(b) {
t.Fatalf("键 %q 默认值不一致: schema=%s defaults=%s", e.Key, a, b)
}
}
}
// TestKVSchemaNoDuplicates 键名不得重复。
func TestKVSchemaNoDuplicates(t *testing.T) {
seen := map[string]bool{}
for _, e := range KVSchema() {
if seen[e.Key] {
t.Fatalf("schema 键 %q 重复定义", e.Key)
}
seen[e.Key] = true
}
}
// TestV2NewKeysPresent v2 新增键必须在 schema 与 defaults 中同时存在。
func TestV2NewKeysPresent(t *testing.T) {
def := defaults()
newKeys := []string{
KeyBackgroundURL, KeyFooterText, KeyFooterBeian,
KeyNotifyEnabled, KeyMaxSaveCount, KeyMaxFileSize,
}
for _, k := range newKeys {
if _, ok := def[k]; !ok {
t.Fatalf("v2 新键 %q 缺少默认值", k)
}
}
}
// TestV2AccessorDefaults v2 便捷访问器默认语义。
func TestV2AccessorDefaults(t *testing.T) {
t.Setenv("FCB_DB_DRIVER", "sqlite")
t.Setenv("FCB_DB_DSN", "")
c, err := New()
if err != nil {
t.Fatalf("New: %v", err)
}
// 背景图:background_url 与 background 均空 → 空
if c.BackgroundURL() != "" {
t.Fatalf("背景图默认应为空: %q", c.BackgroundURL())
}
// legacy background 键兜底
c.ApplyKV(map[string]any{KeyBackground: "/legacy/bg.jpg"})
if c.BackgroundURL() != "/legacy/bg.jpg" {
t.Fatalf("legacy background 应回落生效: %q", c.BackgroundURL())
}
// max_file_size > 0 时优先于 uploadSize
c.ApplyKV(map[string]any{KeyMaxFileSize: int64(1024), KeyUploadSize: int64(2048)})
if c.MaxFileSize() != 1024 {
t.Fatalf("max_file_size 应优先: %d", c.MaxFileSize())
}
// max_file_size = 0 回落 uploadSize
c.ApplyKV(map[string]any{KeyMaxFileSize: int64(0)})
if c.MaxFileSize() != 2048 {
t.Fatalf("max_file_size=0 应回落 uploadSize: %d", c.MaxFileSize())
}
// 通知默认开启
if !c.NotifyEnabled() {
t.Fatal("notify_enabled 默认应开启")
}
// 保存次数上限默认不限制
if c.MaxSaveCount() != 0 {
t.Fatalf("max_save_count 默认应 0: %d", c.MaxSaveCount())
}
// 页脚默认空
if c.FooterText() != "" || c.FooterBeian() != "" {
t.Fatal("页脚默认应为空")
}
}
+139
View File
@@ -0,0 +1,139 @@
// Package database 负责数据库连接与迁移(需求 ⑧:双方言):
// - sqlite(默认):modernc.org/sqlite 纯 Go 驱动(GORM 封装 glebarez/sqlite),零 CGO、零外部依赖;
// - postgres:可选,配置 FCB_DB_DRIVER=postgres + FCB_DB_DSN 后启用。
//
// 两方言共用 GORM 抽象层,AutoMigrate 与全部业务查询保持方言无关;
// 唯一的原生 SQL(migrates 建表)已改为双方言分支。
package database
import (
"context"
"fmt"
"os"
"path/filepath"
"strings"
"time"
"github.com/glebarez/sqlite"
"gorm.io/driver/postgres"
"gorm.io/gorm"
"gorm.io/gorm/logger"
"filecodebox/internal/config"
"filecodebox/internal/model"
)
// Options 连接选项(main.go 从 config.Env 装配)。
type Options struct {
Driver string // sqlite | postgres(空按 sqlite 处理)
DSN string // postgres 连接串;sqlite 为文件路径(空回退 config.DefaultSQLitePath
}
// Open 按驱动连接数据库并执行连接池设置与探活。
func Open(ctx context.Context, opts Options) (*gorm.DB, error) {
driver := strings.ToLower(strings.TrimSpace(opts.Driver))
if driver == "" {
driver = config.DBDriverSQLite
}
var dialector gorm.Dialector
switch driver {
case config.DBDriverSQLite:
path := strings.TrimSpace(opts.DSN)
if path == "" {
path = config.DefaultSQLitePath
}
// 自动创建父目录(如 ./data),对齐参考实现 data_root 语义
if dir := filepath.Dir(path); dir != "" && dir != "." {
if err := os.MkdirAll(dir, 0o755); err != nil {
return nil, fmt.Errorf("database: 创建 SQLite 目录 %s 失败: %w", dir, err)
}
}
// DSN 参数:busy_timeout 防写锁竞态;WAL 提升并发读写(Query 参数形式,驱动原生支持)
dsn := path + "?_pragma=busy_timeout(10000)&_pragma=journal_mode(WAL)&_pragma=foreign_keys(1)"
dialector = sqlite.Open(dsn)
case config.DBDriverPostgres:
if strings.TrimSpace(opts.DSN) == "" {
return nil, fmt.Errorf("database: FCB_DB_DRIVER=postgres 需要提供 FCB_DB_DSN")
}
dialector = postgres.Open(opts.DSN)
default:
return nil, fmt.Errorf("database: 不支持的数据库驱动 %q(仅支持 sqlite|postgres", driver)
}
db, err := gorm.Open(dialector, &gorm.Config{
Logger: logger.Default.LogMode(logger.Warn),
// 避免 GORM 生成方言特有子句;时间语义由应用层统一(容器本地时区)
NowFunc: time.Now,
})
if err != nil {
return nil, fmt.Errorf("database: 连接 %s 失败: %w", driver, err)
}
sqlDB, err := db.DB()
if err != nil {
return nil, err
}
// 连接池:SQLite 单文件场景保守设置;Postgres 沿用 v1 参数
switch driver {
case config.DBDriverSQLite:
sqlDB.SetMaxOpenConns(8)
sqlDB.SetMaxIdleConns(4)
sqlDB.SetConnMaxLifetime(0) // 长连接文件句柄,无需轮换
case config.DBDriverPostgres:
sqlDB.SetMaxOpenConns(32)
sqlDB.SetMaxIdleConns(8)
sqlDB.SetConnMaxLifetime(time.Hour)
}
// 连接探活(带超时)
pingCtx, cancel := context.WithTimeout(ctx, 10*time.Second)
defer cancel()
if err := sqlDB.PingContext(pingCtx); err != nil {
return nil, fmt.Errorf("database: %s 探活失败: %w", driver, err)
}
return db, nil
}
// Migrate 执行迁移:先建迁移台账表(双方言分支),再 AutoMigrate 全部模型。
func Migrate(ctx context.Context, db *gorm.DB) error {
if err := createMigratesTable(ctx, db); err != nil {
return err
}
if err := model.AutoMigrate(db); err != nil {
return fmt.Errorf("database: AutoMigrate 失败: %w", err)
}
return nil
}
// createMigratesTable 创建迁移台账表。
// 双方言差异:自增主键 postgres 用 BIGSERIAL、sqlite 用 INTEGER PRIMARY KEY AUTOINCREMENT
// 时间戳默认值 postgres 用 CURRENT_TIMESTAMP、sqlite 用 CURRENT_TIMESTAMP(等价)。
func createMigratesTable(ctx context.Context, db *gorm.DB) error {
ddl := `
CREATE TABLE IF NOT EXISTS migrates (
id INTEGER PRIMARY KEY AUTOINCREMENT,
migration_file VARCHAR(255) NOT NULL UNIQUE,
executed_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
)`
if db.Dialector.Name() == config.DBDriverPostgres {
ddl = `
CREATE TABLE IF NOT EXISTS migrates (
id BIGSERIAL PRIMARY KEY,
migration_file VARCHAR(255) NOT NULL UNIQUE,
executed_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
)`
}
if err := db.WithContext(ctx).Exec(ddl).Error; err != nil {
return fmt.Errorf("database: 创建 migrates 表失败: %w", err)
}
return nil
}
// Close 关闭底层连接。
func Close(db *gorm.DB) error {
sqlDB, err := db.DB()
if err != nil {
return err
}
return sqlDB.Close()
}
+237
View File
@@ -0,0 +1,237 @@
// 数据库双方言测试(需求 ⑧):
// - sqlite:始终执行(纯 Go,临时目录建库);
// - postgres:设置 FCB_TEST_PG_DSN(真实连接串)后执行,未设置时跳过。
//
// 覆盖:Open/Migrate 全表建立、settings KV 读写、JSON 字段往返、
// 分页查询(LIMIT/OFFSET 语义)、布尔/时间字段往返 —— 双方言逐项比对。
package database_test
import (
"context"
"encoding/json"
"fmt"
"os"
"path/filepath"
"strings"
"testing"
"time"
"gorm.io/gorm"
"filecodebox/internal/database"
"filecodebox/internal/model"
)
// pgTestDSN 返回 Postgres 测试连接串;未设置 FCB_TEST_PG_DSN 时返回空。
func pgTestDSN(t *testing.T) string {
t.Helper()
dsn := os.Getenv("FCB_TEST_PG_DSN")
if dsn == "" {
t.Skip("未设置 FCB_TEST_PG_DSN,跳过 Postgres 双方言用例(sqlite 用例仍执行)")
}
return dsn
}
// openTestDB 按方言打开数据库并执行迁移;返回 gorm 实例与关闭函数。
func openTestDB(t *testing.T, driver, dsn string) (*gorm.DB, func()) {
t.Helper()
if dsn == "" {
// sqlite:临时文件库
dir := t.TempDir()
dsn = filepath.Join(dir, "test.db")
}
ctx := context.Background()
db, err := database.Open(ctx, database.Options{Driver: driver, DSN: dsn})
if err != nil {
t.Fatalf("[%s] Open 失败: %v", driver, err)
}
if err := database.Migrate(ctx, db); err != nil {
_ = database.Close(db)
t.Fatalf("[%s] Migrate 失败: %v", driver, err)
}
return db, func() { _ = database.Close(db) }
}
// runDialectSuite 双方言共用的行为断言集。
func runDialectSuite(t *testing.T, db *gorm.DB) {
t.Helper()
ctx := context.Background()
// —— 1. 全表建立 ——
for _, m := range model.AllModels() {
if !db.Migrator().HasTable(m) {
t.Fatalf("表 %T 未创建", m)
}
}
// —— 2. settings KV 读写 + JSON 字段往返 ——
// GORM 软特性:KeyValue.Value 为 *stringJSON 文本),双方言 text 类型
// 可重跑:先清掉同键旧行(共享测试库场景)
if err := db.WithContext(ctx).Where(model.KeyValue{Key: "settings"}).Delete(&model.KeyValue{}).Error; err != nil {
t.Fatalf("KV 旧数据清理失败: %v", err)
}
kv := map[string]any{"background_url": "https://example.com/bg.jpg", "footer_beian": "京ICP备2024000001号-1", "max_save_seconds": 3600}
raw, err := json.Marshal(kv)
if err != nil {
t.Fatalf("marshal KV: %v", err)
}
row := model.KeyValue{Key: "settings", Value: strPtr(string(raw))}
if err := db.WithContext(ctx).Create(&row).Error; err != nil {
t.Fatalf("KV 写入失败: %v", err)
}
var got model.KeyValue
if err := db.WithContext(ctx).Where(model.KeyValue{Key: "settings"}).First(&got).Error; err != nil {
t.Fatalf("KV 读取失败: %v", err)
}
parsed := map[string]any{}
if err := json.Unmarshal([]byte(*got.Value), &parsed); err != nil {
t.Fatalf("KV JSON 解析失败: %v", err)
}
if parsed["background_url"] != "https://example.com/bg.jpg" {
t.Fatalf("KV JSON 字段往返不一致: %v", parsed)
}
// 更新(先查后改,方言无关)
if err := db.WithContext(ctx).Model(&got).Update("value", strPtr(`{"notify_enabled":0}`)).Error; err != nil {
t.Fatalf("KV 更新失败: %v", err)
}
var got2 model.KeyValue
_ = db.WithContext(ctx).Where(model.KeyValue{Key: "settings"}).First(&got2)
if *got2.Value != `{"notify_enabled":0}` {
t.Fatalf("KV 更新未生效: %s", *got2.Value)
}
// —— 3. 分页查询(LIMIT/OFFSET)——
// 每次运行用随机前缀避免脏数据互相影响
prefix := fmt.Sprintf("pg%d_", time.Now().UnixNano())
for i := 0; i < 25; i++ {
fc := model.FileCodes{
Code: fmt.Sprintf("%s%03d", prefix, i),
ExpiredCount: -1,
IsChunked: i%2 == 0, // 布尔字段往返
}
if err := db.WithContext(ctx).Create(&fc).Error; err != nil {
t.Fatalf("FileCodes 写入失败: %v", err)
}
}
var page []model.FileCodes
if err := db.WithContext(ctx).
Where("code LIKE ?", prefix+"%").
Order("id ASC").
Limit(10).Offset(20).
Find(&page).Error; err != nil {
t.Fatalf("分页查询失败: %v", err)
}
if len(page) != 5 {
t.Fatalf("第二页应剩 5 条,实际 %d", len(page))
}
if page[0].Code != prefix+"020" {
t.Fatalf("分页偏移错误: %s", page[0].Code)
}
var total int64
if err := db.WithContext(ctx).Model(&model.FileCodes{}).
Where("code LIKE ?", prefix+"%").Count(&total).Error; err != nil {
t.Fatalf("计数查询失败: %v", err)
}
if total != 25 {
t.Fatalf("总数应 25,实际 %d", total)
}
// —— 4. 布尔/时间/可空字段往返 ——
now := time.Now().Truncate(time.Second) // sqlite 秒级精度
fc := model.FileCodes{
Code: prefix + "special",
ExpiredAt: &now,
ExpiredCount: 5,
Text: strPtr("你好 FileCodeBox"),
FileHash: strPtr("abc123"),
IsChunked: true,
}
if err := db.WithContext(ctx).Create(&fc).Error; err != nil {
t.Fatalf("完整字段写入失败: %v", err)
}
var back model.FileCodes
if err := db.WithContext(ctx).Where(model.FileCodes{Code: fc.Code}).First(&back).Error; err != nil {
t.Fatalf("完整字段读取失败: %v", err)
}
if back.Text == nil || *back.Text != "你好 FileCodeBox" {
t.Fatalf("text 字段往返不一致: %v", back.Text)
}
if !back.IsChunked {
t.Fatal("布尔字段往返不一致")
}
if back.ExpiredAt == nil {
t.Fatal("时间字段往返丢失")
}
if diff := back.ExpiredAt.Sub(now); diff > time.Second || diff < -time.Second {
t.Fatalf("时间字段偏差过大: %v", diff)
}
if back.FileHash == nil || *back.FileHash != "abc123" {
t.Fatalf("可空字段往返不一致: %v", back.FileHash)
}
// LOWER + LIKEadmin 列表检索路径:真实代码先对关键词小写化再拼 LIKE 模式,
// 对齐 admin.go 的 "LOWER(code) LIKE ?" 用法,双方言均支持)
var hits int64
lowerPattern := "%" + strings.ToLower(prefix+"SPECIAL") + "%"
if err := db.WithContext(ctx).Model(&model.FileCodes{}).
Where("LOWER(code) LIKE ?", lowerPattern).Count(&hits).Error; err != nil {
t.Fatalf("LOWER/LIKE 查询失败: %v", err)
}
if hits != 1 {
t.Fatalf("LOWER/LIKE 命中数应 1,实际 %d", hits)
}
// 可重跑:清理本前缀数据(共享测试库场景)
if err := db.WithContext(ctx).Where("code LIKE ?", prefix+"%").Delete(&model.FileCodes{}).Error; err != nil {
t.Fatalf("清理测试数据失败: %v", err)
}
}
func strPtr(s string) *string { return &s }
// TestSQLiteDialect sqlite(默认模式):临时文件库全流程。
func TestSQLiteDialect(t *testing.T) {
db, closeFn := openTestDB(t, "sqlite", "")
defer closeFn()
runDialectSuite(t, db)
}
// TestSQLiteInMemoryDialect sqlite 内存库(DSN 为 :memory: 等价路径场景)。
func TestSQLiteInMemoryDialect(t *testing.T) {
dir := t.TempDir()
db, closeFn := openTestDB(t, "sqlite", filepath.Join(dir, "mem.db"))
defer closeFn()
runDialectSuite(t, db)
}
// TestPostgresDialect postgres(可选模式):FCB_TEST_PG_DSN 指向真实实例。
func TestPostgresDialect(t *testing.T) {
dsn := pgTestDSN(t)
db, closeFn := openTestDB(t, "postgres", dsn)
defer closeFn()
runDialectSuite(t, db)
}
// TestOpenRejectsUnknownDriver 非法驱动应报错。
func TestOpenRejectsUnknownDriver(t *testing.T) {
if _, err := database.Open(context.Background(), database.Options{Driver: "mysql", DSN: "x"}); err == nil {
t.Fatal("非法驱动应报错")
}
}
// TestOpenPostgresRequiresDSN postgres 模式缺 DSN 应报错。
func TestOpenPostgresRequiresDSN(t *testing.T) {
if _, err := database.Open(context.Background(), database.Options{Driver: "postgres", DSN: ""}); err == nil {
t.Fatal("postgres 缺 DSN 应报错")
}
}
// TestSQLiteAutoCreatesDataDir sqlite 默认相对路径下自动创建父目录。
func TestSQLiteAutoCreatesDataDir(t *testing.T) {
dir := t.TempDir()
nested := filepath.Join(dir, "deep", "data", "fcb.db")
db, closeFn := openTestDB(t, "sqlite", nested)
defer closeFn()
if _, err := os.Stat(nested); err != nil {
t.Fatalf("数据库文件应已创建: %v", err)
}
runDialectSuite(t, db)
}
+124
View File
@@ -0,0 +1,124 @@
// Package janitor 后台清理循环(安全审计 M5):
// 回收过期容量预留、超时未完成的上传会话(含其分片对象)与过期预签名会话
// (direct 模式残留对象一并删除)。此前这些资源仅在同 token 复用/显式取消时
// 释放,恶意 init 可长期占用容量预留或累积垃圾数据。
package janitor
import (
"context"
"errors"
"log"
"time"
"gorm.io/gorm"
"filecodebox/internal/model"
"filecodebox/internal/storage"
)
// chunkSessionMaxAge 未完成分片会话的最大保留时长(预留 TTL 为 2h,
// 会话保留 24h 以支持断点续传;超时后由本循环清理)。
const chunkSessionMaxAge = 24 * time.Hour
// presignGrace 过期预签名会话的宽限时长(到点即删,避免与在途 confirm 竞争)。
const presignGrace = time.Hour
// Start 启动周期清理循环;ctx 取消时退出。
func Start(ctx context.Context, db *gorm.DB, store *storage.Manager, interval time.Duration) {
go func() {
ticker := time.NewTicker(interval)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
Run(ctx, db, store)
}
}
}()
}
// Run 执行一轮清理;单项失败仅记日志,不影响其他项。
func Run(ctx context.Context, db *gorm.DB, store *storage.Manager) {
now := time.Now()
cleanExpiredReservations(ctx, db, now)
cleanExpiredChunkSessions(ctx, db, store, now)
cleanExpiredPresignSessions(ctx, db, store, now)
}
// cleanExpiredReservations 删除全部过期容量预留。
func cleanExpiredReservations(ctx context.Context, db *gorm.DB, now time.Time) {
if err := db.WithContext(ctx).
Where("expires_at <= ?", now).
Delete(&model.StorageReservation{}).Error; err != nil {
log.Printf("[janitor] 清理过期容量预留失败: %v", err)
}
}
// engineFor 按归属引擎取回实例;空/未知引擎回落当前引擎(对齐 API 层 storeFor 语义)。
func engineFor(store *storage.Manager, name string) (storage.Storage, error) {
if name != "" && storage.ValidEngine(name) {
if s, err := store.EngineOf(name); err == nil {
return s, nil
}
}
return store.Current(), nil
}
// cleanExpiredChunkSessions 清理超时未完成的分片会话及其分片对象。
func cleanExpiredChunkSessions(ctx context.Context, db *gorm.DB, store *storage.Manager, now time.Time) {
var sessions []model.UploadChunk
if err := db.WithContext(ctx).
Where("chunk_index = -1 AND created_at < ?", now.Add(-chunkSessionMaxAge)).
Limit(200).
Find(&sessions).Error; err != nil {
log.Printf("[janitor] 查询过期分片会话失败: %v", err)
return
}
for _, s := range sessions {
engine, err := engineFor(store, s.Engine)
if err == nil && s.SavePath != "" {
if err := engine.CleanChunks(ctx, s.UploadID, s.SavePath); err != nil &&
!errors.Is(err, storage.ErrNotFound) && !errors.Is(err, storage.ErrInvalidPath) {
log.Printf("[janitor] 清理分片对象失败 upload_id=%s: %v", s.UploadID, err)
}
}
if err := db.WithContext(ctx).
Where("upload_id = ?", s.UploadID).
Delete(&model.UploadChunk{}).Error; err != nil {
log.Printf("[janitor] 删除过期分片会话失败 upload_id=%s: %v", s.UploadID, err)
continue
}
log.Printf("[janitor] 已清理超时分片会话 upload_id=%s file=%s", s.UploadID, s.FileName)
}
}
// cleanExpiredPresignSessions 清理过期预签名会话;direct 模式残留对象一并删除。
func cleanExpiredPresignSessions(ctx context.Context, db *gorm.DB, store *storage.Manager, now time.Time) {
var sessions []model.PresignUploadSession
if err := db.WithContext(ctx).
Where("expires_at < ?", now.Add(-presignGrace)).
Limit(200).
Find(&sessions).Error; err != nil {
log.Printf("[janitor] 查询过期预签名会话失败: %v", err)
return
}
for _, s := range sessions {
if s.Mode == "direct" && s.SavePath != "" {
if engine, err := engineFor(store, s.Engine); err == nil {
if err := engine.DeleteFile(ctx, s.SavePath); err != nil &&
!errors.Is(err, storage.ErrNotFound) && !errors.Is(err, storage.ErrInvalidPath) {
log.Printf("[janitor] 删除直传残留对象失败 upload_id=%s: %v", s.UploadID, err)
}
}
}
if err := db.WithContext(ctx).
Where("upload_id = ?", s.UploadID).
Delete(&model.PresignUploadSession{}).Error; err != nil {
log.Printf("[janitor] 删除过期预签名会话失败 upload_id=%s: %v", s.UploadID, err)
continue
}
log.Printf("[janitor] 已清理过期预签名会话 upload_id=%s mode=%s", s.UploadID, s.Mode)
}
}
+261
View File
@@ -0,0 +1,261 @@
package middleware
import (
"net/http"
"strings"
"time"
"github.com/gin-gonic/gin"
"filecodebox/internal/audit"
"filecodebox/internal/model"
"filecodebox/internal/response"
)
// auditHooks 审计钩子:由 API 层在响应前后填充与落库。
// 中间件负责计时与公共字段(IP/UA/设备/耗时),业务上下文通过 auditEntry 传递。
type auditEntry struct {
Entry audit.Entry
// start 请求进入审计中间件的时刻,用于计算耗时。
start time.Time
// writer 下载动作时包装的响应计数器。
writer *bytesCountWriter
// skip 为 true 表示业务 handler 显式跳过审计(AuditSkip)。
skip bool
// recorded 防止重复落库。
recorded bool
}
// bytesCountWriter 统计响应体写出字节数(用于下载审计)。
type bytesCountWriter struct {
gin.ResponseWriter
count int64
}
func (w *bytesCountWriter) Write(b []byte) (int, error) {
n, err := w.ResponseWriter.Write(b)
w.count += int64(n)
return n, err
}
func (w *bytesCountWriter) WriteString(s string) (int, error) {
n, err := w.ResponseWriter.WriteString(s)
w.count += int64(n)
return n, err
}
// Classifier 判定请求是否属于需审计的动作;返回动作名与是否命中。
type Classifier func(c *gin.Context) (action string, ok bool)
// DefaultClassifier 按任务合同的默认路由语义分类:
// - 上传:POST /share/file、/share/text、/chunk/upload*、/presign*
// - 下载:GET /share/download、/share/select、/share/metadata
// - 管理(L5):POST/PATCH/DELETE 的敏感管理操作——登录/登出、配置与密码
// 修改、存储引擎切换、文件更新/删除/策略动作
//
// API 层可传入自定义分类器覆盖。
func DefaultClassifier(c *gin.Context) (string, bool) {
path := c.FullPath()
if path == "" {
path = c.Request.URL.Path
}
p := strings.TrimRight(path, "/")
switch c.Request.Method {
case http.MethodPost, http.MethodPut:
switch {
case p == "/share/file" || p == "/share/text":
return audit.ActionUpload, true
case strings.HasPrefix(p, "/chunk/upload"):
return audit.ActionUpload, true
case strings.HasPrefix(p, "/presign"):
return audit.ActionUpload, true
}
if adminAuditActions[p] {
return audit.ActionAdmin, true
}
case http.MethodPatch, http.MethodDelete:
if adminAuditActions[p] {
return audit.ActionAdmin, true
}
case http.MethodGet:
switch p {
case "/share/download", "/share/select", "/share/metadata":
return audit.ActionDownload, true
}
}
return "", false
}
// adminAuditActions 需要审计的管理端敏感操作路由(L5)。
var adminAuditActions = map[string]bool{
"/admin/login": true,
"/admin/logout": true,
"/admin/config/update": true,
"/admin/settings/password": true,
"/admin/storage/switch": true,
"/admin/file/update": true,
"/admin/file/delete": true,
"/admin/file/batch-delete": true,
"/admin/file/batch-update": true,
"/admin/file/policy-action": true,
"/admin/file/batch-policy-action": true,
}
// Audit 审计中间件:对分类器命中的 upload/download/admin 动作写审计日志。
// handler 通过 AuditSet 填充取件码/文件名/字节数等业务字段;
// handler 未显式 AuditRecordRequest 时按 HTTP 状态兜底落库。
func Audit(service *audit.Service, classify Classifier) gin.HandlerFunc {
if classify == nil {
classify = DefaultClassifier
}
return func(c *gin.Context) {
start := time.Now()
action, ok := classify(c)
// 未命中审计动作的请求直接放行,不产生审计记录。
// L5admin 类动作同样需要建 auditEntry 并落库(登录失败/配置变更等)。
if !ok {
c.Next()
return
}
entry := audit.Entry{
Action: action,
IP: GetClientIP(c),
UserAgent: c.Request.UserAgent(),
}
info := audit.ParseUserAgent(entry.UserAgent)
entry.DeviceOS = info.OS
entry.DeviceBrowser = info.Browser
entry.DeviceType = info.Type
// 交给后续 handler 填充
state := &auditEntry{Entry: entry, start: start}
c.Set("auditEntry", state)
// 下载动作:包装 Writer 以捕获实际写出字节数(必须在 c.Next() 前替换)
if action == audit.ActionDownload {
state.writer = &bytesCountWriter{ResponseWriter: c.Writer}
c.Writer = state.writer
}
c.Next()
// 下载兜底统计:handler 未填 TransferredBytes 时取响应写出字节
if action == audit.ActionDownload && state.Entry.TransferredBytes == 0 &&
!state.recorded && !state.skip && state.writer != nil {
state.Entry.TransferredBytes = state.writer.count
}
// handler 未显式落库时兜底记录
ae, exists := c.Get("auditEntry")
if !exists {
return
}
state, isState := ae.(*auditEntry)
if !isState || state.recorded || state.skip {
return
}
state.Entry.Duration = time.Since(start)
state.Entry.Actor = resolveActor(c)
status := c.Writer.Status()
switch {
case state.Entry.Result != "":
// handler 已给出结论
case status >= 500:
state.Entry.Result = model.AuditResultFailed
case status == 401 || status == 403 || status == 423 || status == 429 || status == 428:
state.Entry.Result = model.AuditResultDenied
case status >= 400:
state.Entry.Result = model.AuditResultFailed
default:
state.Entry.Result = model.AuditResultSuccess
}
switch {
case state.Entry.ErrorMsg != "":
// handler 已给出错误信息
case c.Errors.String() != "":
state.Entry.ErrorMsg = c.Errors.String()
case status >= 400:
// 兜底:记录 HTTP 状态
state.Entry.ErrorMsg = "HTTP " + itoa64(int64(status))
}
service.Record(state.Entry)
state.recorded = true
}
}
// AuditEntry 获取当前请求的审计状态(由 Audit 中间件创建)。
func AuditEntry(c *gin.Context) *auditEntry {
if v, ok := c.Get("auditEntry"); ok {
if ae, ok := v.(*auditEntry); ok {
return ae
}
}
return nil
}
// AuditSet 填充当前请求的审计字段;仅对已启用审计的请求生效。
func AuditSet(c *gin.Context, fn func(e *audit.Entry)) {
if ae := AuditEntry(c); ae != nil && fn != nil {
fn(&ae.Entry)
}
}
// AuditRecordRequest 显式触发落库(含耗时);由 handler 在响应前调用。
func AuditRecordRequest(c *gin.Context, service *audit.Service, result, errMsg string) {
ae := AuditEntry(c)
if ae == nil || ae.recorded || ae.skip {
return
}
ae.Entry.Duration = time.Since(ae.start)
ae.Entry.Result = result
ae.Entry.ErrorMsg = errMsg
ae.Entry.Actor = resolveActor(c)
service.Record(ae.Entry)
ae.recorded = true
}
// AuditSkip 标记当前请求不写审计。
func AuditSkip(c *gin.Context) {
if ae := AuditEntry(c); ae != nil {
ae.skip = true
}
}
// resolveActor 判断请求者角色:管理员 JWT 有效 → admin,否则 guest。
func resolveActor(c *gin.Context) string {
header := c.GetHeader("Authorization")
if len(header) > 7 && header[:7] == "Bearer " {
// 仅检查声明是否有效,不重复校验签名逻辑(AdminAuth 已处理受保护路由)
if _, ok := c.Get("claims"); ok {
return audit.ActorAdmin
}
}
return audit.ActorGuest
}
// AuditRecord 显式按结果落库;duration 由中间件按起始时间计算。
func AuditRecord(c *gin.Context, service *audit.Service, result, errMsg string) {
ae := AuditEntry(c)
if ae == nil || ae.recorded || ae.skip {
return
}
AuditRecordRequest(c, service, result, errMsg)
}
// GuardNotInitialized 系统未初始化守卫:除 setup/health 外返回 428。
func GuardNotInitialized(isInit func() bool) gin.HandlerFunc {
return func(c *gin.Context) {
if isInit() {
c.Next()
return
}
path := c.Request.URL.Path
if path == "/setup" || path == "/api/v1/health" {
c.Next()
return
}
response.Fail(c, 428, "系统未初始化,请先完成初始化")
}
}
@@ -0,0 +1,89 @@
// audit_l5_test.go — L5 回归:admin 类动作(如登录失败)必须落审计。
package middleware
import (
"context"
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/gin-gonic/gin"
"filecodebox/internal/audit"
"filecodebox/internal/model"
)
type captureSink struct {
logs []model.AuditLog
}
func (s *captureSink) Save(_ context.Context, logs []model.AuditLog) error {
s.logs = append(s.logs, logs...)
return nil
}
// TestAuditRecordsAdminActions L5/admin/login 失败(401)后应产生一条
// result=denied 的 admin 审计记录(此前 skip 条件把 admin 动作整体跳过)。
func TestAuditRecordsAdminActions(t *testing.T) {
gin.SetMode(gin.TestMode)
sink := &captureSink{}
svc := audit.NewService(sink)
r := gin.New()
r.Use(Audit(svc, nil)) // DefaultClassifier
r.POST("/admin/login", func(c *gin.Context) {
c.JSON(http.StatusUnauthorized, gin.H{"code": 401})
})
r.POST("/share/text", func(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"code": 200})
})
r.GET("/healthz", func(c *gin.Context) {
c.Status(http.StatusOK) // 未分类动作:不应产生审计
})
// 管理端:401 → denied
w := httptest.NewRecorder()
r.ServeHTTP(w, httptest.NewRequest("POST", "/admin/login", nil))
if w.Code != http.StatusUnauthorized {
t.Fatalf("login should 401, got %d", w.Code)
}
// 上传类:200 → success
w2 := httptest.NewRecorder()
r.ServeHTTP(w2, httptest.NewRequest("POST", "/share/text", nil))
// 未分类:不落库
w3 := httptest.NewRecorder()
r.ServeHTTP(w3, httptest.NewRequest("GET", "/healthz", nil))
// audit.Service 异步落库,轮询等待
var actions []string
for i := 0; i < 50; i++ {
if len(sink.logs) >= 2 {
break
}
waitMillis(20)
}
if len(sink.logs) != 2 {
t.Fatalf("应恰好 2 条审计记录, got %d", len(sink.logs))
}
for _, l := range sink.logs {
actions = append(actions, l.Action)
switch l.Action {
case audit.ActionAdmin:
if l.Result != model.AuditResultDenied {
t.Fatalf("admin 401 应记 denied, got %q", l.Result)
}
case audit.ActionUpload:
if l.Result != model.AuditResultSuccess {
t.Fatalf("upload 200 应记 success, got %q", l.Result)
}
default:
t.Fatalf("意外动作 %q", l.Action)
}
}
_ = actions
}
func waitMillis(ms int) {
time.Sleep(time.Duration(ms) * time.Millisecond)
}
+198
View File
@@ -0,0 +1,198 @@
package middleware
import (
"context"
"net/http/httptest"
"sync"
"testing"
"time"
"github.com/gin-gonic/gin"
"filecodebox/internal/audit"
"filecodebox/internal/model"
)
// memSink 测试用内存落库实现。
type memSink struct {
mu sync.Mutex
logs []model.AuditLog
notif chan struct{}
}
func newMemSink() *memSink { return &memSink{notif: make(chan struct{}, 16)} }
func (m *memSink) Save(_ context.Context, logs []model.AuditLog) error {
m.mu.Lock()
m.logs = append(m.logs, logs...)
m.mu.Unlock()
m.notif <- struct{}{}
return nil
}
func (m *memSink) snapshot() []model.AuditLog {
m.mu.Lock()
defer m.mu.Unlock()
out := make([]model.AuditLog, len(m.logs))
copy(out, m.logs)
return out
}
// waitFor 等待 sink 收到 n 条记录(带超时)。
func (m *memSink) waitFor(t *testing.T, n int) []model.AuditLog {
t.Helper()
deadline := time.After(2 * time.Second)
for {
logs := m.snapshot()
if len(logs) >= n {
return logs
}
select {
case <-m.notif:
case <-deadline:
t.Fatalf("等待审计记录超时: 已收到 %d 条", len(logs))
}
}
}
func auditRouter(svc *audit.Service) *gin.Engine {
r := gin.New()
r.Use(ClientIP(nil))
r.Use(Audit(svc, nil)) // 默认分类器
// 上传路由:命中默认分类器(POST /share/file
r.POST("/share/file", func(c *gin.Context) {
// 模拟 handler 填充业务字段并显式落库
AuditSet(c, func(e *audit.Entry) {
e.FileCode = "Ab3xY"
e.FileName = "hello.zip"
e.SizeBytes = 1024
})
AuditRecordRequest(c, svc, model.AuditResultSuccess, "")
c.JSON(200, gin.H{"ok": true})
})
// 下载路由:命中默认分类器(GET /share/download
r.GET("/share/download", func(c *gin.Context) {
AuditSet(c, func(e *audit.Entry) { e.FileCode = "Xy12Z" })
c.JSON(404, gin.H{"msg": "文件已过期删除"}) // 未显式落库 → 状态码兜底
})
// 普通路由:不命中,不应产生审计
r.GET("/plain", func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) })
return r
}
func TestAuditMiddlewareRecordsUpload(t *testing.T) {
sink := newMemSink()
svc := audit.NewService(sink)
r := auditRouter(svc)
w := httptest.NewRecorder()
req := httptest.NewRequest("POST", "/share/file", nil)
req.Header.Set("User-Agent", "Mozilla/5.0 (Windows NT 10.0; Win64; x64) Chrome/120.0.0.0 Safari/537.36")
r.ServeHTTP(w, req)
if w.Code != 200 {
t.Fatalf("上传应成功: %d", w.Code)
}
logs := sink.waitFor(t, 1)
e := logs[0]
if e.Action != audit.ActionUpload {
t.Errorf("action = %s", e.Action)
}
if e.FileCode != "Ab3xY" || e.FileName != "hello.zip" {
t.Errorf("file fields = %s/%s", e.FileCode, e.FileName)
}
if e.SizeBytes != 1024 {
t.Errorf("size = %d", e.SizeBytes)
}
if e.Result != model.AuditResultSuccess {
t.Errorf("result = %s", e.Result)
}
if e.DeviceOS != "Windows" || e.DeviceBrowser != "Chrome" || e.DeviceType != "desktop" {
t.Errorf("device = %s/%s/%s", e.DeviceOS, e.DeviceBrowser, e.DeviceType)
}
if e.DurationMs < 0 {
t.Errorf("duration = %d", e.DurationMs)
}
if e.Actor != audit.ActorGuest {
t.Errorf("actor = %s", e.Actor)
}
}
func TestAuditMiddlewareSkipsPlainRoutes(t *testing.T) {
sink := newMemSink()
svc := audit.NewService(sink)
r := auditRouter(svc)
w := httptest.NewRecorder()
r.ServeHTTP(w, httptest.NewRequest("GET", "/plain", nil))
if w.Code != 200 {
t.Fatalf("plain 路由应成功: %d", w.Code)
}
time.Sleep(150 * time.Millisecond)
if logs := sink.snapshot(); len(logs) != 0 {
t.Fatalf("普通路由不应产生审计记录: %v", logs)
}
}
func TestAuditFailedDownloadFallback(t *testing.T) {
sink := newMemSink()
svc := audit.NewService(sink)
r := auditRouter(svc)
w := httptest.NewRecorder()
req := httptest.NewRequest("GET", "/share/download?code=xyz", nil)
req.Header.Set("User-Agent", "curl/8.4.0")
r.ServeHTTP(w, req)
logs := sink.waitFor(t, 1)
e := logs[0]
if e.Action != audit.ActionDownload {
t.Errorf("action = %s", e.Action)
}
if e.Result != model.AuditResultFailed {
t.Errorf("4xx 兜底 result = %s", e.Result)
}
if e.DeviceType != "bot" {
t.Errorf("curl 应识别为 bot: %s", e.DeviceType)
}
if e.ErrorMsg == "" {
t.Error("失败记录应包含错误信息")
}
}
func TestAuditDeniedStatusMapping(t *testing.T) {
sink := newMemSink()
svc := audit.NewService(sink)
r := gin.New()
r.Use(ClientIP(nil))
r.Use(Audit(svc, nil))
r.GET("/share/select", func(c *gin.Context) { c.AbortWithStatus(429) })
w := httptest.NewRecorder()
r.ServeHTTP(w, httptest.NewRequest("GET", "/share/select?code=abc", nil))
logs := sink.waitFor(t, 1)
if logs[0].Result != model.AuditResultDenied {
t.Errorf("429 应映射为 denied: %s", logs[0].Result)
}
}
func TestAuditDownloadBytesCounted(t *testing.T) {
sink := newMemSink()
svc := audit.NewService(sink)
r := gin.New()
r.Use(ClientIP(nil))
r.Use(Audit(svc, nil))
r.GET("/share/download", func(c *gin.Context) {
payload := []byte("0123456789abcdef") // 16 字节
c.Data(200, "application/octet-stream", payload)
// 未显式落库 → 中间件兜底;TransferredBytes 应等于写出字节
})
w := httptest.NewRecorder()
r.ServeHTTP(w, httptest.NewRequest("GET", "/share/download?code=bytes", nil))
logs := sink.waitFor(t, 1)
if logs[0].TransferredBytes != 16 {
t.Errorf("下载字节数 = %d, want 16", logs[0].TransferredBytes)
}
if logs[0].Result != model.AuditResultSuccess {
t.Errorf("result = %s", logs[0].Result)
}
}
+30
View File
@@ -0,0 +1,30 @@
// Package middleware — bodylimit.go:全局请求体大小限制。
//
// 安全审计 M3:此前服务对请求体无任何上限,text/plain 形态的绑定会先
// io.ReadAll 全量进内存、multipart 大文件先完整落盘后才被策略拒绝,
// 未认证攻击者可借大 body 消耗内存/磁盘/带宽。
// 本中间件用 http.MaxBytesReader 按「路径类别」给出上限:
// - 管理端(/admin/*):1MiB
// - 文本/元信息/取件类(/share/text|metadata|select):1MiB
// - 其余(含上传):maxFileSize0=回落 uploadSize,仍为 0 时 64MiB 兜底)+ 2MiB 表单开销。
//
// 超限时后续读取返回错误,统一被 handler 的 bind 错误路径映射为 400。
package middleware
import (
"net/http"
"github.com/gin-gonic/gin"
)
// BodyLimit 按请求路径动态限制请求体大小(limit<=0 表示不限制)。
func BodyLimit(limitFn func(c *gin.Context) int64) gin.HandlerFunc {
return func(c *gin.Context) {
if c.Request.Body != nil && limitFn != nil {
if limit := limitFn(c); limit > 0 {
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, limit)
}
}
c.Next()
}
}
+74
View File
@@ -0,0 +1,74 @@
package middleware
import (
"net/url"
"strings"
"github.com/gin-gonic/gin"
)
// Cors 跨域中间件(L6 收紧):
// - 公开接口:维持 allow_origins=*Bearer Token 认证,无 Cookie CSRF 面);
// - 管理端(/admin/*):当请求携带 Origin 且既不同源也不在允许域名列表时,
// 不回 CORS 头(浏览器将拦截跨域读取)。防止管理端 token 泄露后
// 被任意第三方页面直接跨域调用。无 Origin 的非浏览器请求不受影响。
//
// extraAllowedOrigins:管理端额外允许的来源(如 site_domain 配置的对外域名)。
func Cors(extraAllowedOrigins ...string) gin.HandlerFunc {
allowedHosts := map[string]bool{}
for _, o := range extraAllowedOrigins {
if o == "" {
continue
}
raw := strings.TrimSpace(o)
if !strings.Contains(raw, "://") {
raw = "https://" + raw
}
if u, err := url.Parse(raw); err == nil && u.Host != "" {
allowedHosts[u.Host] = true
}
}
// adminCrossOriginBlocked 判断 /admin 请求是否应拒绝跨域:
// 仅在「带 Origin 且 Origin 既不同源也不在白名单」时为 true。
adminBlocked := func(c *gin.Context) bool {
p := c.Request.URL.Path
if p != "/admin" && !strings.HasPrefix(p, "/admin/") {
return false
}
origin := c.GetHeader("Origin")
if origin == "" {
return false
}
o, err := url.Parse(origin)
if err != nil || o.Host == "" {
return true // Origin 非法:按跨域拒绝处理
}
if o.Host == c.Request.Host || allowedHosts[o.Host] {
return false
}
return true
}
return func(c *gin.Context) {
if adminBlocked(c) {
// 不回 ACAO;预检直接 204(浏览器会因无 CORS 头拦截后续请求)
if c.Request.Method == "OPTIONS" {
c.AbortWithStatus(204)
return
}
c.Next()
return
}
c.Header("Access-Control-Allow-Origin", "*")
c.Header("Access-Control-Allow-Methods", "GET, POST, PUT, DELETE, OPTIONS, HEAD")
c.Header("Access-Control-Allow-Headers", "Authorization, Content-Type, Content-Disposition, X-Requested-With")
c.Header("Access-Control-Expose-Headers", "Content-Disposition, Content-Length")
c.Header("Access-Control-Max-Age", "86400")
if c.Request.Method == "OPTIONS" {
c.AbortWithStatus(204)
return
}
c.Next()
}
}
+99
View File
@@ -0,0 +1,99 @@
package middleware
import (
"errors"
"strings"
"time"
"github.com/gin-gonic/gin"
"github.com/golang-jwt/jwt/v5"
"filecodebox/internal/response"
)
// jwtClaims 自定义声明:对齐参考实现(payload 含 is_admin 与 exp)。
type jwtClaims struct {
IsAdmin bool `json:"is_admin"`
jwt.RegisteredClaims
}
// 签发/校验相关错误。
var (
ErrTokenExpired = errors.New("token已过期")
ErrTokenInvalid = errors.New("无效的签名")
ErrNotAdmin = errors.New("未授权或授权校验失败")
)
// SignAdminToken 用 HS256 签发管理员 JWT。
// secret 为数据库 settings 中的 jwt_secretexpires 为会话有效期。
func SignAdminToken(secret string, expires time.Duration) (string, time.Time, error) {
if strings.TrimSpace(secret) == "" {
return "", time.Time{}, errors.New("JWT签名密钥未初始化")
}
expiresAt := time.Now().Add(expires)
claims := jwtClaims{
IsAdmin: true,
RegisteredClaims: jwt.RegisteredClaims{
ExpiresAt: jwt.NewNumericDate(expiresAt),
IssuedAt: jwt.NewNumericDate(time.Now()),
Issuer: "filecodebox",
},
}
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
signed, err := token.SignedString([]byte(secret))
return signed, expiresAt, err
}
// VerifyAdminToken 校验管理员 JWT:签名、过期时间与 is_admin 声明。
func VerifyAdminToken(secret, token string) (*jwtClaims, error) {
if strings.TrimSpace(secret) == "" {
return nil, errors.New("JWT签名密钥未初始化")
}
parsed, err := jwt.ParseWithClaims(token, &jwtClaims{}, func(t *jwt.Token) (any, error) {
if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok {
return nil, ErrTokenInvalid
}
return []byte(secret), nil
}, jwt.WithValidMethods([]string{"HS256"}))
if err != nil {
if errors.Is(err, jwt.ErrTokenExpired) {
return nil, ErrTokenExpired
}
return nil, ErrTokenInvalid
}
claims, ok := parsed.Claims.(*jwtClaims)
if !ok || !parsed.Valid {
return nil, ErrTokenInvalid
}
if !claims.IsAdmin {
return nil, ErrNotAdmin
}
return claims, nil
}
// SecretProvider 动态提供当前 jwt_secretsettings KV 运行时可变)。
type SecretProvider func() string
// AdminAuth 管理员鉴权中间件:校验 Authorization: Bearer <token>。
// 成功后把声明写入 gin 上下文(ctxClaims)。
func AdminAuth(secret SecretProvider) gin.HandlerFunc {
return func(c *gin.Context) {
header := c.GetHeader("Authorization")
if !strings.HasPrefix(header, "Bearer ") {
response.Fail(c, 401, "未授权或授权校验失败")
return
}
token := strings.TrimSpace(strings.TrimPrefix(header, "Bearer "))
if token == "" {
response.Fail(c, 401, "未授权或授权校验失败")
return
}
claims, err := VerifyAdminToken(secret(), token)
if err != nil {
response.Fail(c, 401, err.Error())
return
}
c.Set("claims", claims)
c.Next()
}
}
+83
View File
@@ -0,0 +1,83 @@
package middleware
import (
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/gin-gonic/gin"
"filecodebox/internal/response"
)
const testSecret = "unit-test-secret-0123456789abcdef"
func init() { gin.SetMode(gin.TestMode) }
func TestSignAndVerifyAdminToken(t *testing.T) {
token, expiresAt, err := SignAdminToken(testSecret, time.Hour)
if err != nil {
t.Fatalf("签发失败: %v", err)
}
if expiresAt.Before(time.Now()) {
t.Fatal("过期时间不合理")
}
claims, err := VerifyAdminToken(testSecret, token)
if err != nil {
t.Fatalf("校验失败: %v", err)
}
if !claims.IsAdmin {
t.Fatal("is_admin 应为 true")
}
}
func TestVerifyTamperedToken(t *testing.T) {
token, _, _ := SignAdminToken(testSecret, time.Hour)
claims, err := VerifyAdminToken(testSecret+"-wrong", token)
if err == nil || claims != nil {
t.Fatal("密钥不匹配应校验失败")
}
// 篡改 payload
tampered := token[:len(token)-3] + "abc"
if _, err := VerifyAdminToken(testSecret, tampered); err == nil {
t.Fatal("篡改的 token 应校验失败")
}
// 非 HMAC 算法拒绝
algNone := "eyJhbGciOiJub25lIiwidHlwIjoiSldUIn0.eyJpc19hZG1pbiI6dHJ1ZX0."
if _, err := VerifyAdminToken(testSecret, algNone); err == nil {
t.Fatal("none 算法应被拒绝")
}
}
func TestAdminAuthMiddleware(t *testing.T) {
r := gin.New()
r.GET("/protected", AdminAuth(func() string { return testSecret }), func(c *gin.Context) {
response.OK(c, gin.H{"ok": true})
})
// 无 token → 401
w := httptest.NewRecorder()
r.ServeHTTP(w, httptest.NewRequest("GET", "/protected", nil))
if w.Code != http.StatusUnauthorized {
t.Fatalf("无 token 应 401: %d", w.Code)
}
// 有效 token → 200
token, _, _ := SignAdminToken(testSecret, time.Hour)
w = httptest.NewRecorder()
req := httptest.NewRequest("GET", "/protected", nil)
req.Header.Set("Authorization", "Bearer "+token)
r.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("有效 token 应 200: %d %s", w.Code, w.Body.String())
}
// 过期 token → 401
expired, _, _ := SignAdminToken(testSecret, -time.Minute)
w = httptest.NewRecorder()
req = httptest.NewRequest("GET", "/protected", nil)
req.Header.Set("Authorization", "Bearer "+expired)
r.ServeHTTP(w, req)
if w.Code != http.StatusUnauthorized {
t.Fatalf("过期 token 应 401: %d", w.Code)
}
}
+295
View File
@@ -0,0 +1,295 @@
package middleware
import (
"context"
"errors"
"net"
"net/http"
"strings"
"sync"
"time"
"github.com/gin-gonic/gin"
"filecodebox/internal/cache"
"filecodebox/internal/response"
)
// 限流类别(对齐参考 apps/base/utils.py 的 ip_limit)。
const (
LimitError = "error" // 取件错误(密码错误、取件失败)
LimitUpload = "upload" // 上传次数
LimitLogin = "login" // 管理员登录失败
LimitMeta = "metadata" // 分享元信息查询
)
// LimitRule 限流规则:window 内最多 count 次。
type LimitRule struct {
Count int // 允许次数
Window time.Duration // 时间窗口
}
// clientIP 解析客户端真实 IP:仅当直连地址属于可信代理时才采信 X-Forwarded-For / X-Real-IP。
// 语义对齐参考 apps/base/dependencies.py 的 get_client_ip。
func clientIP(c *gin.Context, trustedProxies []*net.IPNet) string {
remote := net.ParseIP(c.RemoteIP())
parse := func(s string) net.IP {
ip := net.ParseIP(strings.TrimSpace(s))
return ip
}
isTrusted := func(ip net.IP) bool {
if ip == nil {
return false
}
for _, n := range trustedProxies {
if n.Contains(ip) {
return true
}
}
return false
}
if !isTrusted(remote) {
return remote.String()
}
// X-Forwarded-For:从右往左找第一个非可信代理地址
if xff := c.GetHeader("X-Forwarded-For"); xff != "" {
parts := strings.Split(xff, ",")
for i := len(parts) - 1; i >= 0; i-- {
candidate := parse(parts[i])
if candidate == nil {
return remote.String()
}
if !isTrusted(candidate) {
return candidate.String()
}
}
return strings.TrimSpace(parts[0])
}
if xr := c.GetHeader("X-Real-IP"); xr != "" {
if ip := parse(xr); ip != nil {
return ip.String()
}
}
return remote.String()
}
// ParseTrustedProxies 把 CIDR/单 IP 字符串解析为网络列表。
func ParseTrustedProxies(items []string) []*net.IPNet {
var out []*net.IPNet
for _, item := range items {
item = strings.TrimSpace(item)
if item == "" {
continue
}
if !strings.Contains(item, "/") {
item += "/32"
if strings.Contains(item, ":") { // IPv6
item = item[:len(item)-3] + "/128"
}
}
_, network, err := net.ParseCIDR(item)
if err != nil {
continue
}
out = append(out, network)
}
return out
}
// ClientIP 中间件:解析真实 IP 并写入上下文(ctxClientIP)。
func ClientIP(trustedProxies []*net.IPNet) gin.HandlerFunc {
return func(c *gin.Context) {
c.Set("ctxClientIP", clientIP(c, trustedProxies))
c.Next()
}
}
// GetClientIP 从 gin 上下文取解析后的客户端 IP。
func GetClientIP(c *gin.Context) string {
if v, ok := c.Get("ctxClientIP"); ok {
if s, ok := v.(string); ok {
return s
}
}
if ip := net.ParseIP(c.RemoteIP()); ip != nil {
return ip.String()
}
return c.RemoteIP()
}
// RateLimiter 基于 cache.Cache 的固定窗口 IP 限流器。
// 计数语义对齐参考实现:check 通过时放行,业务方在发生"计数事件"(如失败/成功上传)后调用 Add。
//
// L9:缓存故障降级——此前 cache.Get/Incr 失败(如 Redis 宕机)时一律放行,
// 登录爆破防护随之失效。现降级为进程内固定窗口计数(单实例语义),
// 缓存恢复后自动回到共享缓存计数。降级期间计数独立于缓存,不叠加。
type RateLimiter struct {
cache cache.Cache
limits map[string]LimitRule
prefix string
fbMu sync.Mutex
fallback map[string]*fallbackEntry // 进程内降级计数
lastPrune time.Time
}
type fallbackEntry struct {
count int64
expires time.Time
}
// fallbackMaxEntries 降级计数表上限(超出即整体重置,防内存增长)。
const fallbackMaxEntries = 8192
// NewRateLimiter 构造限流器;limits 为各类别规则(来自 settings 的 errorCount/errorMinute 等)。
func NewRateLimiter(cache cache.Cache, limits map[string]LimitRule) *RateLimiter {
if limits == nil {
limits = map[string]LimitRule{}
}
return &RateLimiter{cache: cache, limits: limits, prefix: "fcb:rl", fallback: map[string]*fallbackEntry{}}
}
// SetRule 运行时更新规则(settings KV 变更后调用)。
func (r *RateLimiter) SetRule(kind string, rule LimitRule) {
r.limits[kind] = rule
}
func (r *RateLimiter) windowKey(kind, ip string, now time.Time) string {
// 固定窗口:按窗口起点分桶
bucket := now.Unix() / int64(r.limits[kind].Window/time.Second)
return r.prefix + ":" + kind + ":" + ip + ":" + itoa64(bucket)
}
// Check 只读检查该 IP 在当前窗口内是否仍被允许(不计数)。
// 对齐参考 check_ip:已用次数 >= 上限即拒绝。
// 缓存键不存在(ErrNotFound)视为 0 次;缓存故障时降级为进程内计数。
func (r *RateLimiter) Check(c *gin.Context, kind string) (bool, int64) {
rule, ok := r.limits[kind]
if !ok || rule.Count <= 0 || rule.Window <= 0 {
return true, 0
}
ip := GetClientIP(c)
ctx, cancel := context.WithTimeout(c.Request.Context(), 2*time.Second)
defer cancel()
now := time.Now()
raw, err := r.cache.Get(ctx, r.windowKey(kind, ip, now))
if err != nil {
if errors.Is(err, cache.ErrNotFound) {
return true, 0 // 键不存在:窗口内尚无计数
}
// 缓存故障:降级进程内计数判定
return r.fallbackCount(kind, ip, rule, now) < int64(rule.Count), 0
}
n := parseInt64(raw)
return n < int64(rule.Count), n
}
// Add 记录一次计数事件(对齐参考 add_ip:调用即计数,如上传成功/登录失败/取件错误)。
// 缓存故障时降级为进程内计数。
func (r *RateLimiter) Add(c *gin.Context, kind string) {
rule, ok := r.limits[kind]
if !ok || rule.Count <= 0 || rule.Window <= 0 {
return
}
ip := GetClientIP(c)
ctx, cancel := context.WithTimeout(c.Request.Context(), 2*time.Second)
defer cancel()
if _, err := r.cache.Incr(ctx, r.windowKey(kind, ip, time.Now()), rule.Window); err != nil &&
!errors.Is(err, cache.ErrNotFound) {
// Incr 正常情况下不会因键不存在失败(缺键即从 0 起);
// 其余错误视为缓存故障 → 进程内计数
r.fallbackIncr(kind, ip, rule, time.Now())
}
}
// —— 进程内降级计数(L9)——
func (r *RateLimiter) fallbackIncr(kind, ip string, rule LimitRule, now time.Time) {
key := r.windowKey(kind, ip, now)
expires := now.Add(rule.Window)
r.fbMu.Lock()
defer r.fbMu.Unlock()
r.pruneFallbackLocked(now)
if len(r.fallback) >= fallbackMaxEntries {
r.fallback = map[string]*fallbackEntry{} // 极端情况整体重置,防内存无限增长
}
e, ok := r.fallback[key]
if !ok || now.After(e.expires) {
r.fallback[key] = &fallbackEntry{count: 1, expires: expires}
return
}
e.count++
}
func (r *RateLimiter) fallbackCount(kind, ip string, rule LimitRule, now time.Time) int64 {
key := r.windowKey(kind, ip, now)
r.fbMu.Lock()
defer r.fbMu.Unlock()
e, ok := r.fallback[key]
if !ok || now.After(e.expires) {
return 0
}
return e.count
}
// pruneFallbackLocked 清理已过窗口的降级计数(低频触发:每 1000 条或 10 分钟一次)。
func (r *RateLimiter) pruneFallbackLocked(now time.Time) {
if r.lastPrune.IsZero() || len(r.fallback) >= 1024 || now.Sub(r.lastPrune) >= 10*time.Minute {
for k, e := range r.fallback {
if now.After(e.expires) {
delete(r.fallback, k)
}
}
r.lastPrune = now
}
}
// RequireRateLimit 中间件:请求进入即检查,请求完成即计数。
// 适用于"每次访问都计数"的类别(如 metadata 查询);
// 上传/登录等"仅成功/失败才计数"的场景由 handler 显式调用 Check/Add。
func (r *RateLimiter) RequireRateLimit(kind string) gin.HandlerFunc {
return func(c *gin.Context) {
allowed, _ := r.Check(c, kind)
if !allowed {
response.Fail(c, http.StatusLocked, "请求次数过多,请稍后再试")
return
}
c.Next()
r.Add(c, kind)
}
}
// parseInt64 解析十进制整数字符串,非法输入返回 0。
func parseInt64(s string) int64 {
var n int64
for _, ch := range s {
if ch < '0' || ch > '9' {
return 0
}
n = n*10 + int64(ch-'0')
}
return n
}
func itoa64(n int64) string {
if n == 0 {
return "0"
}
neg := n < 0
if neg {
n = -n
}
var buf [21]byte
i := len(buf)
for n > 0 {
i--
buf[i] = byte('0' + n%10)
n /= 10
}
if neg {
i--
buf[i] = '-'
}
return string(buf[i:])
}
@@ -0,0 +1,66 @@
package middleware
import (
"context"
"net/http/httptest"
"testing"
"time"
"github.com/gin-gonic/gin"
"filecodebox/internal/cache"
)
// failingCache 模拟缓存故障(Get/Incr 均返回非 ErrNotFound 错误)。
type failingCache struct{ cache.Cache }
func (f *failingCache) Get(_ context.Context, _ string) (string, error) {
return "", context.DeadlineExceeded
}
func (f *failingCache) Set(_ context.Context, _, _ string, _ time.Duration) error {
return context.DeadlineExceeded
}
func (f *failingCache) Incr(_ context.Context, _ string, _ time.Duration) (int64, error) {
return 0, context.DeadlineExceeded
}
func newLimiterTest(c *gin.Context, cacheImpl cache.Cache, count int) *RateLimiter {
return NewRateLimiter(cacheImpl, map[string]LimitRule{
LimitLogin: {Count: count, Window: time.Minute},
})
}
func ginTestContext() *gin.Context {
gin.SetMode(gin.TestMode)
c, _ := gin.CreateTestContext(httptest.NewRecorder())
c.Request = httptest.NewRequest("POST", "/admin/login", nil)
return c
}
// TestRateLimiterFallbackOnCacheFailure L9:缓存故障时限流降级为进程内计数,
// 超过上限后 Check 拒绝(此前 fail-open 会一直放行)。
func TestRateLimiterFallbackOnCacheFailure(t *testing.T) {
c := ginTestContext()
rl := newLimiterTest(c, &failingCache{cache.NewMemory()}, 3)
for i := 0; i < 3; i++ {
if ok, _ := rl.Check(c, LimitLogin); !ok {
t.Fatalf("第 %d 次检查不应拒绝", i+1)
}
rl.Add(c, LimitLogin)
}
if ok, _ := rl.Check(c, LimitLogin); ok {
t.Fatal("缓存故障降级下,超过上限后 Check 应拒绝(fail-close")
}
}
// TestRateLimiterNormalCacheCounting 正常缓存路径行为不变。
func TestRateLimiterNormalCacheCounting(t *testing.T) {
c := ginTestContext()
rl := newLimiterTest(c, cache.NewMemory(), 2)
for i := 0; i < 2; i++ {
rl.Add(c, LimitLogin)
}
if ok, _ := rl.Check(c, LimitLogin); ok {
t.Fatal("达到上限后 Check 应拒绝")
}
}
@@ -0,0 +1,126 @@
package middleware
import (
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/gin-gonic/gin"
"filecodebox/internal/cache"
)
func rateLimitRouter(rl *RateLimiter, kind string) *gin.Engine {
r := gin.New()
r.Use(ClientIP(nil))
r.GET("/limited", rl.RequireRateLimit(kind), func(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"ok": true})
})
return r
}
func TestRateLimiterBlocksAfterCount(t *testing.T) {
mem := cache.NewMemory()
defer mem.Close()
rl := NewRateLimiter(mem, map[string]LimitRule{
LimitMeta: {Count: 3, Window: time.Minute},
})
r := rateLimitRouter(rl, LimitMeta)
// 前 3 次通过
for i := 0; i < 3; i++ {
w := httptest.NewRecorder()
r.ServeHTTP(w, httptest.NewRequest("GET", "/limited", nil))
if w.Code != http.StatusOK {
t.Fatalf("第 %d 次应通过: %d", i+1, w.Code)
}
}
// 第 4 次 423
w := httptest.NewRecorder()
r.ServeHTTP(w, httptest.NewRequest("GET", "/limited", nil))
if w.Code != http.StatusLocked {
t.Fatalf("超限应 423: %d", w.Code)
}
}
func TestRateLimiterAddAfterSuccess(t *testing.T) {
// 模拟 upload 语义:Check 放行 + handler 成功后 Add
mem := cache.NewMemory()
defer mem.Close()
rl := NewRateLimiter(mem, map[string]LimitRule{
LimitUpload: {Count: 2, Window: time.Minute},
})
r := gin.New()
r.Use(ClientIP(nil))
r.POST("/upload", func(c *gin.Context) {
if allowed, _ := rl.Check(c, LimitUpload); !allowed {
c.JSON(http.StatusLocked, gin.H{"err": "too many"})
return
}
rl.Add(c, LimitUpload) // 成功上传计数
c.JSON(http.StatusOK, gin.H{"ok": true})
})
for i := 0; i < 2; i++ {
w := httptest.NewRecorder()
r.ServeHTTP(w, httptest.NewRequest("POST", "/upload", nil))
if w.Code != http.StatusOK {
t.Fatalf("第 %d 次上传应通过", i+1)
}
}
w := httptest.NewRecorder()
r.ServeHTTP(w, httptest.NewRequest("POST", "/upload", nil))
if w.Code != http.StatusLocked {
t.Fatalf("第 3 次上传应被拒绝: %d", w.Code)
}
}
func TestParseTrustedProxies(t *testing.T) {
nets := ParseTrustedProxies([]string{"10.0.0.0/8", "192.168.1.1", "", "bad-input"})
if len(nets) != 2 {
t.Fatalf("应解析出 2 个可信网段: %d", len(nets))
}
}
func TestClientIPFromTrustedProxy(t *testing.T) {
// 对齐参考语义:仅当直连地址可信时才解析 XFF,
// 且从右往左返回第一个"非可信代理"地址(该地址即最近可信代理看到的客户端)。
nets := ParseTrustedProxies([]string{"127.0.0.0/8", "10.0.0.0/8"})
r := gin.New()
r.Use(ClientIP(nets))
var seen string
r.GET("/ip", func(c *gin.Context) { seen = GetClientIP(c) })
req := httptest.NewRequest("GET", "/ip", nil)
req.RemoteAddr = "127.0.0.1:5000"
req.Header.Set("X-Forwarded-For", "203.0.113.9, 10.0.0.1")
r.ServeHTTP(httptest.NewRecorder(), req)
if seen != "203.0.113.9" {
t.Fatalf("多级可信代理下应取最左非可信地址: %s", seen)
}
// 仅直连可信:XFF 右起第一个(10.0.0.1)非可信 → 取它
nets2 := ParseTrustedProxies([]string{"127.0.0.0/8"})
r2 := gin.New()
r2.Use(ClientIP(nets2))
var seen2 string
r2.GET("/ip", func(c *gin.Context) { seen2 = GetClientIP(c) })
req2 := httptest.NewRequest("GET", "/ip", nil)
req2.RemoteAddr = "127.0.0.1:5000"
req2.Header.Set("X-Forwarded-For", "203.0.113.9, 10.0.0.1")
r2.ServeHTTP(httptest.NewRecorder(), req2)
if seen2 != "10.0.0.1" {
t.Fatalf("右起第一个非可信地址应为 10.0.0.1: %s", seen2)
}
// 非可信直连:忽略伪造头
seen = ""
req = httptest.NewRequest("GET", "/ip", nil)
req.RemoteAddr = "8.8.8.8:1234"
req.Header.Set("X-Forwarded-For", "1.2.3.4")
r.ServeHTTP(httptest.NewRecorder(), req)
if seen != "8.8.8.8" {
t.Fatalf("非可信直连应忽略 XFF: %s", seen)
}
}
+152
View File
@@ -0,0 +1,152 @@
// Package model 定义 GORM 数据模型与 Postgres 自动迁移。
// 字段对齐参考实现 apps/base/models.py,并新增审计日志表。
package model
import (
"time"
"gorm.io/gorm"
)
// FileCodes 文件/文本分享记录(对齐参考 FileCodes)。
type FileCodes struct {
ID int64 `gorm:"primaryKey;autoIncrement" json:"id"`
Code string `gorm:"column:code;size:255;uniqueIndex;not null" json:"code"` // 取件码
Prefix string `gorm:"size:255;default:''" json:"prefix"` // 文件名前缀/文本分享标记
Suffix string `gorm:"size:255;default:''" json:"suffix"` // 文件名后缀(含扩展名)
UUIDFileName *string `gorm:"size:255" json:"uuid_file_name"` // 存储侧 UUID 文件名
FilePath *string `gorm:"size:255" json:"file_path"` // 存储侧相对路径
Size int64 `gorm:"default:0" json:"size"` // 字节数;文本为字符数
Text *string `gorm:"type:text" json:"text"` // 文本分享内容
ExpiredAt *time.Time `json:"expired_at"` // 过期时间;永久分享为 NULL
ExpiredCount int `gorm:"default:0" json:"expired_count"` // 剩余可取次数;<0 表示按时间过期
UsedCount int `gorm:"default:0" json:"used_count"` // 已取次数
CreatedAt time.Time `json:"created_at"`
FileHash *string `gorm:"size:64" json:"file_hash"` // SHA256
IsChunked bool `gorm:"default:false" json:"is_chunked"`
UploadID *string `gorm:"size:36" json:"upload_id"` // 分片上传会话 ID
Engine string `gorm:"size:16;default:''" json:"engine"` // 归属存储引擎(v3local|s3|webdav;空=历史数据按当前引擎取)
}
// TableName 表名。
func (FileCodes) TableName() string { return "file_codes" }
// Expired 判断是否已过期(对齐参考语义:expired_count<0 按时间,否则按次数)。
func (f *FileCodes) Expired(now time.Time) bool {
if f.ExpiredAt == nil {
return false
}
if f.ExpiredCount < 0 {
return f.ExpiredAt.Before(now)
}
return f.ExpiredCount <= 0
}
// UploadChunk 分片上传记录(对齐参考 UploadChunk)。
type UploadChunk struct {
ID int64 `gorm:"primaryKey;autoIncrement" json:"id"`
UploadID string `gorm:"size:36;index:idx_upload_chunk,unique,priority:1;not null" json:"upload_id"`
ChunkIndex int `gorm:"index:idx_upload_chunk,unique,priority:2;not null" json:"chunk_index"`
ChunkHash string `gorm:"size:64;not null" json:"chunk_hash"` // 分片 SHA256
TotalChunks int `json:"total_chunks"`
FileSize int64 `json:"file_size"`
ChunkSize int `json:"chunk_size"`
FileName string `gorm:"size:255" json:"file_name"`
SavePath string `gorm:"size:512" json:"save_path"`
CreatedAt time.Time `json:"created_at"`
Completed bool `gorm:"default:false" json:"completed"`
Engine string `gorm:"size:16;default:''" json:"engine"` // 归属存储引擎(v3:分片与会话记录当时引擎,合并走同一引擎)
}
// TableName 表名。
func (UploadChunk) TableName() string { return "upload_chunks" }
// KeyValue 运行时配置键值(对齐参考 KeyValue)。value 存 JSON。
type KeyValue struct {
ID int64 `gorm:"primaryKey;autoIncrement" json:"id"`
Key string `gorm:"size:255;uniqueIndex;not null" json:"key"`
Value *string `gorm:"type:text" json:"value"` // JSON 字符串
CreatedAt time.Time `json:"created_at"`
}
// TableName 表名。
func (KeyValue) TableName() string { return "key_values" }
// PresignUploadSession 预签名直传会话(对齐参考 PresignUploadSession)。
type PresignUploadSession struct {
ID int64 `gorm:"primaryKey;autoIncrement" json:"id"`
UploadID string `gorm:"size:36;uniqueIndex;not null" json:"upload_id"`
FileName string `gorm:"size:255" json:"file_name"`
FileSize int64 `json:"file_size"`
SavePath string `gorm:"size:512" json:"save_path"`
Mode string `gorm:"size:10" json:"mode"` // direct=客户端直传 | proxy=服务器代理
ExpireValue int `json:"expire_value"`
ExpireStyle string `gorm:"size:20;default:day" json:"expire_style"`
CreatedAt time.Time `json:"created_at"`
ExpiresAt time.Time `json:"expires_at"`
Engine string `gorm:"size:16;default:''" json:"engine"` // 归属存储引擎(v3:直传/代理完成走同一引擎取回)
}
// TableName 表名。
func (PresignUploadSession) TableName() string { return "presign_upload_sessions" }
// IsExpired 会话是否已过期。
func (p *PresignUploadSession) IsExpired(now time.Time) bool { return p.ExpiresAt.Before(now) }
// StorageReservation 上传容量预留(尚未写入 file_codes 的占位)。
type StorageReservation struct {
ID int64 `gorm:"primaryKey;autoIncrement" json:"id"`
Token string `gorm:"size:64;uniqueIndex;not null" json:"token"`
Size int64 `json:"size"`
ExpiresAt time.Time `gorm:"index" json:"expires_at"`
}
// TableName 表名。
func (StorageReservation) TableName() string { return "storage_reservations" }
// 审计结果常量。
const (
AuditResultSuccess = "success" // 操作成功
AuditResultDenied = "denied" // 被拒绝(限流/鉴权/策略)
AuditResultFailed = "failed" // 执行失败(服务端/客户端错误)
)
// AuditLog 上传/下载审计日志(需求 ③)。
type AuditLog struct {
ID int64 `gorm:"primaryKey;autoIncrement" json:"id"`
Action string `gorm:"size:32;index" json:"action"` // upload | download
FileCode string `gorm:"size:64;index" json:"file_code"` // 取件码(上传时为生成的码)
FileName string `gorm:"size:255" json:"file_name"` // 原始文件名/文本标记
SizeBytes int64 `json:"size_bytes"` // 文件总字节数
TransferredBytes int64 `json:"transferred_bytes"` // 本次实际传输字节数
IP string `gorm:"size:64;index" json:"ip"`
UserAgent string `gorm:"size:512" json:"user_agent"`
DeviceOS string `gorm:"size:64" json:"device_os"` // Windows/macOS/Android/iOS/Linux/Unknown
DeviceBrowser string `gorm:"size:64" json:"device_browser"` // Chrome/Firefox/Safari/Edge/...
DeviceType string `gorm:"size:32" json:"device_type"` // desktop/mobile/tablet/bot/other
Actor string `gorm:"size:64" json:"actor"` // admin | guest
Result string `gorm:"size:16;index" json:"result"` // success | denied | failed
ErrorMsg string `gorm:"size:512" json:"error_msg"`
DurationMs int64 `json:"duration_ms"`
CreatedAt time.Time `gorm:"index" json:"created_at"` // 操作时间
}
// TableName 表名。
func (AuditLog) TableName() string { return "audit_logs" }
// AllModels 全部需要迁移的模型。
func AllModels() []any {
return []any{
&FileCodes{},
&UploadChunk{},
&KeyValue{},
&PresignUploadSession{},
&StorageReservation{},
&AuditLog{},
}
}
// AutoMigrate 在 Postgres 上建表/补列;服务启动时调用。
func AutoMigrate(db *gorm.DB) error {
return db.AutoMigrate(AllModels()...)
}
+25
View File
@@ -0,0 +1,25 @@
// Package response 提供统一响应封装:{"code":200,"msg":"...","data":...}。
package response
import (
"net/http"
"github.com/gin-gonic/gin"
)
// Body 统一响应体。
type Body struct {
Code int `json:"code"`
Msg string `json:"msg"`
Data any `json:"data,omitempty"`
}
// OK 成功响应(code=200)。
func OK(c *gin.Context, data any) {
c.JSON(http.StatusOK, Body{Code: 200, Msg: "ok", Data: data})
}
// Fail 失败响应,httpStatus 与 code 语义一致(404 过期/不存在、403 拒绝、429 限流、500 服务端错误)。
func Fail(c *gin.Context, httpStatus int, msg string) {
c.AbortWithStatusJSON(httpStatus, Body{Code: httpStatus, Msg: msg})
}
+150
View File
@@ -0,0 +1,150 @@
// settings 包双方言测试:Manager 全流程(ensure 行、KV 读写合并、Reload、
// UpdateKV 屏蔽内部键、SystemStart)分别在 sqlite(默认)与 postgresFCB_TEST_PG_DSN)上执行。
package settings_test
import (
"context"
"os"
"path/filepath"
"testing"
"gorm.io/gorm"
"filecodebox/internal/config"
"filecodebox/internal/database"
"filecodebox/internal/settings"
)
// pgTestDSN 返回 Postgres 测试连接串;未设置 FCB_TEST_PG_DSN 时跳过。
func pgTestDSN(t *testing.T) string {
t.Helper()
dsn := os.Getenv("FCB_TEST_PG_DSN")
if dsn == "" {
t.Skip("未设置 FCB_TEST_PG_DSN,跳过 Postgres 双方言用例")
}
return dsn
}
// newTestManager 按方言构造库 + Manager(已完成 Migrate)。
func newTestManager(t *testing.T, driver, dsn string) (*settings.Manager, *gorm.DB, func()) {
t.Helper()
if dsn == "" {
dsn = filepath.Join(t.TempDir(), "settings-test.db")
}
ctx := context.Background()
db, err := database.Open(ctx, database.Options{Driver: driver, DSN: dsn})
if err != nil {
t.Fatalf("[%s] Open: %v", driver, err)
}
if err := database.Migrate(ctx, db); err != nil {
_ = database.Close(db)
t.Fatalf("[%s] Migrate: %v", driver, err)
}
// Config 仅作内存载体(驱动不回连数据库),postgres 模式给占位 DSN 以通过校验
t.Setenv("FCB_DB_DRIVER", driver)
if driver == "postgres" {
t.Setenv("FCB_DB_DSN", dsn)
} else {
t.Setenv("FCB_DB_DSN", "")
}
cfg, err := config.New()
if err != nil {
_ = database.Close(db)
t.Fatalf("[%s] config.New: %v", driver, err)
}
mgr, err := settings.NewManager(ctx, db, cfg)
if err != nil {
_ = database.Close(db)
t.Fatalf("[%s] NewManager: %v", driver, err)
}
return mgr, db, func() { _ = database.Close(db) }
}
// runManagerSuite 双方言共用的 Manager 行为断言。
func runManagerSuite(t *testing.T, mgr *settings.Manager) {
t.Helper()
ctx := context.Background()
// 1. 初始:未初始化(admin_token 空)
if mgr.IsInitialized() {
t.Fatal("初始 admin_token 为空应视为未初始化")
}
// 2. UpdateKV 写入策略键 → Reload 后读取生效
patch := map[string]any{
settings.KeyBackgroundURL: "https://example.com/bg.png",
settings.KeyFooterText: "自建部署,仅供内部演示",
settings.KeyFooterBeian: "京ICP备2024000001号-1",
settings.KeyNotifyEnabled: 0,
settings.KeyMaxSaveSeconds: 86400,
"_internal_secret": "must-drop", // 下划线内部键必须被拒
}
if err := mgr.UpdateKV(ctx, patch); err != nil {
t.Fatalf("UpdateKV 失败: %v", err)
}
if err := mgr.Reload(ctx); err != nil {
t.Fatalf("Reload 失败: %v", err)
}
cfg := mgr.Get()
if got := cfg.GetString(settings.KeyBackgroundURL); got != "https://example.com/bg.png" {
t.Fatalf("background_url 未生效: %q", got)
}
if got := cfg.GetString(settings.KeyFooterBeian); got != "京ICP备2024000001号-1" {
t.Fatalf("footer_beian 未生效: %q", got)
}
if cfg.GetBool(settings.KeyNotifyEnabled) {
t.Fatal("notify_enabled=0 应生效")
}
if got := cfg.MaxSaveSeconds(); got != 86400 {
t.Fatalf("max_save_seconds 未生效: %d", got)
}
if _, ok := cfg.Get("_internal_secret"); ok {
t.Fatal("下划线内部键不应进入运行时配置")
}
// 3. KV 合并语义:二次 UpdateKV 不覆盖未提及键
if err := mgr.UpdateKV(ctx, map[string]any{settings.KeyNotifyEnabled: 1}); err != nil {
t.Fatalf("二次 UpdateKV: %v", err)
}
if err := mgr.Reload(ctx); err != nil {
t.Fatalf("二次 Reload: %v", err)
}
cfg = mgr.Get()
if !cfg.GetBool(settings.KeyNotifyEnabled) {
t.Fatal("notify_enabled 二次写入应生效")
}
if got := cfg.GetString(settings.KeyFooterText); got == "" {
t.Fatal("二次写入不应清空 footer_text")
}
// 4. SystemStartsys_start 键写入且为毫秒时间戳
mgr.SystemStart(ctx)
// 5. 敏感键判定(双模式一致)
if !settings.IsSensitiveKey("admin_token") || !settings.IsSensitiveKey("jwt_secret") {
t.Fatal("admin_token/jwt_secret 应为敏感键")
}
if settings.IsSensitiveKey("footer_text") {
t.Fatal("footer_text 不应为敏感键")
}
// 6. KV schema 表完整性:全部键可从默认值读取
for _, e := range settings.KVSchema() {
if _, ok := cfg.Get(e.Key); !ok {
t.Fatalf("schema 键 %q 在默认配置中不存在", e.Key)
}
}
}
func TestManagerSQLite(t *testing.T) {
mgr, _, closeFn := newTestManager(t, "sqlite", "")
defer closeFn()
runManagerSuite(t, mgr)
}
func TestManagerPostgres(t *testing.T) {
dsn := pgTestDSN(t)
mgr, _, closeFn := newTestManager(t, "postgres", dsn)
defer closeFn()
runManagerSuite(t, mgr)
}
+98
View File
@@ -0,0 +1,98 @@
// Package settings 密码哈希与校验:
// 新密码使用 bcrypt(格式 bcrypt$<bcrypt原生哈希串>);同时兼容两代旧格式——
// sha256$salt$hash(上一版)与旧版明文(迁移校验)。
// 安全审计 M1:单轮 SHA256+盐抗 GPU 爆破不足,新哈希统一升级 bcrypt。
// 兼容策略:VerifyPassword 支持全部三代格式;调用方可用 NeedsRehash 判定
// 登录成功后是否需要用新算法重哈希写回(登录升级路径见 api.adminLogin)。
package settings
import (
"crypto/hmac"
"crypto/rand"
"crypto/sha256"
"encoding/hex"
"strconv"
"strings"
"golang.org/x/crypto/bcrypt"
)
// bcryptCost bcrypt 工作因子:122026 年桌面 CPU 单次校验约 100-250ms
// 离线爆破成本相比单轮 SHA256 提升数个数量级)。
const bcryptCost = 12
// bcryptMaxLen bcrypt 算法只取前 72 字节;超长输入统一截断,
// 避免 GenerateFromPassword/CompareHashAndPassword 对 >72 字节返回错误。
const bcryptMaxLen = 72
func bcryptBytes(password string) []byte {
b := []byte(password)
if len(b) > bcryptMaxLen {
b = b[:bcryptMaxLen]
}
return b
}
// HashPassword 生成 bcrypt$<hash> 格式密码哈希(<hash> 为 bcrypt 原生串
// `$2a$<cost>$<salt><hash>`cost 内嵌于哈希串中)。
func HashPassword(password string) string {
sum, err := bcrypt.GenerateFromPassword(bcryptBytes(password), bcryptCost)
if err != nil {
// 截断后仅剩非法 cost 等实现级错误:确定性失败优于弱哈希回落
panic("settings: bcrypt 哈希失败: " + err.Error())
}
return "bcrypt$" + string(sum)
}
// VerifyPassword 校验密码:支持 bcrypt$、sha256$salt$hash 与旧版明文三种格式。
func VerifyPassword(password, hashed string) bool {
if hashed == "" {
return false
}
switch {
case strings.HasPrefix(hashed, "bcrypt$"):
return bcrypt.CompareHashAndPassword([]byte(hashed[len("bcrypt$"):]), bcryptBytes(password)) == nil
case strings.HasPrefix(hashed, "sha256$"):
parts := strings.Split(hashed, "$")
if len(parts) != 3 {
return false
}
salt, stored := parts[1], parts[2]
sum := sha256.Sum256([]byte(salt + password))
return hmac.Equal([]byte(hex.EncodeToString(sum[:])), []byte(stored))
}
// 旧版明文比较(兼容迁移)
return hmac.Equal([]byte(password), []byte(hashed))
}
// NeedsRehash 判断哈希是否需要升级为当前算法/成本(登录成功后判定,透明迁移)。
// sha256 与明文一律 truebcrypt 成本低于当前 bcryptCost 时 true。
func NeedsRehash(hashed string) bool {
if !strings.HasPrefix(hashed, "bcrypt$") {
return true
}
// bcrypt 原生串格式:$2a$<cost>$<salt><hash>
parts := strings.Split(hashed[len("bcrypt$"):], "$")
if len(parts) < 4 {
return true
}
cost, err := strconv.Atoi(parts[2])
if err != nil {
return true
}
return cost < bcryptCost
}
// IsPasswordHashed 判断是否为受支持的哈希格式(bcrypt / sha256)。
func IsPasswordHashed(s string) bool {
return strings.HasPrefix(s, "bcrypt$") || strings.HasPrefix(s, "sha256$")
}
// GenerateJWTSecret 生成 64 字符十六进制随机密钥。
func GenerateJWTSecret() string {
b := make([]byte, 32)
if _, err := rand.Read(b); err != nil {
panic("settings: crypto/rand 不可用: " + err.Error())
}
return hex.EncodeToString(b)
}
+38
View File
@@ -0,0 +1,38 @@
package settings
import "testing"
func TestHashPasswordRoundTrip(t *testing.T) {
h := HashPassword("s3cret-密码")
if !IsPasswordHashed(h) {
t.Fatalf("哈希格式不对: %s", h)
}
if !VerifyPassword("s3cret-密码", h) {
t.Fatal("正确密码校验失败")
}
if VerifyPassword("wrong", h) {
t.Fatal("错误密码竟通过校验")
}
if h == HashPassword("s3cret-密码") {
t.Fatal("盐值未随机化")
}
}
func TestVerifyLegacyPlaintext(t *testing.T) {
if !VerifyPassword("FileCodeBox2023", "FileCodeBox2023") {
t.Fatal("旧版明文兼容校验失败")
}
if VerifyPassword("nope", "FileCodeBox2023") {
t.Fatal("明文比较不应放行其他密码")
}
}
func TestGenerateJWTSecretLength(t *testing.T) {
s := GenerateJWTSecret()
if len(s) < 32 {
t.Fatalf("密钥太短: %d", len(s))
}
if s == GenerateJWTSecret() {
t.Fatal("密钥未随机化")
}
}
+258
View File
@@ -0,0 +1,258 @@
// Package settings — sanitize.go:受控 HTML 白名单净化(安全审计 L7)。
//
// notify_content 设计上「允许 <a> 等受控 HTML」,此前由管理端任意写入并经
// 前端 v-html 直出——管理员账号一旦被盗即可对全站访客注入脚本。
// 本净化器只保留纯文本与 <a href="http(s)|/|#">,其余标签连同其内层内容
// 一并丢弃(不做 HTML 转义输出,避免脚本字面量进入页面 DOM),
// 在公开配置读取与保存两处调用(双保险,覆盖历史存量数据)。
package settings
import (
"strings"
)
// dropContentTags 标签内部内容也一并丢弃的危险标签(script/style 等)。
var dropContentTags = map[string]bool{
"script": true, "style": true, "iframe": true, "object": true, "embed": true,
"title": true, "textarea": true, "noscript": true, "template": true,
"svg": true, "math": true, "xmp": true, "noembed": true, "noframes": true,
}
// SanitizeInlineHTML 白名单净化内联 HTML
// - <script>/<style>/<iframe> 等危险标签连同内部内容整体丢弃;
// - 其他非 <a> 标签仅丢弃标签本身、保留其内层文本(如 <b>加粗</b> → 加粗);
// - <a> 仅保留 href 属性,且值必须以 http://、https://、/ 或 # 开头;
// - HTML 注释(<!-- -->)丢弃,未闭合的危险标签丢弃其后全部内容;
// - 文本片段原样保留(不含 '<',渲染时为安全文本节点)。
func SanitizeInlineHTML(input string) string {
if input == "" {
return ""
}
var b strings.Builder
b.Grow(len(input))
i := 0
pendingAnchor := false
writeClose := func() {
if pendingAnchor {
b.WriteString("</a>")
pendingAnchor = false
}
}
for i < len(input) {
lt := strings.IndexByte(input[i:], '<')
if lt < 0 {
b.WriteString(input[i:])
break
}
b.WriteString(input[i : i+lt])
rest := input[i+lt:]
// 注释:整体丢弃
if strings.HasPrefix(rest, "<!--") {
end := strings.Index(rest, "-->")
if end < 0 {
break // 未闭合注释:丢弃剩余全部
}
i += lt + end + 3
continue
}
end := findTagEnd(rest)
if end < 0 {
break // 未闭合标签:丢弃剩余全部(不当作文本,防 < 绕过)
}
rawTag := rest[:end+1] // 形如 "<a href=..>"、"</div>"、"<img .../>"
name, closing, _ := parseTagName(rawTag)
if name != "" && !closing {
if dropContentTags[name] {
// 危险标签:连内层跳到对应闭合标签;无闭合(如 <script> 到结尾)则全丢
closeIdx := findClosingTag(input, i+lt+end+1, name)
if closeIdx < 0 {
writeClose()
return b.String()
}
i = closeIdx
continue
}
if name == "a" {
writeClose()
if href, ok := parseAllowedAnchor(rawTag); ok {
b.WriteString(`<a href="` + escapeAttr(href) + `">`)
pendingAnchor = true
}
// href 非法的 <a>:标签丢弃,但内层文本仍保留
}
// 其余开标签:丢弃标签本身,保留内层文本
i += lt + end + 1
continue
}
if name != "" && closing && name == "a" {
writeClose() // 仅在存在未闭合的合法 <a> 时输出
}
// 其余闭标签:丢弃
i += lt + end + 1
}
writeClose()
return b.String()
}
// findTagEnd 返回标签结束 '>' 的下标(跳过引号内的 '>',如 href="a<b">);未找到返回 -1。
func findTagEnd(s string) int {
inQuote := byte(0)
for i := 0; i < len(s); i++ {
c := s[i]
if inQuote != 0 {
if c == inQuote {
inQuote = 0
}
continue
}
switch c {
case '"', '\'':
inQuote = c
case '>':
return i
}
}
return -1
}
// parseTagName 解析标签名:返回 (小写名, 是否闭合标签, 是否自闭合 "/>")。
func parseTagName(tag string) (name string, closing, selfClosing bool) {
if len(tag) < 3 || tag[0] != '<' || tag[len(tag)-1] != '>' {
return "", false, false
}
inner := tag[1 : len(tag)-1]
if strings.HasSuffix(inner, "/") {
selfClosing = true
inner = inner[:len(inner)-1]
}
if strings.HasPrefix(inner, "/") {
closing = true
inner = inner[1:]
}
end := 0
for end < len(inner) {
r := inner[end]
if r == ' ' || r == '\t' || r == '\n' || r == '\r' {
break
}
end++
}
if end == 0 {
return "", closing, selfClosing
}
return strings.ToLower(inner[:end]), closing, selfClosing
}
// findClosingTag 从 from 开始查找 </name>,返回闭合标签结束位置(不含);找不到返回 -1。
func findClosingTag(s string, from int, name string) int {
needle := "</" + name
lower := strings.ToLower(s)
pos := from
for {
idx := strings.Index(lower[pos:], needle)
if idx < 0 {
return -1
}
at := pos + idx
after := at + len(needle)
if after < len(s) {
r := lower[after]
if r != '>' && r != ' ' && r != '\t' && r != '\n' && r != '\r' && r != '/' {
pos = after
continue // 形如 </scriptx> 的伪闭合,继续找
}
}
end := strings.IndexByte(s[after:], '>')
if end < 0 {
return -1
}
return after + end + 1
}
}
// parseAllowedAnchor 解析 <a ...> 标签:仅当 href 合法时返回 (href, true)。
func parseAllowedAnchor(tag string) (string, bool) {
inner := tag[1 : len(tag)-1]
// 标签名
nameEnd := 0
for nameEnd < len(inner) {
r := inner[nameEnd]
if r == ' ' || r == '\t' || r == '\n' || r == '\r' {
break
}
nameEnd++
}
href, found := scanAttr(inner[nameEnd:], "href")
if !found {
return "", false // 无 href 的 <a> 不放行(避免依赖默认行为)
}
href = strings.TrimSpace(href)
lower := strings.ToLower(href)
if !(strings.HasPrefix(lower, "http://") || strings.HasPrefix(lower, "https://") ||
strings.HasPrefix(href, "/") || strings.HasPrefix(href, "#")) {
return "", false // javascript:/data: 等一律拒绝
}
return href, true
}
// scanAttr 扫描属性串中的目标属性(支持双引号/单引号/无引号值)。
func scanAttr(s, name string) (string, bool) {
lower := strings.ToLower(s)
want := strings.ToLower(name)
for i := 0; i < len(lower); {
// 跳过空白
for i < len(lower) && (lower[i] == ' ' || lower[i] == '\t' || lower[i] == '\n' || lower[i] == '\r') {
i++
}
if i >= len(lower) {
break
}
// 属性名
start := i
for i < len(lower) && lower[i] != '=' && lower[i] != ' ' && lower[i] != '\t' && lower[i] != '\n' && lower[i] != '\r' {
i++
}
attrName := lower[start:i]
// 跳过空白
for i < len(lower) && (lower[i] == ' ' || lower[i] == '\t' || lower[i] == '\n' || lower[i] == '\r') {
i++
}
if i < len(lower) && lower[i] == '=' {
i++ // 跳过 '='
for i < len(lower) && (lower[i] == ' ' || lower[i] == '\t' || lower[i] == '\n' || lower[i] == '\r') {
i++
}
var val string
if i < len(lower) && (lower[i] == '"' || lower[i] == '\'') {
q := lower[i]
i++
vs := i
for i < len(lower) && lower[i] != q {
i++
}
val = s[vs:i]
if i < len(lower) {
i++ // 跳过闭合引号
}
} else {
vs := i
for i < len(lower) && lower[i] != ' ' && lower[i] != '\t' && lower[i] != '\n' && lower[i] != '\r' {
i++
}
val = s[vs:i]
}
if attrName == want {
return val, true
}
} else if attrName == want {
return "", true // 布尔属性:存在即命中(值空,调用方按非法处理)
}
}
return "", false
}
// escapeAttr HTML 属性转义。
func escapeAttr(s string) string {
r := strings.NewReplacer("&", "&amp;", "<", "&lt;", ">", "&gt;", `"`, "&#34;", "'", "&#39;")
return r.Replace(s)
}
+44
View File
@@ -0,0 +1,44 @@
// sanitize_test.go — SanitizeInlineHTML 单测(安全审计 L7)。
package settings
import "testing"
func TestSanitizeInlineHTML(t *testing.T) {
cases := []struct {
name string
in string
want string
}{
{"空串", "", ""},
{"纯文本保留", "欢迎使用文件快传", "欢迎使用文件快传"},
{"合法链接保留", `<a href="https://example.com">官网</a>`, `<a href="https://example.com">官网</a>`},
{"相对路径链接", `<a href="/docs">文档</a>`, `<a href="/docs">文档</a>`},
{"锚点链接", `<a href="#top">顶部</a>`, `<a href="#top">顶部</a>`},
{"script 整体丢弃", `hello<script>alert(1)</script>world`, "helloworld"},
{"img 丢弃保留文本", `a<img src=x onerror=alert(1)>b`, "ab"},
{"javascript href 拒绝", `<a href="javascript:alert(1)">x</a>`, "x"},
{"data href 拒绝", `<a href="data:text/html,<script>">x</a>`, "x"},
{"事件属性不透传", `<a href="/x" onclick="evil()">y</a>`, `<a href="/x">y</a>`},
{"注释丢弃", `a<!-- secret -->b`, "ab"},
{"未闭合标签丢弃剩余", `ok<script>alert(1)`, "ok"},
{"iframe 丢弃", `<iframe src="//evil"></iframe>text`, "text"},
{"样式标签丢弃", `<style>*{}</style>plain`, "plain"},
{"嵌套危险标签", `<div onclick=e><b>bold</b></div>`, "bold"},
{"大小写标签", `<A HREF="https://e.com">L</A>`, `<a href="https://e.com">L</a>`},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
if got := SanitizeInlineHTML(tc.in); got != tc.want {
t.Fatalf("SanitizeInlineHTML(%q) = %q, want %q", tc.in, got, tc.want)
}
})
}
}
func TestSanitizeInlineHTMLNoScriptContent(t *testing.T) {
// script 内部文本也必须丢弃(不做 HTML 转义输出,避免 alert 字样进入页面 DOM)
got := SanitizeInlineHTML(`<script>var x = "</b>"; alert(1)</script>fine`)
if got != "fine" {
t.Fatalf("script 内容应整体丢弃, got %q", got)
}
}
+87
View File
@@ -0,0 +1,87 @@
// Package settings — schema.gov2 配置键 schema 常量与元数据表。
//
// 键名常量的单一事实来源在 internal/config/schema.godefaults() 需引用);
// 本文件 re-export 供 API/管理层使用,并提供「键名/类型/默认值」全量表,
// 供管理端设置页与文档生成(t4)对齐。新增键必须同步:
// 1. config/schema.go 键名与边界常量
// 2. config/config.go defaults() 默认值
// 3. 本文件 KVSchema() 元数据行
// 4. schema 同步测试(config schema_test / settings schema_test
package settings
import "filecodebox/internal/config"
// —— 键名 re-export(与 config 包保持同一字符串,避免魔法值散落)——
const (
// 需求 ① 背景图
KeyBackground = config.KeyBackground
KeyBackgroundURL = config.KeyBackgroundURL
// 需求 ② 页脚
KeyFooterText = config.KeyFooterText
KeyFooterBeian = config.KeyFooterBeian
// 需求 ③ 系统通知
KeyNotifyEnabled = config.KeyNotifyEnabled
KeyNotifyTitle = config.KeyNotifyTitle
KeyNotifyContent = config.KeyNotifyContent
// 需求 ④ 保存策略与上传频率限制
KeyMaxSaveSeconds = config.KeyMaxSaveSeconds
KeyMaxSaveCount = config.KeyMaxSaveCount
KeyExpireStyle = config.KeyExpireStyle
KeyUploadCount = config.KeyUploadCount
KeyUploadMinute = config.KeyUploadMinute
// 需求 ④⑩ 存储策略
KeyUploadSize = config.KeyUploadSize
KeyMaxFileSize = config.KeyMaxFileSize
KeyAllowedTypes = config.KeyAllowedTypes
KeyStorageLimit = config.KeyStorageLimit
KeyOpenUpload = config.KeyOpenUpload
// v3 存储引擎
KeyStorageEngine = config.KeyStorageEngine
)
// —— 取值边界 re-export ——
const (
MaxSaveSecondsMax = config.MaxSaveSecondsMax // 最长保存秒数上限(365 天)
MaxSaveCountMax = config.MaxSaveCountMax // 保存次数上限
MaxFileSizeMax = config.MaxFileSizeMax // 单文件大小上限(10 GiB
BackgroundURLMaxLen = config.BackgroundURLMaxLen // 背景图 URL 长度上限
FooterTextMaxLen = config.FooterTextMaxLen // 页脚内容长度上限
FooterBeianMaxLen = config.FooterBeianMaxLen // 备案号长度上限
NotifyTitleMaxLen = config.NotifyTitleMaxLen // 通知标题长度上限
NotifyContentMaxLen = config.NotifyContentMaxLen // 通知内容长度上限
)
// 敏感键:不允许出现在管理端 config get 下发/前端可见集合中(双模式下一致生效)。
// v3:引擎凭据(webdav_password/s3_secret_access_key/aws_session_token)加入敏感集——
// 管理端 get 返回掩码占位,update 时空串/掩码=不修改;公开 config 永不下发。
var SensitiveKeys = []string{
"admin_token", "jwt_secret",
"webdav_password", "s3_secret_access_key", "aws_session_token",
}
// SensitiveMaskValue 敏感键掩码占位(管理端 get 展示用)。
const SensitiveMaskValue = "******"
// IsSensitiveKey 判断键是否为敏感键(config get 必须屏蔽)。
func IsSensitiveKey(key string) bool {
for _, k := range SensitiveKeys {
if k == key {
return true
}
}
return false
}
// KVSchema 返回 v2 全量配置键元数据(键名/类型/默认值/边界/说明)。
// 默认值必须与 config defaults() 一致(schema 同步测试保证)。
func KVSchema() []config.KVSchemaEntry { return config.KVSchema() }
// KVSchemaByKey 以键名为索引查看 schema;未知键返回 nil。
func KVSchemaByKey(key string) *config.KVSchemaEntry {
for i := range config.KVSchema() {
if config.KVSchema()[i].Key == key {
return &config.KVSchema()[i]
}
}
return nil
}
+172
View File
@@ -0,0 +1,172 @@
// Package settings 提供数据库 settings KV 的运行时读写:
// envFCB_*)提供基线,DB KV 覆盖可变项;管理端修改后立即生效。
package settings
import (
"context"
"encoding/json"
"errors"
"log"
"sync"
"time"
"gorm.io/gorm"
"filecodebox/internal/config"
"filecodebox/internal/model"
)
// settingsKey 数据库中的配置键(对齐参考实现)。
const settingsKey = "settings"
// Manager 设置管理器:线程安全,缓存 KV 覆盖到内存。
type Manager struct {
db *gorm.DB
cfg *config.Config
mu sync.RWMutex
secret string // jwt_secret(频繁使用,单独缓存)
initPwd string // admin_token 哈希(频繁使用,单独缓存)
}
// NewManager 构造设置管理器并加载 DB KV。
// ensure 默认配置行(首次启动时写入 settings 键)。
func NewManager(ctx context.Context, db *gorm.DB, cfg *config.Config) (*Manager, error) {
m := &Manager{db: db, cfg: cfg}
if err := m.ensureSettingsRow(ctx); err != nil {
return nil, err
}
if err := m.Reload(ctx); err != nil {
return nil, err
}
return m, nil
}
// ensureSettingsRow 首次启动时把默认安全配置写入 KV(对齐 ensure_settings_row)。
func (m *Manager) ensureSettingsRow(ctx context.Context) error {
var row model.KeyValue
err := m.db.WithContext(ctx).Where(model.KeyValue{Key: settingsKey}).First(&row).Error
if err == nil {
return nil
}
// 双方言:必须用 errors.Is 判定(sqlite 驱动错误链与字符串消息与 postgres 不同)
if !errors.Is(err, gorm.ErrRecordNotFound) {
return err
}
// 不存在:写入初始配置(不含 admin_token/jwt_secret,保持未初始化状态)
initial := map[string]any{}
raw, _ := json.Marshal(initial)
row = model.KeyValue{Key: settingsKey, Value: strPtr(string(raw))}
if err := m.db.WithContext(ctx).Create(&row).Error; err != nil {
return err
}
log.Println("[settings] 系统尚未初始化,请在浏览器中打开站点并完成管理员密码设置")
return nil
}
// Reload 从数据库加载 settings KV 并覆盖到运行时配置。
func (m *Manager) Reload(ctx context.Context) error {
var row model.KeyValue
err := m.db.WithContext(ctx).Where(model.KeyValue{Key: settingsKey}).First(&row).Error
if err != nil {
// 行不存在时保持现有覆盖(双方言:errors.Is 判定)
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil
}
return err
}
kv := map[string]any{}
if row.Value != nil && *row.Value != "" {
if err := json.Unmarshal([]byte(*row.Value), &kv); err != nil {
log.Printf("[settings] settings KV 解析失败: %v", err)
}
}
// 内部键不允许通过 KV 覆盖(_ 开头)
safe := map[string]any{}
for k, v := range kv {
if len(k) > 0 && k[0] == '_' {
continue
}
safe[k] = v
}
m.mu.Lock()
m.cfg.ApplyKV(safe)
m.secret, _ = safe["jwt_secret"].(string)
m.initPwd, _ = safe["admin_token"].(string)
m.mu.Unlock()
return nil
}
// Get 返回当前配置(只读使用;不要修改返回值)。
func (m *Manager) Get() *config.Config { return m.cfg }
// SecretProvider 返回 jwt_secret 读取函数(JWT 中间件用)。
func (m *Manager) SecretProvider() func() string {
return func() string {
m.mu.RLock()
defer m.mu.RUnlock()
return m.secret
}
}
// IsInitialized 系统是否已完成初始化(管理员密码已设置且非默认密码)。
func (m *Manager) IsInitialized() bool {
m.mu.RLock()
defer m.mu.RUnlock()
if m.initPwd == "" {
return false
}
// 旧版默认密码视为未初始化(对齐 LEGACY_DEFAULT_ADMIN_TOKEN 检查)
return !verifyLegacyDefault(m.initPwd)
}
// legacyDefaultToken 参考实现的旧默认管理员密码。
const legacyDefaultToken = "FileCodeBox2023"
// verifyLegacyDefault 检查哈希是否对应旧默认密码。
func verifyLegacyDefault(hashed string) bool {
if hashed == "" {
return false
}
return VerifyPassword(legacyDefaultToken, hashed)
}
// UpdateKV 合并更新 settings KV(管理端保存配置)。
func (m *Manager) UpdateKV(ctx context.Context, patch map[string]any) error {
m.mu.Lock()
defer m.mu.Unlock()
// 读现有值
var row model.KeyValue
err := m.db.WithContext(ctx).Where(model.KeyValue{Key: settingsKey}).First(&row).Error
kv := map[string]any{}
if err == nil && row.Value != nil {
_ = json.Unmarshal([]byte(*row.Value), &kv)
}
for k, v := range patch {
if len(k) > 0 && k[0] == '_' {
continue
}
kv[k] = v
}
raw, err := json.Marshal(kv)
if err != nil {
return err
}
if err == nil && row.ID > 0 {
row.Value = strPtr(string(raw))
return m.db.WithContext(ctx).Model(&row).Update("value", row.Value).Error
}
row = model.KeyValue{Key: settingsKey, Value: strPtr(string(raw))}
return m.db.WithContext(ctx).Create(&row).Error
}
// SystemStart 记录系统启动时间(对齐 sys_start 键)。
func (m *Manager) SystemStart(ctx context.Context) {
now := time.Now().UnixMilli()
raw, _ := json.Marshal(now)
_ = m.db.WithContext(ctx).Where(model.KeyValue{Key: "sys_start"}).
Assign(model.KeyValue{Value: strPtr(string(raw))}).
FirstOrCreate(&model.KeyValue{Key: "sys_start"}).Error
}
func strPtr(s string) *string { return &s }
+10
View File
@@ -0,0 +1,10 @@
package storage
import "errors"
// ErrRangeNotSatisfiable 请求的字节范围超出文件大小(对应 HTTP 416)。
// 属于增量错误定义,不改动 interface.go 的既有签名。
var ErrRangeNotSatisfiable = errors.New("storage: 请求范围超出文件大小")
// ErrHashMismatch 分片哈希校验失败(合并分片时检测到数据损坏)。
var ErrHashMismatch = errors.New("storage: 分片哈希校验失败")
+26
View File
@@ -0,0 +1,26 @@
package storage
import (
"context"
"fmt"
)
// Factory 按配置构造存储引擎。由 go-storage 提供 New* 实现后接入。
// 这里提供注册表模式:各引擎实现注册自己的构造函数,main.go 按名称选择。
type Factory func(ctx context.Context) (Storage, error)
var registry = map[string]Factory{}
// RegisterEngine 注册引擎构造函数(init 时调用,名称:local|s3|webdav)。
func RegisterEngine(name string, f Factory) {
registry[name] = f
}
// NewEngine 按名称构造引擎。
func NewEngine(ctx context.Context, name string) (Storage, error) {
f, ok := registry[name]
if !ok {
return nil, fmt.Errorf("storage: 未知存储引擎 %q(仅支持 local|s3|webdav", name)
}
return f(ctx)
}
+111
View File
@@ -0,0 +1,111 @@
// Package storage 定义存储引擎统一契约。
//
// 本文件是 go-storage 并行开发的接口契约:签名一经定义不再改动。
// 三种引擎(local/s3/webdav)都要实现该接口;工厂按 FCB_STORAGE_ENGINE 选择。
package storage
import (
"context"
"errors"
"io"
)
// 错误定义:实现方应返回这些哨兵错误(可用 %w 包装),便于 API 层映射 HTTP 状态码。
var (
// ErrNotFound 文件不存在(HTTP 404)。
ErrNotFound = errors.New("storage: 文件不存在")
// ErrInvalidPath 非法路径(路径穿越等,HTTP 400)。
ErrInvalidPath = errors.New("storage: 非法文件路径")
// ErrUnavailable 存储服务不可用(连接失败等,HTTP 503)。
ErrUnavailable = errors.New("storage: 存储服务不可用")
)
// FileMeta 文件元信息(大小等)。
type FileMeta struct {
Size int64 // 字节数
ContentType string // MIME 类型,可为空
AcceptRanges bool // 是否支持 Range 请求
}
// Download 流式下载句柄。调用方负责 Close。
type Download struct {
// ReadCloser 文件内容流(已按 Range 重定位)。
io.ReadCloser
// Meta 文件元信息。
Meta FileMeta
// Start 当前流的起始字节偏移(Range 请求时为 rangeStart)。
Start int64
// End 流的结束字节偏移(含);未知为 -1。
End int64
// Total 文件总大小(字节);未知为 -1。
Total int64
}
// Range 字节范围(对齐 HTTP Range 语义)。
// nil 指针表示完整文件。
type Range struct {
Start int64 // 起始字节(含)
End int64 // 结束字节(含);-1 表示到文件末尾
}
// Storage 存储引擎统一接口。
//
// 约定:
// - savePath 为存储侧相对路径(引擎内部负责安全解析,拒绝 .. 穿越);
// - 所有方法必须是并发安全的;
// - 实现方遇到不可恢复错误时返回本包哨兵错误(或用 %w 包装)。
type Storage interface {
// SaveFile 流式保存文件:r 读取到 EOF 即完成,返回实际写入字节数。
// 引擎必须按 256KB 级别分块读取,不得将整个文件读入内存。
SaveFile(ctx context.Context, r io.Reader, savePath string) (int64, error)
// DeleteFile 删除文件;文件不存在时返回 ErrNotFound 或 nil 均可接受。
DeleteFile(ctx context.Context, savePath string) error
// Open 以下载模式打开文件,支持 HTTP Range 请求语义:
// - rng 为 nil:返回完整文件流(Start=0End=Total-1);
// - rng 非 nil:返回 [Start, End] 区间流。
// 引擎应尽量透传 RangeWebDAV/S3)或按块 seeklocal)。
Open(ctx context.Context, savePath string, rng *Range) (*Download, error)
// Stat 获取文件元信息;不存在返回 ErrNotFound。
Stat(ctx context.Context, savePath string) (*FileMeta, error)
// SaveChunk 保存一个分片到临时区(upload_id 隔离),返回分片字节数。
SaveChunk(ctx context.Context, uploadID string, chunkIndex int, r io.Reader, savePath string) (int64, error)
// MergeChunks 按索引 0..total-1 有序合并分片并落为正式文件。
// verifyHash 为nil 时不校验;否则为分片 SHA256 校验函数(输入索引,输出期望哈希,空串表示跳过)。
// 返回 (最终文件大小, 整个文件 SHA256)。
MergeChunks(ctx context.Context, uploadID string, total int, verifyHash func(index int) (string, error), savePath string) (int64, string, error)
// CleanChunks 清理分片临时区;不存在时静默成功。
CleanChunks(ctx context.Context, uploadID string, savePath string) error
// FileExists 检查文件是否存在。
FileExists(ctx context.Context, savePath string) (bool, error)
// HeadMeta 读取对象元信息与头部字节(可选能力,供直传 confirm 校验实际
// 大小与内容;不支持时返回 ErrNotSupported)。
// meta 允许为 nil(仅取头部);head 为对象前 headBytes 字节(不足时取实际长度)。
HeadMeta(ctx context.Context, savePath string, headBytes int64) (*FileMeta, []byte, error)
// PresignGetURL 生成限时直链(下载);不支持直链的引擎返回 ErrNotSupported。
PresignGetURL(ctx context.Context, savePath string, expires int64) (string, error)
// PresignPutURL 生成限时直传(上传)URL;不支持直传的引擎返回 ErrNotSupported。
PresignPutURL(ctx context.Context, savePath string, expires int64) (string, error)
// HealthCheck 引擎健康检查(启动时与 /health 使用)。
HealthCheck(ctx context.Context) error
}
// ErrNotSupported 当前引擎不支持该能力(如本地引擎不支持预签名)。
var ErrNotSupported = errors.New("storage: 当前引擎不支持该操作")
// ChunkPath 返回分片临时路径(约定统一为 <dir>/chunks/<upload_id>/<index>.part)。
// 引擎可使用 ChunkDir 拼接自身路径。
type PathBuilder interface {
// ChunkDir 分片临时目录(相对 savePath 所在目录)。
ChunkDir(savePath, uploadID string) string
}
+449
View File
@@ -0,0 +1,449 @@
package storage
import (
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"fmt"
"io"
"mime"
"os"
"path/filepath"
"strings"
"sync/atomic"
)
// 每次读写使用的缓冲大小:256KB,对齐参考实现 SystemFileStorage.chunk_size。
const localChunkSize = 256 * 1024
// LocalStorage 本地文件系统引擎。
//
// 相比参考实现(SystemFileStorage)的改进:
// - 双重路径防护:清洗相对路径 + 根目录前缀校验 + 符号链接逃逸校验;
// - 全部落盘走「临时文件 + fsync + 原子重命名」,断电/中断不产生半截文件;
// - 下载使用 io.NewSectionReader 支持任意 Range,无需整文件读入内存。
type LocalStorage struct {
// root 存储根目录(绝对路径)。
root string
// rootReal 经符号链接解析后的真实根目录,用于逃逸校验。
rootReal string
}
// NewLocalStorage 构造本地引擎。root 为空时使用系统临时目录下的 filecodebox_storage。
func NewLocalStorage(root string) (*LocalStorage, error) {
if strings.TrimSpace(root) == "" {
root = filepath.Join(os.TempDir(), "filecodebox_storage")
}
abs, err := filepath.Abs(root)
if err != nil {
return nil, fmt.Errorf("storage/local: 解析根目录失败: %w", err)
}
if err := os.MkdirAll(abs, 0o755); err != nil {
return nil, fmt.Errorf("storage/local: 创建根目录失败: %w", err)
}
real := abs
if resolved, err := filepath.EvalSymlinks(abs); err == nil {
real = resolved
}
return &LocalStorage{root: abs, rootReal: real}, nil
}
func init() {
RegisterEngine("local", func(ctx context.Context) (Storage, error) {
return NewLocalStorage(engineOptions.Local.Root)
})
}
// withinRoot 判断路径 p 是否位于 root 内(含 root 本身)。
func withinRoot(p, root string) bool {
p = filepath.Clean(p)
root = filepath.Clean(root)
if p == root {
return true
}
return strings.HasPrefix(p, root+string(os.PathSeparator))
}
// absPath 将存储侧相对路径解析为根目录内的绝对路径。
// 任何路径穿越或符号链接逃逸都会返回 ErrInvalidPath。
func (l *LocalStorage) absPath(savePath string) (string, error) {
cleaned, ok := SanitizePath(savePath)
if !ok {
return "", fmt.Errorf("%w: %q", ErrInvalidPath, savePath)
}
full := filepath.Join(l.root, filepath.FromSlash(cleaned))
if !withinRoot(full, l.root) {
return "", fmt.Errorf("%w: %q", ErrInvalidPath, savePath)
}
// 符号链接逃逸校验:文件已存在时解析真实路径;不存在时校验最深已存在的父目录。
if real, err := filepath.EvalSymlinks(full); err == nil {
if !withinRoot(real, l.rootReal) {
return "", fmt.Errorf("%w: 符号链接逃逸 %q", ErrInvalidPath, savePath)
}
} else {
dir := filepath.Dir(full)
if realDir, err := filepath.EvalSymlinks(dir); err == nil && !withinRoot(realDir, l.rootReal) {
return "", fmt.Errorf("%w: 符号链接逃逸 %q", ErrInvalidPath, savePath)
}
}
return full, nil
}
// SaveFile 流式保存:256KB 分块读取写入临时文件,fsync 后原子重命名。
func (l *LocalStorage) SaveFile(ctx context.Context, r io.Reader, savePath string) (int64, error) {
full, err := l.absPath(savePath)
if err != nil {
return 0, err
}
src := &countingReader{r: r}
if err := writeFileAtomic(full, src); err != nil {
return src.count(), err
}
return src.count(), nil
}
// DeleteFile 删除文件;文件不存在时静默成功(对齐契约)。
func (l *LocalStorage) DeleteFile(ctx context.Context, savePath string) error {
full, err := l.absPath(savePath)
if err != nil {
return err
}
if err := os.Remove(full); err != nil && !errors.Is(err, os.ErrNotExist) {
return fmt.Errorf("storage/local: 删除失败: %w", err)
}
return nil
}
// Open 打开文件下载流;rng 非 nil 时用 SectionReader 实现 Range 语义。
func (l *LocalStorage) Open(ctx context.Context, savePath string, rng *Range) (*Download, error) {
full, err := l.absPath(savePath)
if err != nil {
return nil, err
}
f, err := os.Open(full)
if err != nil {
if errors.Is(err, os.ErrNotExist) {
return nil, ErrNotFound
}
return nil, fmt.Errorf("storage/local: 打开文件失败: %w", err)
}
info, err := f.Stat()
if err != nil {
_ = f.Close()
return nil, fmt.Errorf("storage/local: 获取文件信息失败: %w", err)
}
if info.IsDir() {
_ = f.Close()
return nil, ErrNotFound
}
size := info.Size()
start, end := int64(0), size-1
if rng != nil {
if rng.Start < 0 || (rng.End != -1 && rng.End < rng.Start) {
_ = f.Close()
return nil, fmt.Errorf("%w: 非法 Range", ErrRangeNotSatisfiable)
}
if rng.Start >= size {
_ = f.Close()
return nil, ErrRangeNotSatisfiable
}
start = rng.Start
end = size - 1
if rng.End != -1 && rng.End < end {
end = rng.End
}
}
section := io.NewSectionReader(f, start, end-start+1)
dl := &Download{
ReadCloser: &fileSection{Reader: section, closer: f},
Start: start,
End: end,
Total: size,
Meta: FileMeta{
Size: size,
ContentType: mime.TypeByExtension(strings.ToLower(filepath.Ext(full))),
AcceptRanges: true,
},
}
if end < 0 { // 空文件:End 语义上等于 -1(未知),Total=0 已表达大小
dl.End = -1
}
return dl, nil
}
// fileSection 组合 SectionReader 与文件关闭器。
type fileSection struct {
io.Reader
closer io.Closer
}
func (f *fileSection) Close() error { return f.closer.Close() }
// Stat 获取文件元信息。
func (l *LocalStorage) Stat(ctx context.Context, savePath string) (*FileMeta, error) {
full, err := l.absPath(savePath)
if err != nil {
return nil, err
}
info, err := os.Stat(full)
if err != nil {
if errors.Is(err, os.ErrNotExist) {
return nil, ErrNotFound
}
return nil, fmt.Errorf("storage/local: Stat 失败: %w", err)
}
if info.IsDir() {
return nil, ErrNotFound
}
return &FileMeta{
Size: info.Size(),
ContentType: mime.TypeByExtension(strings.ToLower(filepath.Ext(full))),
AcceptRanges: true,
}, nil
}
// SaveChunk 保存分片到 <父目录>/chunks/<uploadID>/<index>.part,原子写入。
func (l *LocalStorage) SaveChunk(ctx context.Context, uploadID string, chunkIndex int, r io.Reader, savePath string) (int64, error) {
// 先校验目标路径合法性,分片目录随合法路径派生。
if _, err := l.absPath(savePath); err != nil {
return 0, err
}
chunkRel, ok := SanitizePath(ChunkPartPath(savePath, uploadID, chunkIndex))
if !ok {
return 0, fmt.Errorf("%w: 分片路径非法 %q", ErrInvalidPath, savePath)
}
full := filepath.Join(l.root, filepath.FromSlash(chunkRel))
if !withinRoot(full, l.root) {
return 0, fmt.Errorf("%w: 分片路径非法 %q", ErrInvalidPath, chunkRel)
}
src := &countingReader{r: r}
if err := writeFileAtomic(full, src); err != nil {
return src.count(), err
}
return src.count(), nil
}
// MergeChunks 按索引 0..total-1 有序合并分片:
// - 逐分片流式拷贝到临时输出(边拷贝边计算整文件与分片 SHA256);
// - verifyHash 非 nil 时校验分片哈希(空串跳过);
// - 全部通过后 fsync + 原子重命名,并清理分片临时目录。
func (l *LocalStorage) MergeChunks(ctx context.Context, uploadID string, total int, verifyHash func(index int) (string, error), savePath string) (int64, string, error) {
if total <= 0 {
return 0, "", fmt.Errorf("storage/local: 非法分片总数 %d", total)
}
full, err := l.absPath(savePath)
if err != nil {
return 0, "", err
}
if err := os.MkdirAll(filepath.Dir(full), 0o755); err != nil {
return 0, "", fmt.Errorf("storage/local: 创建目标目录失败: %w", err)
}
tmp, err := os.CreateTemp(filepath.Dir(full), "."+filepath.Base(full)+".merging-*")
if err != nil {
return 0, "", fmt.Errorf("storage/local: 创建临时文件失败: %w", err)
}
tmpName := tmp.Name()
defer func() {
_ = tmp.Close()
_ = os.Remove(tmpName) // 成功时已被重命名,删除静默失败
}()
totalHash := sha256.New()
buf := make([]byte, localChunkSize)
var size int64
for i := 0; i < total; i++ {
if err := ctx.Err(); err != nil {
return 0, "", err
}
partRel, ok := SanitizePath(ChunkPartPath(savePath, uploadID, i))
if !ok {
return 0, "", fmt.Errorf("%w: 分片 %d 路径非法", ErrInvalidPath, i)
}
partPath := filepath.Join(l.root, filepath.FromSlash(partRel))
in, err := os.Open(partPath)
if err != nil {
return 0, "", fmt.Errorf("storage/local: 分片 %d 不存在: %w", i, err)
}
chunkHash := sha256.New()
n, err := io.CopyBuffer(io.MultiWriter(tmp, totalHash, chunkHash), in, buf)
_ = in.Close()
if err != nil {
return 0, "", fmt.Errorf("storage/local: 读取分片 %d 失败: %w", i, err)
}
if verifyHash != nil {
expected, err := verifyHash(i)
if err != nil {
return 0, "", err
}
if expected != "" && expected != hex.EncodeToString(chunkHash.Sum(nil)) {
return 0, "", fmt.Errorf("%w: 分片 %d 期望 %s", ErrHashMismatch, i, expected)
}
}
size += n
}
if err := tmp.Sync(); err != nil {
return 0, "", fmt.Errorf("storage/local: 落盘失败: %w", err)
}
if err := tmp.Close(); err != nil {
return 0, "", fmt.Errorf("storage/local: 关闭临时文件失败: %w", err)
}
if err := os.Rename(tmpName, full); err != nil {
return 0, "", fmt.Errorf("storage/local: 原子重命名失败: %w", err)
}
// 合并成功后清理分片临时目录(静默容错,不掩盖成功结果)。
_ = l.CleanChunks(ctx, uploadID, savePath)
return size, hex.EncodeToString(totalHash.Sum(nil)), nil
}
// CleanChunks 清理分片临时目录;不存在时静默成功,并尝试移除空 chunks 父目录。
func (l *LocalStorage) CleanChunks(ctx context.Context, uploadID string, savePath string) error {
dirRel, ok := SanitizePath(chunkDirOf(savePath, uploadID))
if !ok {
return fmt.Errorf("%w: 分片目录非法 %q", ErrInvalidPath, savePath)
}
dir := filepath.Join(l.root, filepath.FromSlash(dirRel))
if !withinRoot(dir, l.root) {
return fmt.Errorf("%w: 分片目录非法 %q", ErrInvalidPath, dirRel)
}
if err := os.RemoveAll(dir); err != nil {
return fmt.Errorf("storage/local: 清理分片目录失败: %w", err)
}
// 父级 chunks 目录为空则一并清理(对齐参考实现)。
chunksParent := filepath.Dir(dir)
if entries, err := os.ReadDir(chunksParent); err == nil && len(entries) == 0 {
_ = os.Remove(chunksParent)
}
return nil
}
// chunkDirOf 返回分片目录(去掉文件名部分):<父目录>/chunks/<uploadID>。
func chunkDirOf(savePath, uploadID string) string {
cd := ChunkDir(savePath, uploadID)
// ChunkDir 返回 "<dir>/chunks/<uploadID>/<name>",去掉末段文件名即目录。
if idx := strings.LastIndex(cd, "/"); idx > 0 {
return cd[:idx]
}
return cd
}
// FileExists 检查文件是否存在;非法路径按不存在处理(对齐参考实现)。
func (l *LocalStorage) FileExists(ctx context.Context, savePath string) (bool, error) {
full, err := l.absPath(savePath)
if err != nil {
return false, nil
}
info, err := os.Stat(full)
if err != nil {
if errors.Is(err, os.ErrNotExist) {
return false, nil
}
return false, fmt.Errorf("storage/local: Stat 失败: %w", err)
}
return !info.IsDir(), nil
}
// HeadMeta 读取文件元信息与前 n 字节(本地引擎实现)。
func (l *LocalStorage) HeadMeta(ctx context.Context, savePath string, headBytes int64) (*FileMeta, []byte, error) {
full, err := l.absPath(savePath)
if err != nil {
return nil, nil, err
}
f, err := os.Open(full)
if err != nil {
if errors.Is(err, os.ErrNotExist) {
return nil, nil, ErrNotFound
}
return nil, nil, fmt.Errorf("storage/local: 打开文件失败: %w", err)
}
defer func() { _ = f.Close() }()
info, err := f.Stat()
if err != nil {
return nil, nil, fmt.Errorf("storage/local: Stat 失败: %w", err)
}
if info.IsDir() {
return nil, nil, ErrNotFound
}
head := make([]byte, headBytes)
n, _ := io.ReadFull(f, head)
return &FileMeta{
Size: info.Size(),
ContentType: mime.TypeByExtension(strings.ToLower(filepath.Ext(full))),
AcceptRanges: true,
}, head[:n], nil
}
// PresignGetURL 本地引擎不支持直链。
func (l *LocalStorage) PresignGetURL(ctx context.Context, savePath string, expires int64) (string, error) {
return "", ErrNotSupported
}
// PresignPutURL 本地引擎不支持直传。
func (l *LocalStorage) PresignPutURL(ctx context.Context, savePath string, expires int64) (string, error) {
return "", ErrNotSupported
}
// HealthCheck 健康检查:根目录可写(写入并删除探针文件)。
func (l *LocalStorage) HealthCheck(ctx context.Context) error {
if err := os.MkdirAll(l.root, 0o755); err != nil {
return fmt.Errorf("%w: 本地存储根目录不可创建: %v", ErrUnavailable, err)
}
probe := filepath.Join(l.root, ".health-probe")
if err := os.WriteFile(probe, []byte("ok"), 0o644); err != nil {
return fmt.Errorf("%w: 本地存储不可写: %v", ErrUnavailable, err)
}
_ = os.Remove(probe)
return nil
}
// writeFileAtomic 临时文件 + fsync + rename 的原子落盘。
func writeFileAtomic(dst string, src io.Reader) error {
dir := filepath.Dir(dst)
if err := os.MkdirAll(dir, 0o755); err != nil {
return fmt.Errorf("storage/local: 创建目录失败: %w", err)
}
tmp, err := os.CreateTemp(dir, "."+filepath.Base(dst)+".tmp-*")
if err != nil {
return fmt.Errorf("storage/local: 创建临时文件失败: %w", err)
}
tmpName := tmp.Name()
cleanup := func() { _ = tmp.Close(); _ = os.Remove(tmpName) }
if _, err := io.CopyBuffer(tmp, src, make([]byte, localChunkSize)); err != nil {
cleanup()
return fmt.Errorf("storage/local: 写入失败: %w", err)
}
if err := tmp.Sync(); err != nil {
cleanup()
return fmt.Errorf("storage/local: fsync 失败: %w", err)
}
if err := tmp.Close(); err != nil {
_ = os.Remove(tmpName)
return fmt.Errorf("storage/local: 关闭临时文件失败: %w", err)
}
if err := os.Rename(tmpName, dst); err != nil {
_ = os.Remove(tmpName)
return fmt.Errorf("storage/local: 原子重命名失败: %w", err)
}
return nil
}
// countingReader 统计累计读取字节数(并发安全)。
type countingReader struct {
r io.Reader
n atomic.Int64
}
func (c *countingReader) Read(p []byte) (int, error) {
n, err := c.r.Read(p)
c.n.Add(int64(n))
return n, err
}
// count 返回累计字节数。
func (c *countingReader) count() int64 { return c.n.Load() }
// reset 归零计数(请求体重放时使用)。
func (c *countingReader) reset() { c.n.Store(0) }
// 接口编译期断言。
var _ Storage = (*LocalStorage)(nil)
+346
View File
@@ -0,0 +1,346 @@
package storage
import (
"bytes"
"context"
"crypto/sha256"
"encoding/hex"
"io"
"os"
"path/filepath"
"strings"
"sync"
"testing"
)
// newTestLocal 构造以临时目录为根的本地引擎。
func newTestLocal(t *testing.T) *LocalStorage {
t.Helper()
st, err := NewLocalStorage(t.TempDir())
if err != nil {
t.Fatalf("NewLocalStorage: %v", err)
}
return st
}
func sha256Hex(b []byte) string {
sum := sha256.Sum256(b)
return hex.EncodeToString(sum[:])
}
// TestLocalSaveOpenRange 保存/Stat/完整与 Range 下载。
func TestLocalSaveOpenRange(t *testing.T) {
st := newTestLocal(t)
ctx := context.Background()
data := []byte("hello filecodebox 本地引擎 0123456789")
n, err := st.SaveFile(ctx, bytes.NewReader(data), "2025/08/测试文件.bin")
if err != nil {
t.Fatalf("SaveFile: %v", err)
}
if n != int64(len(data)) {
t.Fatalf("SaveFile n = %d, want %d", n, len(data))
}
// Stat
meta, err := st.Stat(ctx, "2025/08/测试文件.bin")
if err != nil {
t.Fatalf("Stat: %v", err)
}
if meta.Size != int64(len(data)) || !meta.AcceptRanges {
t.Fatalf("Stat = %+v", meta)
}
// 完整下载:对齐 go-api 约定 Start=0、End=Total-1
dl, err := st.Open(ctx, "2025/08/测试文件.bin", nil)
if err != nil {
t.Fatalf("Open: %v", err)
}
got, err := io.ReadAll(dl)
_ = dl.Close()
if err != nil {
t.Fatalf("read full: %v", err)
}
if !bytes.Equal(got, data) {
t.Fatalf("full content mismatch")
}
if dl.Start != 0 || dl.End != int64(len(data))-1 || dl.Total != int64(len(data)) {
t.Fatalf("full download offsets = %d,%d,%d", dl.Start, dl.End, dl.Total)
}
// Range 下载 [2, 7]
dl, err = st.Open(ctx, "2025/08/测试文件.bin", &Range{Start: 2, End: 7})
if err != nil {
t.Fatalf("Open range: %v", err)
}
got, err = io.ReadAll(dl)
_ = dl.Close()
if err != nil {
t.Fatalf("read range: %v", err)
}
if !bytes.Equal(got, data[2:8]) {
t.Fatalf("range content = %q, want %q", got, data[2:8])
}
// Range end 越界自动钳制
dl, err = st.Open(ctx, "2025/08/测试文件.bin", &Range{Start: 5, End: 99999})
if err != nil {
t.Fatalf("Open clamp range: %v", err)
}
got, _ = io.ReadAll(dl)
_ = dl.Close()
if !bytes.Equal(got, data[5:]) {
t.Fatalf("clamp range mismatch")
}
// 起点越界 → 416
if _, err := st.Open(ctx, "2025/08/测试文件.bin", &Range{Start: int64(len(data)) + 1, End: -1}); err == nil {
t.Fatalf("out-of-range start should fail")
} else if !strings.Contains(err.Error(), ErrRangeNotSatisfiable.Error()) {
t.Fatalf("want ErrRangeNotSatisfiable, got %v", err)
}
// 不存在 → 404
if _, err := st.Open(ctx, "no/such/file.bin", nil); err == nil || !strings.Contains(err.Error(), ErrNotFound.Error()) {
t.Fatalf("want ErrNotFound, got %v", err)
}
}
// TestLocalSaveFileAtomic 原子写:目录自动创建,无临时残留。
func TestLocalSaveFileAtomic(t *testing.T) {
ctx := context.Background()
dir := t.TempDir()
st2, _ := NewLocalStorage(dir)
big := bytes.Repeat([]byte("abc123"), 128*1024) // 768KB,跨多个 256KB 缓冲
if _, err := st2.SaveFile(ctx, bytes.NewReader(big), "a/b/c/big.bin"); err != nil {
t.Fatalf("SaveFile: %v", err)
}
// 无临时残留
entries, _ := os.ReadDir(filepath.Join(dir, "a", "b", "c"))
for _, e := range entries {
if strings.HasPrefix(e.Name(), ".") {
t.Fatalf("临时文件残留: %s", e.Name())
}
}
got, _ := os.ReadFile(filepath.Join(dir, "a", "b", "c", "big.bin"))
if !bytes.Equal(got, big) {
t.Fatalf("content mismatch")
}
}
// TestLocalTraversal 防路径穿越(含符号链接逃逸)。
func TestLocalTraversal(t *testing.T) {
st := newTestLocal(t)
ctx := context.Background()
for _, p := range []string{"../escape.txt", "a/../../escape", "..", "/../x"} {
if _, err := st.SaveFile(ctx, strings.NewReader("x"), p); err == nil ||
!strings.Contains(err.Error(), ErrInvalidPath.Error()) {
t.Fatalf("SaveFile(%q) 应拒绝: %v", p, err)
}
if _, err := st.Open(ctx, p, nil); err == nil || !strings.Contains(err.Error(), ErrInvalidPath.Error()) {
t.Fatalf("Open(%q) 应拒绝: %v", p, err)
}
}
// 符号链接逃逸
outside := filepath.Join(t.TempDir(), "outside.txt")
if err := os.WriteFile(outside, []byte("secret"), 0o644); err != nil {
t.Fatal(err)
}
link := filepath.Join(st.root, "link.txt")
if err := os.Symlink(outside, link); err != nil {
t.Skipf("symlink 不可用: %v", err)
}
if _, err := st.Open(ctx, "link.txt", nil); err == nil ||
!strings.Contains(err.Error(), ErrInvalidPath.Error()) {
t.Fatalf("symlink escape 应拒绝: %v", err)
}
}
// TestLocalChunkLifecycle 分片保存/合并(索引有序 + SHA256 校验 + 清理临时目录)。
func TestLocalChunkLifecycle(t *testing.T) {
st := newTestLocal(t)
ctx := context.Background()
savePath := "2025/09/chunked.bin"
uploadID := "upload-abc"
chunks := [][]byte{[]byte("AAAA"), []byte("BB"), []byte("CCCCCC")}
hashes := make([]string, len(chunks))
var total int64
for i, c := range chunks {
n, err := st.SaveChunk(ctx, uploadID, i, bytes.NewReader(c), savePath)
if err != nil {
t.Fatalf("SaveChunk %d: %v", i, err)
}
if n != int64(len(c)) {
t.Fatalf("SaveChunk %d n = %d", i, n)
}
hashes[i] = sha256Hex(c)
total += int64(len(c))
}
// 分片文件确实存在于临时目录
for i := range chunks {
exists, err := st.FileExists(ctx, ChunkPartPath(savePath, uploadID, i))
if err != nil || !exists {
t.Fatalf("分片 %d 应存在: %v %v", i, exists, err)
}
}
size, fileHash, err := st.MergeChunks(ctx, uploadID, len(chunks), func(i int) (string, error) {
return hashes[i], nil
}, savePath)
if err != nil {
t.Fatalf("MergeChunks: %v", err)
}
if size != total {
t.Fatalf("merged size = %d, want %d", size, total)
}
want := sha256Hex(bytes.Join(chunks, nil))
if fileHash != want {
t.Fatalf("file hash = %s, want %s", fileHash, want)
}
// 合并内容 = 按索引有序拼接
got, err := os.ReadFile(filepath.Join(st.root, filepath.FromSlash(savePath)))
if err != nil {
t.Fatal(err)
}
if !bytes.Equal(got, bytes.Join(chunks, nil)) {
t.Fatalf("merged content mismatch")
}
// 合并后分片目录已清理
exists, _ := st.FileExists(ctx, ChunkPartPath(savePath, uploadID, 0))
if exists {
t.Fatalf("合并后分片应已清理")
}
if _, err := os.Stat(filepath.Join(st.root, "2025/09/chunks")); !os.IsNotExist(err) {
t.Fatalf("chunks 父目录应已清理: %v", err)
}
}
// TestLocalChunkHashMismatch 哈希不匹配 → 合并失败且不留输出文件。
func TestLocalChunkHashMismatch(t *testing.T) {
st := newTestLocal(t)
ctx := context.Background()
savePath := "mismatch.bin"
if _, err := st.SaveChunk(ctx, "uid", 0, strings.NewReader("data"), savePath); err != nil {
t.Fatal(err)
}
_, _, err := st.MergeChunks(ctx, "uid", 1, func(i int) (string, error) {
return sha256Hex([]byte("WRONG")), nil
}, savePath)
if err == nil || !strings.Contains(err.Error(), ErrHashMismatch.Error()) {
t.Fatalf("want ErrHashMismatch, got %v", err)
}
exists, _ := st.FileExists(ctx, savePath)
if exists {
t.Fatalf("校验失败不应产出正式文件")
}
}
// TestLocalVerifyHashAbort verifyHash 返回 error → 合并中止。
func TestLocalVerifyHashAbort(t *testing.T) {
st := newTestLocal(t)
ctx := context.Background()
if _, err := st.SaveChunk(ctx, "uid2", 0, strings.NewReader("data"), "v.bin"); err != nil {
t.Fatal(err)
}
boom := context.Canceled
if _, _, err := st.MergeChunks(ctx, "uid2", 1, func(i int) (string, error) {
return "", boom
}, "v.bin"); err == nil {
t.Fatalf("verifyHash 错误应向上传播")
}
}
// TestLocalDeleteExistsClean 删除/存在性/清理。
func TestLocalDeleteExistsClean(t *testing.T) {
st := newTestLocal(t)
ctx := context.Background()
if _, err := st.SaveFile(ctx, strings.NewReader("x"), "d/f.bin"); err != nil {
t.Fatal(err)
}
if ok, _ := st.FileExists(ctx, "d/f.bin"); !ok {
t.Fatalf("文件应存在")
}
if err := st.DeleteFile(ctx, "d/f.bin"); err != nil {
t.Fatalf("DeleteFile: %v", err)
}
if ok, _ := st.FileExists(ctx, "d/f.bin"); ok {
t.Fatalf("文件应已删除")
}
// 删除不存在 → nil
if err := st.DeleteFile(ctx, "d/f.bin"); err != nil {
t.Fatalf("删除不存在应静默: %v", err)
}
// CleanChunks 幂等
if err := st.CleanChunks(ctx, "uid-x", "y.bin"); err != nil {
t.Fatalf("CleanChunks: %v", err)
}
}
// TestLocalPresignNotSupported 预签名 → ErrNotSupported。
func TestLocalPresignNotSupported(t *testing.T) {
st := newTestLocal(t)
ctx := context.Background()
if _, err := st.PresignGetURL(ctx, "a.bin", 60); err == nil || !strings.Contains(err.Error(), ErrNotSupported.Error()) {
t.Fatalf("PresignGetURL want ErrNotSupported, got %v", err)
}
if _, err := st.PresignPutURL(ctx, "a.bin", 60); err == nil || !strings.Contains(err.Error(), ErrNotSupported.Error()) {
t.Fatalf("PresignPutURL want ErrNotSupported, got %v", err)
}
}
// TestLocalHealthCheck 健康检查。
func TestLocalHealthCheck(t *testing.T) {
st := newTestLocal(t)
if err := st.HealthCheck(context.Background()); err != nil {
t.Fatalf("HealthCheck: %v", err)
}
}
// TestLocalConcurrent 并发安全冒烟。
func TestLocalConcurrent(t *testing.T) {
st := newTestLocal(t)
ctx := context.Background()
var wg sync.WaitGroup
for i := 0; i < 16; i++ {
wg.Add(1)
go func(i int) {
defer wg.Done()
p := "conc/" + itoa(i) + ".bin"
if _, err := st.SaveFile(ctx, strings.NewReader(strings.Repeat("x", i+1)), p); err != nil {
t.Errorf("save %d: %v", i, err)
return
}
dl, err := st.Open(ctx, p, nil)
if err != nil {
t.Errorf("open %d: %v", i, err)
return
}
_, _ = io.ReadAll(dl)
_ = dl.Close()
}(i)
}
wg.Wait()
}
// TestLocalFactoryRegistry 工厂注册与构造。
func TestLocalFactoryRegistry(t *testing.T) {
prev := engineOptions.Local.Root
engineOptions.Local.Root = t.TempDir()
defer func() { engineOptions.Local.Root = prev }()
st, err := NewEngine(context.Background(), "local")
if err != nil {
t.Fatalf("NewEngine(local): %v", err)
}
if err := st.HealthCheck(context.Background()); err != nil {
t.Fatalf("HealthCheck: %v", err)
}
if _, err := NewEngine(context.Background(), "unknown"); err == nil {
t.Fatalf("未知引擎应报错")
}
}
+187
View File
@@ -0,0 +1,187 @@
package storage
import (
"context"
"fmt"
"io"
"sync"
)
// Manager 存储引擎管理器:实现 Storage 全接口并支持运行时热切换。
//
// v3 需求:管理后台可设置存储类型(local|s3|webdav)与各引擎参数,
// 保存后无需重启即生效。设计要点:
// - 读写/保存类操作全部委托到"当前引擎"(原子指针,无锁热路径);
// - Switch 先构建并健康检查新引擎,成功才替换指针,失败保持原引擎;
// - EngineOf 按名字取引擎实例(带缓存),供"按文件归属引擎取回旧文件"使用;
// - 管理端修改引擎参数后调用 Invalidate 使对应实例缓存失效,下次构建生效。
type Manager struct {
// build 构建指定引擎实例(由装配方注入:内部刷新全局 EngineOptions 后走工厂)。
build func(name string) (Storage, error)
mu sync.RWMutex
current Storage
curName string
cache map[string]Storage
}
// validEngines 合法引擎名(与 FCB_STORAGE_ENGINE 枚举一致)。
var validEngines = map[string]bool{"local": true, "s3": true, "webdav": true}
// ValidEngine 校验引擎名是否合法。
func ValidEngine(name string) bool { return validEngines[name] }
// NewManager 创建管理器:current 为启动时已构建的引擎(主装配流已做过健康检查)。
// build 注入构建函数(管理端切换/参数变更时使用,内部须串行——Manager 已加锁)。
func NewManager(name string, current Storage, build func(name string) (Storage, error)) *Manager {
return &Manager{
build: build,
current: current,
curName: name,
cache: map[string]Storage{name: current},
}
}
// —— Storage 接口委托(全部走当前引擎)——
// SaveFile 流式保存文件(委托当前引擎)。
func (m *Manager) SaveFile(ctx context.Context, r io.Reader, savePath string) (int64, error) {
return m.current.SaveFile(ctx, r, savePath)
}
// DeleteFile 删除文件(委托当前引擎)。
func (m *Manager) DeleteFile(ctx context.Context, savePath string) error {
return m.current.DeleteFile(ctx, savePath)
}
// Open 打开文件流(委托当前引擎;旧文件由 API 层先经 EngineOf 按归属引擎取)。
func (m *Manager) Open(ctx context.Context, savePath string, rng *Range) (*Download, error) {
return m.current.Open(ctx, savePath, rng)
}
// Stat 文件元信息(委托当前引擎)。
func (m *Manager) Stat(ctx context.Context, savePath string) (*FileMeta, error) {
return m.current.Stat(ctx, savePath)
}
// SaveChunk 保存分片(委托当前引擎)。
func (m *Manager) SaveChunk(ctx context.Context, uploadID string, chunkIndex int, r io.Reader, savePath string) (int64, error) {
return m.current.SaveChunk(ctx, uploadID, chunkIndex, r, savePath)
}
// MergeChunks 合并分片(委托当前引擎)。
func (m *Manager) MergeChunks(ctx context.Context, uploadID string, total int, verifyHash func(index int) (string, error), savePath string) (int64, string, error) {
return m.current.MergeChunks(ctx, uploadID, total, verifyHash, savePath)
}
// CleanChunks 清理分片临时区(委托当前引擎)。
func (m *Manager) CleanChunks(ctx context.Context, uploadID string, savePath string) error {
return m.current.CleanChunks(ctx, uploadID, savePath)
}
// FileExists 文件存在性(委托当前引擎)。
func (m *Manager) FileExists(ctx context.Context, savePath string) (bool, error) {
return m.current.FileExists(ctx, savePath)
}
// HeadMeta 元信息与头部字节(委托当前引擎;供直传 confirm 校验)。
func (m *Manager) HeadMeta(ctx context.Context, savePath string, headBytes int64) (*FileMeta, []byte, error) {
return m.current.HeadMeta(ctx, savePath, headBytes)
}
// PresignGetURL 限时直链下载(委托当前引擎)。
func (m *Manager) PresignGetURL(ctx context.Context, savePath string, expires int64) (string, error) {
return m.current.PresignGetURL(ctx, savePath, expires)
}
// PresignPutURL 限时直传(委托当前引擎)。
func (m *Manager) PresignPutURL(ctx context.Context, savePath string, expires int64) (string, error) {
return m.current.PresignPutURL(ctx, savePath, expires)
}
// HealthCheck 健康检查(委托当前引擎)。
func (m *Manager) HealthCheck(ctx context.Context) error {
return m.current.HealthCheck(ctx)
}
// —— 管理面:当前引擎名 / 按名取实例 / 热切换 / 缓存失效 ——
// CurrentName 当前引擎名(管理端展示与文件归属戳用;并发安全)。
func (m *Manager) CurrentName() string {
m.mu.RLock()
defer m.mu.RUnlock()
return m.curName
}
// Current 当前引擎实例。
func (m *Manager) Current() Storage {
m.mu.RLock()
defer m.mu.RUnlock()
return m.current
}
// EngineOf 按名字取引擎实例(带缓存;用于按文件归属引擎取回旧文件)。
// 实例不存在时现场构建(不健康检查——读旧文件尽力而为,构建失败即报错)。
func (m *Manager) EngineOf(name string) (Storage, error) {
if !validEngines[name] {
return nil, fmt.Errorf("storage: 未知存储引擎 %q(仅支持 local|s3|webdav", name)
}
m.mu.RLock()
if s, ok := m.cache[name]; ok {
m.mu.RUnlock()
return s, nil
}
m.mu.RUnlock()
m.mu.Lock()
defer m.mu.Unlock()
// 双检:拿写锁期间可能已被并发构建
if s, ok := m.cache[name]; ok {
return s, nil
}
s, err := m.build(name)
if err != nil {
return nil, err
}
m.cache[name] = s
return s, nil
}
// Switch 热切换当前引擎:构建新实例 → 健康检查 → 成功才替换指针。
// 任一步失败返回错误且当前引擎保持不变(管理端 503 上报)。
func (m *Manager) Switch(name string) (Storage, error) {
if !validEngines[name] {
return nil, fmt.Errorf("storage: 未知存储引擎 %q(仅支持 local|s3|webdav", name)
}
m.mu.Lock()
defer m.mu.Unlock()
if m.curName == name {
return m.current, nil
}
s, ok := m.cache[name]
if !ok {
var err error
s, err = m.build(name)
if err != nil {
return nil, fmt.Errorf("storage: 构建 %s 引擎失败: %w", name, err)
}
}
if err := s.HealthCheck(context.Background()); err != nil {
return nil, fmt.Errorf("storage: %s 引擎健康检查未通过: %w", name, err)
}
m.current = s
m.curName = name
m.cache[name] = s
return s, nil
}
// Invalidate 引擎参数变更后使对应实例缓存失效(下次 EngineOf/Switch 重建生效)。
// 当前引擎不受影响(运行中实例继续服务,直到显式 Switch)。
func (m *Manager) Invalidate(name string) {
m.mu.Lock()
defer m.mu.Unlock()
if name == m.curName {
return // 当前引擎实例仍被热路径使用,不重建;参数生效由下一次 Switch 完成
}
delete(m.cache, name)
}
+197
View File
@@ -0,0 +1,197 @@
package storage
import (
"context"
"errors"
"io"
"strings"
"sync/atomic"
"testing"
)
// fakeEngine 可配置健康检查结果的桩引擎。
type fakeEngine struct{ failHealth bool }
func (f *fakeEngine) SaveFile(ctx context.Context, r io.Reader, savePath string) (int64, error) {
return 0, nil
}
func (f *fakeEngine) DeleteFile(ctx context.Context, savePath string) error { return nil }
func (f *fakeEngine) Open(ctx context.Context, savePath string, rng *Range) (*Download, error) {
return nil, ErrNotFound
}
func (f *fakeEngine) Stat(ctx context.Context, savePath string) (*FileMeta, error) {
return nil, ErrNotFound
}
func (f *fakeEngine) SaveChunk(ctx context.Context, uploadID string, chunkIndex int, r io.Reader, savePath string) (int64, error) {
return 0, nil
}
func (f *fakeEngine) MergeChunks(ctx context.Context, uploadID string, total int, verifyHash func(index int) (string, error), savePath string) (int64, string, error) {
return 0, "", nil
}
func (f *fakeEngine) CleanChunks(ctx context.Context, uploadID string, savePath string) error {
return nil
}
func (f *fakeEngine) FileExists(ctx context.Context, savePath string) (bool, error) {
return false, nil
}
func (f *fakeEngine) HeadMeta(ctx context.Context, savePath string, headBytes int64) (*FileMeta, []byte, error) {
return nil, nil, ErrNotSupported
}
func (f *fakeEngine) PresignGetURL(ctx context.Context, savePath string, expires int64) (string, error) {
return "", ErrNotSupported
}
func (f *fakeEngine) PresignPutURL(ctx context.Context, savePath string, expires int64) (string, error) {
return "", ErrNotSupported
}
func (f *fakeEngine) HealthCheck(ctx context.Context) error {
if f.failHealth {
return ErrUnavailable
}
return nil
}
// newTestManager 构造测试用 Managerlocal 健康引擎起步;s3/webdav 由计数器控制健康。
func newTestManager(s3Fail *atomic.Bool) *Manager {
build := func(name string) (Storage, error) {
switch name {
case "local":
return &fakeEngine{}, nil
case "s3":
return &fakeEngine{failHealth: s3Fail.Load()}, nil
case "webdav":
return &fakeEngine{}, nil
}
return nil, errors.New("unknown")
}
return NewManager("local", &fakeEngine{}, build)
}
// TestSwitchSuccessAndCurrentName 切换成功后当前引擎名与实例更新。
func TestSwitchSuccessAndCurrentName(t *testing.T) {
var s3Fail atomic.Bool
m := newTestManager(&s3Fail)
if m.CurrentName() != "local" {
t.Fatalf("初始引擎应为 local,得到 %s", m.CurrentName())
}
if _, err := m.Switch("s3"); err != nil {
t.Fatalf("Switch(s3) 失败: %v", err)
}
if m.CurrentName() != "s3" {
t.Fatalf("切换后引擎应为 s3,得到 %s", m.CurrentName())
}
if _, err := m.Switch("webdav"); err != nil {
t.Fatalf("Switch(webdav) 失败: %v", err)
}
if m.CurrentName() != "webdav" {
t.Fatalf("切换后引擎应为 webdav,得到 %s", m.CurrentName())
}
}
// TestSwitchFailureKeepsCurrent 健康检查失败时保持原引擎(v3 核心语义)。
func TestSwitchFailureKeepsCurrent(t *testing.T) {
var s3Fail atomic.Bool
s3Fail.Store(true) // s3 不健康
m := newTestManager(&s3Fail)
if _, err := m.Switch("s3"); err == nil {
t.Fatal("s3 不健康时 Switch 应失败")
}
if m.CurrentName() != "local" {
t.Fatalf("切换失败后应保持 local,得到 %s", m.CurrentName())
}
// 恢复健康后可切换成功
s3Fail.Store(false)
if _, err := m.Switch("s3"); err != nil {
t.Fatalf("恢复健康后 Switch(s3) 应成功: %v", err)
}
if m.CurrentName() != "s3" {
t.Fatalf("恢复后引擎应为 s3,得到 %s", m.CurrentName())
}
}
// TestSwitchInvalidName 非法引擎名拒绝。
func TestSwitchInvalidName(t *testing.T) {
m := newTestManager(&atomic.Bool{})
if _, err := m.Switch("ftp"); err == nil || !strings.Contains(err.Error(), "未知存储引擎") {
t.Fatalf("非法引擎名应报未知存储引擎,得到 %v", err)
}
if !ValidEngine("local") || ValidEngine("ftp") {
t.Fatal("ValidEngine 判定错误")
}
}
// TestEngineOfCacheAndInvalidate EngineOf 缓存命中 + Invalidate 后重建(参数生效路径)。
func TestEngineOfCacheAndInvalidate(t *testing.T) {
var builds atomic.Int64
build := func(name string) (Storage, error) {
builds.Add(1)
return &fakeEngine{failHealth: false}, nil
}
m := NewManager("local", &fakeEngine{}, build)
s1, err := m.EngineOf("s3")
if err != nil {
t.Fatalf("EngineOf(s3): %v", err)
}
s2, err := m.EngineOf("s3")
if err != nil {
t.Fatalf("EngineOf(s3) second: %v", err)
}
if s1 != s2 {
t.Fatal("EngineOf 应命中缓存返回同一实例")
}
if n := builds.Load(); n != 1 {
t.Fatalf("应只构建 1 次,实际 %d", n)
}
// Invalidate 后下次取重建新实例
m.Invalidate("s3")
s3, err := m.EngineOf("s3")
if err != nil {
t.Fatalf("EngineOf(s3) after invalidate: %v", err)
}
if s3 == s1 {
t.Fatal("Invalidate 后应返回重建的新实例")
}
if n := builds.Load(); n != 2 {
t.Fatalf("Invalidate 后应再构建 1 次,实际累计 %d", n)
}
}
// TestInvalidateCurrentNoop Invalidate 当前引擎不生效(热路径实例保持)。
func TestInvalidateCurrentNoop(t *testing.T) {
m := newTestManager(&atomic.Bool{})
cur := m.Current()
m.Invalidate("local") // 当前引擎:应为 no-op
if m.Current() != cur {
t.Fatal("Invalidate 当前引擎不应替换实例")
}
}
// TestSwitchSameNameNoop 同名 Switch 幂等。
func TestSwitchSameNameNoop(t *testing.T) {
m := newTestManager(&atomic.Bool{})
s, err := m.Switch("local")
if err != nil {
t.Fatalf("Switch(local) 同名应成功: %v", err)
}
if s != m.Current() {
t.Fatal("同名 Switch 应返回当前实例")
}
}
// TestDelegateToCurrent 保存/读取类操作委托当前引擎(切换后指向新引擎)。
func TestDelegateToCurrent(t *testing.T) {
var s3Fail atomic.Bool
m := newTestManager(&s3Fail)
ctx := context.Background()
// local 引擎 HealthCheck 健康
if err := m.HealthCheck(ctx); err != nil {
t.Fatalf("委托 HealthCheck(local): %v", err)
}
if _, err := m.Switch("s3"); err != nil {
t.Fatalf("Switch(s3): %v", err)
}
if err := m.HealthCheck(ctx); err != nil {
t.Fatalf("委托 HealthCheck(s3): %v", err)
}
}
+70
View File
@@ -0,0 +1,70 @@
package storage
// EngineOptions 引擎构造选项:由 main.goAPI 层任务)从 config KV 填充。
// 各引擎的 RegisterEngine 工厂读取本结构;零值即安全默认。
type EngineOptions struct {
// Local 本地引擎选项。
Local LocalOptions
// S3 S3 引擎选项。
S3 S3Options
// WebDAV WebDAV 引擎选项。
WebDAV WebDAVOptions
}
// LocalOptions 本地引擎配置(对齐 local_storage_path)。
type LocalOptions struct {
// Root 存储根目录;空则使用系统临时目录。
Root string
}
// S3Options S3 引擎配置(对齐 s3_* 配置键)。
type S3Options struct {
AccessKeyID string // s3_access_key_id
SecretAccessKey string // s3_secret_access_key
SessionToken string // aws_session_token
Bucket string // s3_bucket_name
Endpoint string // s3_endpoint_urlMinIO 等;空则 AWS 默认端点)
Region string // s3_region_name,默认 auto
AddressingStyle string // s3_addressing_style: auto|path|virtual
}
// WebDAVOptions WebDAV 引擎配置(对齐 webdav_* 配置键 + 本次优化项)。
type WebDAVOptions struct {
// BaseURL 服务地址,如 https://dav.example.com/dav/。
BaseURL string
// Username/Password 凭据(Basic 与 Digest 共用)。
Username string
Password string
// RootPath 远端根目录(webdav_root_path),会自动逐级创建。
RootPath string
// MaxRetries 5xx/网络错误最大重试次数(指数退避),0 取默认 3。
MaxRetries int
// BaseBackoff 重试基础退避时长,0 取默认 200ms。
BaseBackoff int64
// Timeout 单请求超时秒数,0 取默认 30s。
Timeout int64
// MaxIdleConnsPerHost 连接池每主机最大空闲连接,0 取默认 16(连接复用优化)。
MaxIdleConnsPerHost int
}
// engineOptions 全局引擎选项(由 main.go 注入;默认零值)。
var engineOptions EngineOptions
// SetEngineOptions 注入引擎构造选项(在 RegisterEngine 工厂执行前调用)。
func SetEngineOptions(opts EngineOptions) { engineOptions = opts }
// applyDefaults 填充零值默认项。
func (o *WebDAVOptions) applyDefaults() {
if o.MaxRetries <= 0 {
o.MaxRetries = 3
}
if o.BaseBackoff <= 0 {
o.BaseBackoff = 200
}
if o.Timeout <= 0 {
o.Timeout = 30
}
if o.MaxIdleConnsPerHost <= 0 {
o.MaxIdleConnsPerHost = 16
}
}
+103
View File
@@ -0,0 +1,103 @@
package storage
import (
"path"
"strings"
)
// ChunkDir 实现默认分片目录约定:<父目录>/chunks/<uploadID>。
// local/s3/webdav 三引擎共用,保持分片路径一致。
func ChunkDir(savePath, uploadID string) string {
dir := path.Dir(savePath)
name := path.Base(savePath)
// 防御:savePath 非法时仍返回明确结构,具体引擎再做安全校验
if name == "." || name == "/" {
name = "file"
}
return path.Join(dir, "chunks", uploadID) + "/" + name
}
// ChunkPartPath 分片对象完整路径(相对存储根)。
func ChunkPartPath(savePath, uploadID string, index int) string {
dir := path.Dir(savePath)
return path.Join(dir, "chunks", uploadID, itoa(index)+".part")
}
// SanitizePath 清理相对路径:统一斜杠、去首尾斜杠、拒绝 .. 穿越。
// 返回清理后的相对路径与是否合法。
func SanitizePath(p string) (string, bool) {
raw := strings.ReplaceAll(strings.TrimSpace(p), "\\", "/")
raw = strings.TrimPrefix(raw, "/")
if raw == "" {
return "", false
}
cleaned := path.Clean(raw)
if cleaned == ".." || strings.HasPrefix(cleaned, "../") || path.IsAbs(cleaned) {
return "", false
}
// 拒绝任何单独的 .. 段
for _, seg := range strings.Split(cleaned, "/") {
if seg == ".." {
return "", false
}
}
return cleaned, true
}
// SanitizeFileName 清理文件名:剥离路径、替换非法字符、限制长度。
// 对齐参考 core/utils.py 的 sanitize_filename。
func SanitizeFileName(name string) string {
// 剥离路径
if idx := strings.LastIndexAny(name, "/\\"); idx >= 0 {
name = name[idx+1:]
}
var b strings.Builder
for _, r := range name {
switch {
case r < 0x20 || r == 0x7f:
b.WriteByte('_')
case strings.ContainsRune(`\*?:"<>|`, r):
b.WriteByte('_')
case r == ' ':
b.WriteByte('_')
default:
b.WriteRune(r)
}
}
cleaned := b.String()
// 压缩连续下划线
for strings.Contains(cleaned, "__") {
cleaned = strings.ReplaceAll(cleaned, "__", "_")
}
cleaned = strings.Trim(cleaned, "._")
if cleaned == "" {
return "unnamed_file"
}
if len(cleaned) > 255 {
cleaned = cleaned[:255]
}
return cleaned
}
// itoa 小整数转字符串。
func itoa(n int) string {
if n == 0 {
return "0"
}
neg := n < 0
if neg {
n = -n
}
var buf [21]byte
i := len(buf)
for n > 0 {
i--
buf[i] = byte('0' + n%10)
n /= 10
}
if neg {
i--
buf[i] = '-'
}
return string(buf[i:])
}
+74
View File
@@ -0,0 +1,74 @@
package storage
import (
"crypto/sha256"
"encoding/hex"
"io"
"testing"
)
// TestSanitizePath 校验路径穿越防护。
func TestSanitizePath(t *testing.T) {
cases := []struct {
in string
ok bool
out string
}{
{"2025/08/uuid.zip", true, "2025/08/uuid.zip"},
{"/2025/08/uuid.zip", true, "2025/08/uuid.zip"},
{"a\\b\\c.txt", true, "a/b/c.txt"},
{"../etc/passwd", false, ""},
{"a/../../b", false, ""},
{"..", false, ""},
{"", false, ""},
}
for _, tc := range cases {
got, ok := SanitizePath(tc.in)
if ok != tc.ok || (ok && got != tc.out) {
t.Errorf("SanitizePath(%q) = (%q, %v), want (%q, %v)", tc.in, got, ok, tc.out, tc.ok)
}
}
}
// TestSanitizeFileName 校验文件名清理。
func TestSanitizeFileName(t *testing.T) {
cases := []struct{ in, want string }{
{"hello world.zip", "hello_world.zip"},
{"/path/to/file.txt", "file.txt"},
{"a<b>:c?.mp4", "a_b_c_.mp4"}, // 连续下划线压缩,对齐参考 re.sub(r"_+", "_")
{"", "unnamed_file"},
{"__..__", "unnamed_file"},
}
for _, tc := range cases {
if got := SanitizeFileName(tc.in); got != tc.want {
t.Errorf("SanitizeFileName(%q) = %q, want %q", tc.in, got, tc.want)
}
}
}
// TestChunkPartPath 校验分片路径约定。
func TestChunkPartPath(t *testing.T) {
got := ChunkPartPath("2025/08/uuid.zip", "upload-1", 3)
want := "2025/08/chunks/upload-1/3.part"
if got != want {
t.Errorf("ChunkPartPath = %q, want %q", got, want)
}
}
// TestChunkDir 校验分片目录约定。
func TestChunkDir(t *testing.T) {
got := ChunkDir("2025/08/uuid.zip", "upload-1")
want := "2025/08/chunks/upload-1/uuid.zip"
if got != want {
t.Errorf("ChunkDir = %q, want %q", got, want)
}
}
// TestSHA256Helper 辅助:确认 sha256 用法一致(合并校验依赖)。
func TestSHA256Helper(t *testing.T) {
h := sha256.Sum256([]byte("abc"))
if got := hex.EncodeToString(h[:]); got != "ba7816bf8f01cfea414140de5dae2223b00361a396177a9cb410ff61f20015ad" {
t.Errorf("sha256(abc) = %s", got)
}
_ = io.EOF
}
+650
View File
@@ -0,0 +1,650 @@
package storage
import (
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"fmt"
"io"
"net"
"net/http"
"os"
"strconv"
"strings"
"time"
"github.com/aws/aws-sdk-go-v2/aws"
awshttp "github.com/aws/aws-sdk-go-v2/aws/transport/http"
"github.com/aws/aws-sdk-go-v2/config"
"github.com/aws/aws-sdk-go-v2/credentials"
"github.com/aws/aws-sdk-go-v2/feature/s3/manager"
"github.com/aws/aws-sdk-go-v2/service/s3"
"github.com/aws/aws-sdk-go-v2/service/s3/types"
"github.com/aws/smithy-go"
)
// S3Storage 基于 aws-sdk-go-v2 的 S3 兼容对象存储引擎(AWS / MinIO / R2 / OSS 等)。
//
// 相比参考实现(S3FileStorageaioboto3)的改进:
// - 单例客户端 + 自定义连接池 Transport(参考实现每次操作新建 session);
// - SaveFile 走 manager.Uploader 分片并发上传(未知长度也可流式,内存占用 ≤ partSize);
// - SaveChunk 落本地临时文件获取精确 Content-Length(参考实现整块读入内存);
// - MergeChunks 用 S3 原生 multipart 流式合并,边读边校验哈希(不落盘、不整块进内存);
// - 5xx/网络错误由 SDK 内置指数退避重试器处理(可配次数)。
type S3Storage struct {
client *s3.Client
presigner *s3.PresignClient
uploader *manager.Uploader
bucket string
}
// NewS3Storage 构造 S3 引擎。
func NewS3Storage(opts S3Options) (*S3Storage, error) {
if strings.TrimSpace(opts.Bucket) == "" {
return nil, fmt.Errorf("storage/s3: 缺少 bucket 配置(s3_bucket_name")
}
region := strings.TrimSpace(opts.Region)
if region == "" {
region = "us-east-1"
}
loadOpts := []func(*config.LoadOptions) error{
config.WithRegion(region),
config.WithCredentialsProvider(credentials.NewStaticCredentialsProvider(
opts.AccessKeyID, opts.SecretAccessKey, opts.SessionToken,
)),
// SDK 内置重试器:标准模式,指数退避 + 抖动,覆盖 5xx 与网络错误。
config.WithRetryMaxAttempts(3),
// 兼容性:仅协议要求时才计算校验和。默认的 trailing CRC32 需要可重放流
// 或 TLSMinIO/R2 等自建端点通常不需要,关闭后 MergeChunks 的
// GET→UploadPart 纯流式转发才能工作。
config.WithRequestChecksumCalculation(aws.RequestChecksumCalculationWhenRequired),
config.WithResponseChecksumValidation(aws.ResponseChecksumValidationWhenRequired),
}
if ep := strings.TrimSpace(opts.Endpoint); ep != "" {
loadOpts = append(loadOpts, config.WithBaseEndpoint(ep))
}
awsCfg, err := config.LoadDefaultConfig(context.Background(), loadOpts...)
if err != nil {
return nil, fmt.Errorf("storage/s3: 初始化 SDK 配置失败: %w", err)
}
client := s3.NewFromConfig(awsCfg, func(o *s3.Options) {
// 寻址风格:path 显式启用;auto 时自定义端点(自建 MinIO 等)默认 path-style。
switch strings.ToLower(strings.TrimSpace(opts.AddressingStyle)) {
case "path":
o.UsePathStyle = true
case "virtual":
o.UsePathStyle = false
default: // auto
o.UsePathStyle = strings.TrimSpace(opts.Endpoint) != ""
}
// 连接复用:自定义 Transport 连接池。
o.HTTPClient = newPooledHTTPClient()
})
st := &S3Storage{
client: client,
presigner: s3.NewPresignClient(client),
bucket: opts.Bucket,
}
st.uploader = manager.NewUploader(client, func(u *manager.Uploader) {
u.PartSize = 5 * 1024 * 1024 // 5MBS3 multipart 最小分片
u.Concurrency = 4
u.LeavePartsOnError = false
})
return st, nil
}
func init() {
RegisterEngine("s3", func(ctx context.Context) (Storage, error) {
return NewS3Storage(engineOptions.S3)
})
}
// newPooledHTTPClient 供 SDK 使用的连接池化 HTTP 客户端。
func newPooledHTTPClient() *http.Client {
return &http.Client{
Transport: &http.Transport{
Proxy: http.ProxyFromEnvironment,
DialContext: (&net.Dialer{
Timeout: 10 * time.Second,
KeepAlive: 30 * time.Second,
}).DialContext,
ForceAttemptHTTP2: true,
MaxIdleConns: 100,
MaxIdleConnsPerHost: 16,
IdleConnTimeout: 90 * time.Second,
TLSHandshakeTimeout: 10 * time.Second,
ExpectContinueTimeout: time.Second,
ResponseHeaderTimeout: 60 * time.Second,
},
}
}
// key 校验并规范化对象键(拒绝穿越,统一斜杠)。
func (s *S3Storage) key(savePath string) (string, error) {
cleaned, ok := SanitizePath(savePath)
if !ok {
return "", fmt.Errorf("%w: %q", ErrInvalidPath, savePath)
}
return cleaned, nil
}
// SaveFile 流式保存:manager.Uploader 按需分片并发上传,内存占用恒定。
func (s *S3Storage) SaveFile(ctx context.Context, r io.Reader, savePath string) (int64, error) {
key, err := s.key(savePath)
if err != nil {
return 0, err
}
src := &countingReader{r: r}
_, err = s.uploader.Upload(ctx, &s3.PutObjectInput{
Bucket: aws.String(s.bucket),
Key: aws.String(key),
Body: src,
ContentType: aws.String("application/octet-stream"),
})
if err != nil {
return src.count(), mapS3Error(err, "PutObject")
}
return src.count(), nil
}
// DeleteFile 删除对象;S3 对不存在的键也返回成功。
func (s *S3Storage) DeleteFile(ctx context.Context, savePath string) error {
key, err := s.key(savePath)
if err != nil {
return err
}
_, err = s.client.DeleteObject(ctx, &s3.DeleteObjectInput{
Bucket: aws.String(s.bucket),
Key: aws.String(key),
})
return mapS3Error(err, "DeleteObject")
}
// Open 获取下载流:Range 直接透传为 GetObject Range 头(流式,不落盘)。
func (s *S3Storage) Open(ctx context.Context, savePath string, rng *Range) (*Download, error) {
key, err := s.key(savePath)
if err != nil {
return nil, err
}
input := &s3.GetObjectInput{
Bucket: aws.String(s.bucket),
Key: aws.String(key),
}
if rng != nil {
input.Range = aws.String(rangeHeaderValue(rng))
}
out, err := s.client.GetObject(ctx, input)
if err != nil {
return nil, mapS3Error(err, "GetObject")
}
size := aws.ToInt64(out.ContentLength)
start, end := int64(0), size-1
if cr := aws.ToString(out.ContentRange); cr != "" { // 服务端按 206 返回了区间
if sr, e, total, ok := parseContentRange(cr); ok {
start, end = sr, e
if total >= 0 {
size = total
}
}
}
return &Download{
ReadCloser: out.Body,
Start: start,
End: end,
Total: size,
Meta: FileMeta{
Size: size,
ContentType: aws.ToString(out.ContentType),
AcceptRanges: true,
},
}, nil
}
// rangeHeaderValue 将 Range 结构转为 HTTP Range 头值。
func rangeHeaderValue(rng *Range) string {
if rng.End < 0 {
return fmt.Sprintf("bytes=%d-", rng.Start)
}
return fmt.Sprintf("bytes=%d-%d", rng.Start, rng.End)
}
// parseContentRange 解析 "bytes 0-99/1000"total 可能为 "*")。
func parseContentRange(v string) (start, end, total int64, ok bool) {
v = strings.TrimSpace(v)
if !strings.HasPrefix(v, "bytes ") {
return 0, 0, -1, false
}
parts := strings.SplitN(strings.TrimPrefix(v, "bytes "), "/", 2)
if len(parts) != 2 {
return 0, 0, -1, false
}
total = -1
if parts[1] != "*" {
t, err := strconv.ParseInt(parts[1], 10, 64)
if err != nil {
return 0, 0, -1, false
}
total = t
}
se := strings.SplitN(parts[0], "-", 2)
if len(se) != 2 {
return 0, 0, -1, false
}
s0, err1 := strconv.ParseInt(se[0], 10, 64)
e0, err2 := strconv.ParseInt(se[1], 10, 64)
if err1 != nil || err2 != nil {
return 0, 0, -1, false
}
return s0, e0, total, true
}
// Stat 获取对象元信息(HeadObject)。
func (s *S3Storage) Stat(ctx context.Context, savePath string) (*FileMeta, error) {
key, err := s.key(savePath)
if err != nil {
return nil, err
}
out, err := s.client.HeadObject(ctx, &s3.HeadObjectInput{
Bucket: aws.String(s.bucket),
Key: aws.String(key),
})
if err != nil {
return nil, mapS3Error(err, "HeadObject")
}
return &FileMeta{
Size: aws.ToInt64(out.ContentLength),
ContentType: aws.ToString(out.ContentType),
AcceptRanges: true,
}, nil
}
// SaveChunk 保存分片对象:落临时文件获取精确长度后 PutObject(可重试)。
func (s *S3Storage) SaveChunk(ctx context.Context, uploadID string, chunkIndex int, r io.Reader, savePath string) (int64, error) {
if _, err := s.key(savePath); err != nil {
return 0, err
}
key, err := s.key(ChunkPartPath(savePath, uploadID, chunkIndex))
if err != nil {
return 0, err
}
// 落临时文件:获得精确 Content-Length 与可重放 Body(网络失败可安全重试)。
tmp, err := os.CreateTemp("", "fcb-s3-chunk-*")
if err != nil {
return 0, fmt.Errorf("storage/s3: 创建临时文件失败: %w", err)
}
tmpName := tmp.Name()
defer func() { _ = tmp.Close(); _ = os.Remove(tmpName) }()
size, err := io.Copy(tmp, r)
if err != nil {
return 0, fmt.Errorf("storage/s3: 缓存分片失败: %w", err)
}
if _, err := tmp.Seek(0, io.SeekStart); err != nil {
return 0, fmt.Errorf("storage/s3: 回卷分片失败: %w", err)
}
_, err = s.client.PutObject(ctx, &s3.PutObjectInput{
Bucket: aws.String(s.bucket),
Key: aws.String(key),
Body: tmp,
ContentLength: aws.Int64(size),
ContentType: aws.String("application/octet-stream"),
})
if err != nil {
return size, mapS3Error(err, "PutObject(分片)")
}
return size, nil
}
// S3 multipart 最小分片限制:除最后一片外每片 ≥5MB,否则 Complete 返回 EntityTooSmall。
// 分片上传的分块大小由服务端配置保证(建议 ≥5MB)。
const s3MinPartSize = 5 * 1024 * 1024
// MergeChunks 用 S3 原生 multipart 流式合并:
// 逐分片 GET → 边流边算哈希 → UploadPart(带精确 Content-Length)→ Complete。
// 任一步失败即 Abort 并返回错误;成功后清理分片对象。
func (s *S3Storage) MergeChunks(ctx context.Context, uploadID string, total int, verifyHash func(index int) (string, error), savePath string) (int64, string, error) {
if total <= 0 {
return 0, "", fmt.Errorf("storage/s3: 非法分片总数 %d", total)
}
key, err := s.key(savePath)
if err != nil {
return 0, "", err
}
chunkPrefix, err := s.key(chunkDirOf(savePath, uploadID))
if err != nil {
return 0, "", err
}
// 单分片快速路径:直接流式 PutObject,绕过 multipart 的 5MB 限制。
if total == 1 {
return s.mergeSingle(ctx, chunkPrefix+"/0.part", key, verifyHash, 0)
}
mpu, err := s.client.CreateMultipartUpload(ctx, &s3.CreateMultipartUploadInput{
Bucket: aws.String(s.bucket),
Key: aws.String(key),
ContentType: aws.String("application/octet-stream"),
})
if err != nil {
return 0, "", mapS3Error(err, "CreateMultipartUpload")
}
_ = aws.ToString(mpu.UploadId) // S3 侧 multipart 会话 IDAbort 时复用 mpu.UploadId
size := int64(0)
totalHash := sha256.New()
parts := make([]types.CompletedPart, 0, total)
defer func() {
// 出错时取消 multipart(避免残留分片产生存储费用)。
if len(parts) < total {
_, _ = s.client.AbortMultipartUpload(ctx, &s3.AbortMultipartUploadInput{
Bucket: aws.String(s.bucket),
Key: aws.String(key),
UploadId: mpu.UploadId,
})
}
}()
for i := 0; i < total; i++ {
if err := ctx.Err(); err != nil {
return 0, "", err
}
var expected string
if verifyHash != nil {
expected, err = verifyHash(i)
if err != nil {
return 0, "", err
}
}
getOut, err := s.client.GetObject(ctx, &s3.GetObjectInput{
Bucket: aws.String(s.bucket),
Key: aws.String(fmt.Sprintf("%s/%d.part", chunkPrefix, i)),
})
if err != nil {
return 0, "", mapS3Error(err, fmt.Sprintf("GetObject(分片 %d)", i))
}
// 分片先流式落临时文件:计算哈希 + 获得可回卷 body(SDK 签名哈希需要 seekable 流,
// 同时为 UploadPart 失败重试保留数据)。
chunkHash := sha256.New()
tmp, err := os.CreateTemp("", "fcb-s3-part-*")
if err != nil {
_ = getOut.Body.Close()
return 0, "", fmt.Errorf("storage/s3: 创建分片临时文件失败: %w", err)
}
partLen, err := io.Copy(io.MultiWriter(tmp, totalHash, chunkHash), getOut.Body)
_ = getOut.Body.Close()
if err != nil {
_ = tmp.Close()
_ = os.Remove(tmp.Name())
return 0, "", fmt.Errorf("storage/s3: 读取分片 %d 失败: %w", i, err)
}
if expected != "" && expected != hex.EncodeToString(chunkHash.Sum(nil)) {
_ = tmp.Close()
_ = os.Remove(tmp.Name())
return 0, "", fmt.Errorf("%w: 分片 %d 期望 %s", ErrHashMismatch, i, expected)
}
if _, err := tmp.Seek(0, io.SeekStart); err != nil {
_ = tmp.Close()
_ = os.Remove(tmp.Name())
return 0, "", fmt.Errorf("storage/s3: 回卷分片 %d 失败: %w", i, err)
}
up, err := s.client.UploadPart(ctx, &s3.UploadPartInput{
Bucket: aws.String(s.bucket),
Key: aws.String(key),
UploadId: mpu.UploadId,
PartNumber: aws.Int32(int32(i + 1)),
Body: tmp,
ContentLength: aws.Int64(partLen),
})
_ = tmp.Close()
_ = os.Remove(tmp.Name())
if err != nil {
return 0, "", mapS3Error(err, fmt.Sprintf("UploadPart(分片 %d)", i))
}
parts = append(parts, types.CompletedPart{
PartNumber: aws.Int32(int32(i + 1)),
ETag: up.ETag,
})
size += partLen
}
if _, err := s.client.CompleteMultipartUpload(ctx, &s3.CompleteMultipartUploadInput{
Bucket: aws.String(s.bucket),
Key: aws.String(key),
UploadId: mpu.UploadId,
MultipartUpload: &types.CompletedMultipartUpload{Parts: parts},
}); err != nil {
return 0, "", mapS3Error(err, "CompleteMultipartUpload")
}
// 合并成功后清理分片对象(静默容错)。
_ = s.CleanChunks(ctx, uploadID, savePath)
return size, hex.EncodeToString(totalHash.Sum(nil)), nil
}
// mergeSingle 单分片合并快速路径:GET 分片 → 落临时文件校验 → PutObject 正式键。
func (s *S3Storage) mergeSingle(ctx context.Context, chunkKey, dstKey string, verifyHash func(index int) (string, error), index int) (int64, string, error) {
getOut, err := s.client.GetObject(ctx, &s3.GetObjectInput{
Bucket: aws.String(s.bucket),
Key: aws.String(chunkKey),
})
if err != nil {
return 0, "", mapS3Error(err, "GetObject(分片)")
}
defer func() { _ = getOut.Body.Close() }()
tmp, err := os.CreateTemp("", "fcb-s3-merge-*")
if err != nil {
return 0, "", fmt.Errorf("storage/s3: 创建临时文件失败: %w", err)
}
tmpName := tmp.Name()
defer func() { _ = tmp.Close(); _ = os.Remove(tmpName) }()
fileHash := sha256.New()
size, err := io.Copy(io.MultiWriter(tmp, fileHash), getOut.Body)
if err != nil {
return 0, "", fmt.Errorf("storage/s3: 读取分片失败: %w", err)
}
if verifyHash != nil {
expected, err := verifyHash(index)
if err != nil {
return 0, "", err
}
if expected != "" && expected != hex.EncodeToString(fileHash.Sum(nil)) {
return 0, "", fmt.Errorf("%w: 分片 %d 期望 %s", ErrHashMismatch, index, expected)
}
}
if _, err := tmp.Seek(0, io.SeekStart); err != nil {
return 0, "", fmt.Errorf("storage/s3: 回卷临时文件失败: %w", err)
}
if _, err := s.client.PutObject(ctx, &s3.PutObjectInput{
Bucket: aws.String(s.bucket),
Key: aws.String(dstKey),
Body: tmp,
ContentLength: aws.Int64(size),
ContentType: aws.String("application/octet-stream"),
}); err != nil {
return 0, "", mapS3Error(err, "PutObject(合并)")
}
return size, hex.EncodeToString(fileHash.Sum(nil)), nil
}
// CleanChunks 列举并批量删除分片对象;前缀不存在时静默成功。
func (s *S3Storage) CleanChunks(ctx context.Context, uploadID string, savePath string) error {
prefix, err := s.key(chunkDirOf(savePath, uploadID))
if err != nil {
return err
}
prefix += "/"
paginator := s3.NewListObjectsV2Paginator(s.client, &s3.ListObjectsV2Input{
Bucket: aws.String(s.bucket),
Prefix: aws.String(prefix),
})
for paginator.HasMorePages() {
page, err := paginator.NextPage(ctx)
if err != nil {
return mapS3Error(err, "ListObjectsV2(分片)")
}
if len(page.Contents) == 0 {
return nil
}
objs := make([]types.ObjectIdentifier, 0, len(page.Contents))
for _, obj := range page.Contents {
objs = append(objs, types.ObjectIdentifier{Key: obj.Key})
}
if _, err := s.client.DeleteObjects(ctx, &s3.DeleteObjectsInput{
Bucket: aws.String(s.bucket),
Delete: &types.Delete{Objects: objs, Quiet: aws.Bool(true)},
}); err != nil {
return mapS3Error(err, "DeleteObjects(分片)")
}
}
return nil
}
// FileExists HeadObject 探测存在性。
func (s *S3Storage) FileExists(ctx context.Context, savePath string) (bool, error) {
key, err := s.key(savePath)
if err != nil {
return false, err
}
_, err = s.client.HeadObject(ctx, &s3.HeadObjectInput{
Bucket: aws.String(s.bucket),
Key: aws.String(key),
})
if err != nil {
if isS3NotFound(err) {
return false, nil
}
return false, mapS3Error(err, "HeadObject")
}
return true, nil
}
// PresignGetURL 生成限时下载直链。
func (s *S3Storage) PresignGetURL(ctx context.Context, savePath string, expires int64) (string, error) {
key, err := s.key(savePath)
if err != nil {
return "", err
}
if expires <= 0 {
expires = 3600
}
out, err := s.presigner.PresignGetObject(ctx, &s3.GetObjectInput{
Bucket: aws.String(s.bucket),
Key: aws.String(key),
}, s3.WithPresignExpires(time.Duration(expires)*time.Second))
if err != nil {
return "", mapS3Error(err, "PresignGetObject")
}
return out.URL, nil
}
// PresignPutURL 生成限时直传 URL。
func (s *S3Storage) PresignPutURL(ctx context.Context, savePath string, expires int64) (string, error) {
key, err := s.key(savePath)
if err != nil {
return "", err
}
if expires <= 0 {
expires = 900
}
out, err := s.presigner.PresignPutObject(ctx, &s3.PutObjectInput{
Bucket: aws.String(s.bucket),
Key: aws.String(key),
}, s3.WithPresignExpires(time.Duration(expires)*time.Second))
if err != nil {
return "", mapS3Error(err, "PresignPutObject")
}
return out.URL, nil
}
// HealthCheck 健康检查:列举 bucketMaxKeys=1),同时校验连通性、凭据与 bucket 存在。
func (s *S3Storage) HealthCheck(ctx context.Context) error {
maxKeys := int32(1)
_, err := s.client.ListObjectsV2(ctx, &s3.ListObjectsV2Input{
Bucket: aws.String(s.bucket),
MaxKeys: aws.Int32(maxKeys),
Prefix: aws.String(""),
})
if err != nil {
return fmt.Errorf("%w: S3 健康检查失败: %v", ErrUnavailable, err)
}
return nil
}
// ReadHead 读取对象前 n 字节(保留的便捷封装:HeadMeta 的仅头部形态)。
func (s *S3Storage) ReadHead(ctx context.Context, savePath string, n int64) ([]byte, error) {
_, head, err := s.HeadMeta(ctx, savePath, n)
return head, err
}
// HeadMeta 读取对象元信息与头部字节(S3 引擎实现:HeadObject + Range GET)。
func (s *S3Storage) HeadMeta(ctx context.Context, savePath string, headBytes int64) (*FileMeta, []byte, error) {
key, err := s.key(savePath)
if err != nil {
return nil, nil, err
}
head, err := s.headBytes(ctx, s.bucket, key, headBytes)
if err != nil {
return nil, nil, err
}
meta, err := s.Stat(ctx, savePath)
if err != nil {
return nil, nil, err
}
return meta, head, nil
}
// headBytes 通过 Range GET 读取对象前 n 字节。
func (s *S3Storage) headBytes(ctx context.Context, bucket, key string, n int64) ([]byte, error) {
if n <= 0 {
return nil, nil
}
out, err := s.client.GetObject(ctx, &s3.GetObjectInput{
Bucket: aws.String(bucket),
Key: aws.String(key),
Range: aws.String(fmt.Sprintf("bytes=0-%d", n-1)),
})
if err != nil {
return nil, mapS3Error(err, "GetObject(head)")
}
defer func() { _ = out.Body.Close() }()
return io.ReadAll(io.LimitReader(out.Body, n))
}
// isS3NotFound 判断错误是否为对象不存在。
func isS3NotFound(err error) bool {
var nf *types.NotFound
if errors.As(err, &nf) {
return true
}
var ae smithy.APIError
if errors.As(err, &ae) {
switch ae.ErrorCode() {
case "NotFound", "NoSuchKey":
return true
}
}
var re *awshttp.ResponseError
if errors.As(err, &re) && re.HTTPStatusCode() == http.StatusNotFound {
return true
}
return false
}
// mapS3Error 将 SDK 错误映射为包内哨兵错误。
func mapS3Error(err error, op string) error {
if err == nil {
return nil
}
if isS3NotFound(err) {
return fmt.Errorf("%w%s", ErrNotFound, op)
}
var re *awshttp.ResponseError
if errors.As(err, &re) && re.HTTPStatusCode() == http.StatusRequestedRangeNotSatisfiable {
return fmt.Errorf("%w%s", ErrRangeNotSatisfiable, op)
}
var ae smithy.APIError
if errors.As(err, &ae) && ae.ErrorCode() == "InvalidRange" {
return fmt.Errorf("%w%s", ErrRangeNotSatisfiable, op)
}
return fmt.Errorf("storage/s3: %s 失败: %w", op, err)
}
// 接口编译期断言。
var _ Storage = (*S3Storage)(nil)
+495
View File
@@ -0,0 +1,495 @@
package storage
import (
"bytes"
"context"
"encoding/xml"
"fmt"
"io"
"net/http"
"net/http/httptest"
"sort"
"strings"
"sync"
"testing"
"time"
)
// ---- 最小 S3 兼容假服务(仅覆盖本引擎用到的 API)----
type fakeS3Upload struct {
key string
parts map[int][]byte
}
type fakeS3 struct {
mu sync.Mutex
objects map[string][]byte
uploads map[string]*fakeS3Upload
nextID int
putCount int
getCount int
headCount int
deleteCount int
listCount int
completeN int
// failNextGet:让接下来 N 次 GET 返回 503(重试测试用)。
failNextGet int
}
func newFakeS3() *fakeS3 {
return &fakeS3{objects: map[string][]byte{}, uploads: map[string]*fakeS3Upload{}}
}
// s3Key 从 path-style 路径剥离 bucket 前缀得到对象键。
func s3Key(r *http.Request) (bucket, key string) {
p := strings.TrimPrefix(r.URL.Path, "/")
if i := strings.Index(p, "/"); i >= 0 {
return p[:i], p[i+1:]
}
return p, ""
}
func s3ErrorXML(w http.ResponseWriter, status int, code, msg string) {
w.Header().Set("Content-Type", "application/xml")
w.WriteHeader(status)
_, _ = w.Write([]byte(fmt.Sprintf(
`<?xml version="1.0" encoding="UTF-8"?><Error><Code>%s</Code><Message>%s</Message></Error>`, code, msg)))
}
func (f *fakeS3) ServeHTTP(w http.ResponseWriter, r *http.Request) {
f.mu.Lock()
defer f.mu.Unlock()
q := r.URL.Query()
bucket, key := s3Key(r)
switch {
// UploadPart
case r.Method == http.MethodPut && q.Get("partNumber") != "" && q.Get("uploadId") != "":
up, ok := f.uploads[q.Get("uploadId")]
if !ok {
s3ErrorXML(w, 404, "NoSuchUpload", "upload not found")
return
}
body, _ := io.ReadAll(r.Body)
var n int
_, _ = fmt.Sscanf(q.Get("partNumber"), "%d", &n)
up.parts[n] = body
w.Header().Set("ETag", fmt.Sprintf(`"part-%d"`, n))
w.WriteHeader(http.StatusOK)
// CreateMultipartUpload
case r.Method == http.MethodPost && q.Has("uploads"):
f.nextID++
id := fmt.Sprintf("mpu-%d", f.nextID)
f.uploads[id] = &fakeS3Upload{key: key, parts: map[int][]byte{}}
w.Header().Set("Content-Type", "application/xml")
_, _ = w.Write([]byte(fmt.Sprintf(
`<?xml version="1.0" encoding="UTF-8"?><InitiateMultipartUploadResult xmlns="http://s3.amazonaws.com/doc/2006-03-01/"><Bucket>%s</Bucket><Key>%s</Key><UploadId>%s</UploadId></InitiateMultipartUploadResult>`,
bucket, key, id)))
// CompleteMultipartUpload
case r.Method == http.MethodPost && q.Get("uploadId") != "":
up, ok := f.uploads[q.Get("uploadId")]
if !ok {
s3ErrorXML(w, 404, "NoSuchUpload", "upload not found")
return
}
// 按 partNumber 有序拼接
nums := make([]int, 0, len(up.parts))
for n := range up.parts {
nums = append(nums, n)
}
sort.Ints(nums)
var merged bytes.Buffer
for _, n := range nums {
merged.Write(up.parts[n])
}
f.objects[up.key] = merged.Bytes()
delete(f.uploads, q.Get("uploadId"))
f.completeN++
w.Header().Set("Content-Type", "application/xml")
_, _ = w.Write([]byte(fmt.Sprintf(
`<?xml version="1.0" encoding="UTF-8"?><CompleteMultipartUploadResult xmlns="http://s3.amazonaws.com/doc/2006-03-01/"><Location>http://%s/%s/%s</Location><Bucket>%s</Bucket><Key>%s</Key><ETag>"merged"</ETag></CompleteMultipartUploadResult>`,
r.Host, bucket, up.key, bucket, up.key)))
// AbortMultipartUpload
case r.Method == http.MethodDelete && q.Get("uploadId") != "":
delete(f.uploads, q.Get("uploadId"))
w.WriteHeader(http.StatusNoContent)
// DeleteObjects(批量)
case r.Method == http.MethodPost && q.Has("delete"):
var req struct {
Objects []struct {
Key string `xml:"Key"`
} `xml:"Object"`
}
_ = xml.NewDecoder(r.Body).Decode(&req)
for _, o := range req.Objects {
delete(f.objects, o.Key)
}
w.Header().Set("Content-Type", "application/xml")
_, _ = w.Write([]byte(`<?xml version="1.0" encoding="UTF-8"?><DeleteResult xmlns="http://s3.amazonaws.com/doc/2006-03-01/"></DeleteResult>`))
// ListObjectsV2
case r.Method == http.MethodGet && q.Get("list-type") == "2":
f.listCount++
prefix := q.Get("prefix")
var body strings.Builder
body.WriteString(`<?xml version="1.0" encoding="UTF-8"?><ListBucketResult xmlns="http://s3.amazonaws.com/doc/2006-03-01/"><Name>` + bucket + `</Name><Prefix>` + prefix + `</Prefix><IsTruncated>false</IsTruncated>`)
keys := make([]string, 0, len(f.objects))
for k := range f.objects {
if strings.HasPrefix(k, prefix) {
keys = append(keys, k)
}
}
sort.Strings(keys)
for _, k := range keys {
body.WriteString(fmt.Sprintf("<Contents><Key>%s</Key><Size>%d</Size></Contents>", k, len(f.objects[k])))
}
body.WriteString("</ListBucketResult>")
w.Header().Set("Content-Type", "application/xml")
_, _ = w.Write([]byte(body.String()))
// PutObject
case r.Method == http.MethodPut:
f.putCount++
body, _ := io.ReadAll(r.Body)
f.objects[key] = body
w.Header().Set("ETag", `"put"`)
w.WriteHeader(http.StatusOK)
// GetObject
case r.Method == http.MethodGet:
if f.failNextGet > 0 {
f.failNextGet--
s3ErrorXML(w, 503, "ServiceUnavailable", "flaky")
return
}
f.getCount++
data, ok := f.objects[key]
if !ok {
s3ErrorXML(w, 404, "NoSuchKey", "not found")
return
}
w.Header().Set("Content-Type", "application/octet-stream")
if rng := r.Header.Get("Range"); rng != "" {
start, end := int64(0), int64(len(data))-1
if _, err := fmt.Sscanf(rng, "bytes=%d-%d", &start, &end); err != nil {
var s int64
if _, err := fmt.Sscanf(rng, "bytes=%d-", &s); err == nil {
start, end = s, int64(len(data))-1
}
}
if start < 0 || start >= int64(len(data)) {
s3ErrorXML(w, 416, "InvalidRange", "range not satisfiable")
return
}
if end >= int64(len(data)) {
end = int64(len(data)) - 1
}
slice := data[start : end+1]
w.Header().Set("Content-Range", fmt.Sprintf("bytes %d-%d/%d", start, end, len(data)))
w.WriteHeader(http.StatusPartialContent)
_, _ = w.Write(slice)
return
}
w.WriteHeader(http.StatusOK)
_, _ = w.Write(data)
// HeadObject
case r.Method == http.MethodHead:
f.headCount++
data, ok := f.objects[key]
if !ok {
w.WriteHeader(http.StatusNotFound)
return
}
w.Header().Set("Content-Type", "application/octet-stream")
w.Header().Set("Content-Length", fmt.Sprintf("%d", len(data)))
w.WriteHeader(http.StatusOK)
// DeleteObject
case r.Method == http.MethodDelete:
f.deleteCount++
delete(f.objects, key)
w.WriteHeader(http.StatusNoContent)
default:
s3ErrorXML(w, 400, "NotImplemented", "unsupported")
}
}
// newTestS3 构造对接假服务的 S3 引擎。
func newTestS3(t *testing.T) (*S3Storage, *fakeS3) {
t.Helper()
f := newFakeS3()
srv := httptest.NewServer(f)
t.Cleanup(srv.Close)
st, err := NewS3Storage(S3Options{
AccessKeyID: "test-ak",
SecretAccessKey: "test-sk",
Bucket: "test-bucket",
Endpoint: srv.URL,
Region: "us-east-1",
AddressingStyle: "path",
})
if err != nil {
t.Fatalf("NewS3Storage: %v", err)
}
return st, f
}
// TestS3SaveStatOpenRange 保存/元信息/完整与 Range 下载。
func TestS3SaveStatOpenRange(t *testing.T) {
st, f := newTestS3(t)
ctx := context.Background()
data := []byte("0123456789abcdef S3 引擎测试数据")
n, err := st.SaveFile(ctx, bytes.NewReader(data), "2025/08/s3.bin")
if err != nil {
t.Fatalf("SaveFile: %v", err)
}
if n != int64(len(data)) {
t.Fatalf("n = %d", n)
}
if got := f.objects["2025/08/s3.bin"]; !bytes.Equal(got, data) {
t.Fatalf("stored mismatch")
}
meta, err := st.Stat(ctx, "2025/08/s3.bin")
if err != nil {
t.Fatalf("Stat: %v", err)
}
if meta.Size != int64(len(data)) || !meta.AcceptRanges {
t.Fatalf("Stat = %+v", meta)
}
dl, err := st.Open(ctx, "2025/08/s3.bin", nil)
if err != nil {
t.Fatalf("Open: %v", err)
}
got, _ := io.ReadAll(dl)
_ = dl.Close()
if !bytes.Equal(got, data) {
t.Fatalf("full content mismatch")
}
if dl.Start != 0 || dl.End != int64(len(data))-1 || dl.Total != int64(len(data)) {
t.Fatalf("full offsets = %d,%d,%d", dl.Start, dl.End, dl.Total)
}
dl, err = st.Open(ctx, "2025/08/s3.bin", &Range{Start: 4, End: 9})
if err != nil {
t.Fatalf("Open range: %v", err)
}
got, _ = io.ReadAll(dl)
_ = dl.Close()
if !bytes.Equal(got, data[4:10]) {
t.Fatalf("range mismatch: %q", got)
}
if dl.Start != 4 || dl.End != 9 || dl.Total != int64(len(data)) {
t.Fatalf("range offsets = %d,%d,%d", dl.Start, dl.End, dl.Total)
}
// 404 / 416
if _, err := st.Open(ctx, "missing.bin", nil); err == nil || !strings.Contains(err.Error(), ErrNotFound.Error()) {
t.Fatalf("want ErrNotFound, got %v", err)
}
if _, err := st.Open(ctx, "2025/08/s3.bin", &Range{Start: int64(len(data)) + 5, End: -1}); err == nil ||
!strings.Contains(err.Error(), ErrRangeNotSatisfiable.Error()) {
t.Fatalf("want ErrRangeNotSatisfiable, got %v", err)
}
}
// TestS3DeleteExists 删除与存在性。
func TestS3DeleteExists(t *testing.T) {
st, _ := newTestS3(t)
ctx := context.Background()
if _, err := st.SaveFile(ctx, strings.NewReader("x"), "del.bin"); err != nil {
t.Fatal(err)
}
ok, err := st.FileExists(ctx, "del.bin")
if err != nil || !ok {
t.Fatalf("exists = %v %v", ok, err)
}
if err := st.DeleteFile(ctx, "del.bin"); err != nil {
t.Fatalf("DeleteFile: %v", err)
}
ok, err = st.FileExists(ctx, "del.bin")
if err != nil || ok {
t.Fatalf("after delete exists = %v %v", ok, err)
}
}
// TestS3ChunkMergeMulti 多分片合并:原生 multipart + 哈希校验 + 分片清理。
func TestS3ChunkMergeMulti(t *testing.T) {
st, f := newTestS3(t)
ctx := context.Background()
savePath := "2025/09/s3-chunked.bin"
uploadID := "uid-s3"
chunks := [][]byte{bytes.Repeat([]byte("A"), 6*1024*1024/3), []byte("BBBB"), []byte("CC")}
// 注意:multipart 除最后一片需 ≥5MB;此处只验证代码路径,真实约束由部署配置保证。
// 为避免 EntityTooSmall,将第一片放大:
chunks[0] = bytes.Repeat([]byte("A"), 5*1024*1024)
hashes := make([]string, len(chunks))
var total int64
for i, c := range chunks {
n, err := st.SaveChunk(ctx, uploadID, i, bytes.NewReader(c), savePath)
if err != nil {
t.Fatalf("SaveChunk %d: %v", i, err)
}
hashes[i] = sha256Hex(c)
total += int64(len(c))
_ = n
}
size, fileHash, err := st.MergeChunks(ctx, uploadID, len(chunks), func(i int) (string, error) {
return hashes[i], nil
}, savePath)
if err != nil {
t.Fatalf("MergeChunks: %v", err)
}
if size != total {
t.Fatalf("size = %d want %d", size, total)
}
if fileHash != sha256Hex(bytes.Join(chunks, nil)) {
t.Fatalf("file hash mismatch")
}
merged := f.objects[savePath]
if !bytes.Equal(merged, bytes.Join(chunks, nil)) {
t.Fatalf("merged object mismatch (len=%d)", len(merged))
}
if f.completeN != 1 {
t.Fatalf("CompleteMultipartUpload 次数 = %d", f.completeN)
}
// 分片对象已清理
for k := range f.objects {
if strings.Contains(k, "chunks/"+uploadID) {
t.Fatalf("分片对象残留: %s", k)
}
}
}
// TestS3ChunkMergeSingle 单分片快速路径。
func TestS3ChunkMergeSingle(t *testing.T) {
st, f := newTestS3(t)
ctx := context.Background()
data := []byte("single-chunk")
if _, err := st.SaveChunk(ctx, "uid1", 0, bytes.NewReader(data), "one.bin"); err != nil {
t.Fatal(err)
}
size, fileHash, err := st.MergeChunks(ctx, "uid1", 1, func(i int) (string, error) {
return sha256Hex(data), nil
}, "one.bin")
if err != nil {
t.Fatalf("MergeChunks: %v", err)
}
if size != int64(len(data)) || fileHash != sha256Hex(data) {
t.Fatalf("size/hash mismatch")
}
if !bytes.Equal(f.objects["one.bin"], data) {
t.Fatalf("object mismatch")
}
}
// TestS3CleanChunks 清理残留分片。
func TestS3CleanChunks(t *testing.T) {
st, f := newTestS3(t)
ctx := context.Background()
savePath := "clean.bin"
for i := 0; i < 3; i++ {
if _, err := st.SaveChunk(ctx, "uidc", i, strings.NewReader("zz"), savePath); err != nil {
t.Fatal(err)
}
}
prefix := chunkDirOf(savePath, "uidc") + "/"
count := 0
for k := range f.objects {
if strings.HasPrefix(k, prefix) {
count++
}
}
if count != 3 {
t.Fatalf("期望 3 个分片对象,实际 %d", count)
}
if err := st.CleanChunks(ctx, "uidc", savePath); err != nil {
t.Fatalf("CleanChunks: %v", err)
}
for k := range f.objects {
if strings.HasPrefix(k, prefix) {
t.Fatalf("分片未清理: %s", k)
}
}
// 幂等:再清理一次不报错
if err := st.CleanChunks(ctx, "uidc", savePath); err != nil {
t.Fatalf("CleanChunks idempotent: %v", err)
}
}
// TestS3Presign 预签名 URL 生成。
func TestS3Presign(t *testing.T) {
st, _ := newTestS3(t)
ctx := context.Background()
getURL, err := st.PresignGetURL(ctx, "presign.bin", 600)
if err != nil {
t.Fatalf("PresignGetURL: %v", err)
}
if !strings.Contains(getURL, "X-Amz-Signature") || !strings.Contains(getURL, "X-Amz-Expires=600") {
t.Fatalf("GET 直链缺少签名参数: %s", getURL)
}
putURL, err := st.PresignPutURL(ctx, "presign.bin", 300)
if err != nil {
t.Fatalf("PresignPutURL: %v", err)
}
if !strings.Contains(putURL, "X-Amz-Signature") {
t.Fatalf("PUT 直链缺少签名参数: %s", putURL)
}
}
// TestS3HealthCheck 健康检查(ListObjectsV2)。
func TestS3HealthCheck(t *testing.T) {
st, _ := newTestS3(t)
if err := st.HealthCheck(context.Background()); err != nil {
t.Fatalf("HealthCheck: %v", err)
}
}
// TestS3RetryOn503 SDK 内置重试器:503 后成功。
func TestS3RetryOn503(t *testing.T) {
st, f := newTestS3(t)
ctx := context.Background()
data := []byte("retry-me")
if _, err := st.SaveFile(ctx, bytes.NewReader(data), "retry.bin"); err != nil {
t.Fatal(err)
}
f.mu.Lock()
f.failNextGet = 1
f.mu.Unlock()
dl, err := st.Open(ctx, "retry.bin", nil)
if err != nil {
t.Fatalf("503 后应重试成功: %v", err)
}
got, _ := io.ReadAll(dl)
_ = dl.Close()
if !bytes.Equal(got, data) {
t.Fatalf("content mismatch")
}
}
// TestS3FactoryRegistry 工厂构造。
func TestS3FactoryRegistry(t *testing.T) {
f := newFakeS3()
srv := httptest.NewServer(f)
defer srv.Close()
prev := engineOptions.S3
engineOptions.S3 = S3Options{
AccessKeyID: "ak", SecretAccessKey: "sk", Bucket: "b",
Endpoint: srv.URL, Region: "us-east-1", AddressingStyle: "path",
}
defer func() { engineOptions.S3 = prev }()
st, err := NewEngine(context.Background(), "s3")
if err != nil {
t.Fatalf("NewEngine(s3): %v", err)
}
if err := st.HealthCheck(context.Background()); err != nil {
t.Fatalf("HealthCheck: %v", err)
}
if _, err := st.PresignGetURL(context.Background(), "x.bin", 60); err != nil {
t.Fatalf("Presign: %v", err)
}
_ = time.Now
}
+889
View File
@@ -0,0 +1,889 @@
package storage
import (
"context"
"crypto/sha256"
"encoding/hex"
"encoding/xml"
"errors"
"fmt"
"io"
"math/rand"
"net"
"net/http"
"net/url"
"os"
"path"
"strconv"
"strings"
"sync"
"time"
)
// WebDAVStorage 基于 net/http 的 WebDAV 引擎(本次重写的重点优化对象)。
//
// 相比参考实现(WebDAVFileStorageaiohttp)的改进:
// - 单例 http.Client + 连接池化 Transport(参考实现每个操作新建 ClientSession,无复用);
// - Basic 与 DigestRFC 2617qop=authMD5/SHA-256)双认证自动协商(参考实现仅 Basic);
// - GET 下载透传 Range 头(参考实现全量 GET,无法断点/分段);
// - 5xx/429/网络错误指数退避重试,可配次数(参考实现无重试);
// - 下载经 io.Pipe 流式转发,全程不落盘;
// - 目录存在性内存缓存,按需逐级 MKCOL,避免每次保存都发 PROPFIND;
// - 非流式操作带可配超时;流式传输由调用方 ctx 管控(可取消)。
type WebDAVStorage struct {
base *url.URL // 服务基址(含可能的路径前缀),以 / 结尾
root string // 远端根目录(webdav_root_path
username string
password string
client *http.Client
transport *http.Transport
auth *authState
maxRetries int // 5xx/网络错误最大重试次数
baseBackoff time.Duration // 退避基数
opTimeout time.Duration // 非流式操作超时
dirMu sync.RWMutex
knownDirs map[string]struct{} // 已确认存在的远端目录(含根前缀)
spacesPool sync.Pool // 256KB 复用缓冲
}
// NewWebDAVStorage 构造 WebDAV 引擎。
func NewWebDAVStorage(opts WebDAVOptions) (*WebDAVStorage, error) {
opts.applyDefaults()
raw := strings.TrimSpace(opts.BaseURL)
if raw == "" {
return nil, fmt.Errorf("storage/webdav: 缺少 webdav_url 配置")
}
if !strings.Contains(raw, "://") {
raw = "http://" + raw
}
u, err := url.Parse(raw)
if err != nil {
return nil, fmt.Errorf("storage/webdav: webdav_url 非法: %w", err)
}
if u.Scheme != "http" && u.Scheme != "https" {
return nil, fmt.Errorf("storage/webdav: webdav_url 仅支持 http/https,收到 %q", u.Scheme)
}
if !strings.HasSuffix(u.Path, "/") {
u.Path += "/"
}
root := strings.Trim(opts.RootPath, "/")
if root == "" {
root = "filebox_storage"
}
root = strings.ReplaceAll(root, "\\", "/")
transport := newPooledTransport(opts.MaxIdleConnsPerHost)
return &WebDAVStorage{
base: u,
root: root,
username: opts.Username,
password: opts.Password,
client: &http.Client{Transport: transport},
transport: transport,
auth: newAuthState(opts.Username, opts.Password),
maxRetries: opts.MaxRetries,
baseBackoff: time.Duration(opts.BaseBackoff) * time.Millisecond,
opTimeout: time.Duration(opts.Timeout) * time.Second,
knownDirs: map[string]struct{}{},
spacesPool: sync.Pool{New: func() any {
b := make([]byte, localChunkSize)
return &b
}},
}, nil
}
func init() {
RegisterEngine("webdav", func(ctx context.Context) (Storage, error) {
return NewWebDAVStorage(engineOptions.WebDAV)
})
}
// newPooledTransport 连接池化 TransportKeep-Alive 连接复用是 WebDAV 优化的核心。
func newPooledTransport(maxIdlePerHost int) *http.Transport {
return &http.Transport{
Proxy: http.ProxyFromEnvironment,
DialContext: (&net.Dialer{
Timeout: 10 * time.Second,
KeepAlive: 30 * time.Second,
}).DialContext,
ForceAttemptHTTP2: true,
MaxIdleConns: 100,
MaxIdleConnsPerHost: maxIdlePerHost,
IdleConnTimeout: 90 * time.Second,
TLSHandshakeTimeout: 10 * time.Second,
ExpectContinueTimeout: time.Second,
ResponseHeaderTimeout: 60 * time.Second,
}
}
// requestOpts 单次 WebDAV 请求参数。
type requestOpts struct {
// body 请求体工厂:每次尝试调用一次(重试时重新获取,可重放)。
body func() (io.Reader, int64, error)
// retryBody 请求体是否可重放(seekable);false 时 PUT 类请求失败不重试。
retryBody bool
// headers 附加请求头。
headers map[string]string
// streaming 流式传输(GET/PUT 大 body):不套 opTimeout,由调用方 ctx 管控。
streaming bool
}
// do 执行一次 WebDAV 请求:认证自动协商 + 指数退避重试。
// 返回的响应由调用方负责关闭(drainClose / readErrorBody)。
//
// 重要:非流式操作的可配超时通过 ctx 实现,cancel 不随 do() 返回而调用,
// 而是挂在 davResponse 上、待响应体读完后再触发——否则取消会提前杀掉
// Keep-Alive 连接,破坏连接复用。
func (w *WebDAVStorage) do(ctx context.Context, method, rawURL string, opts requestOpts) (*davResponse, error) {
// 非流式操作套可配超时(流式由调用方 ctx 管控)。
var cancel context.CancelFunc
if !opts.streaming {
if _, hasDeadline := ctx.Deadline(); !hasDeadline {
ctx, cancel = context.WithTimeout(ctx, w.opTimeout)
}
}
fail := func(err error) (*davResponse, error) {
if cancel != nil {
cancel()
}
return nil, err
}
// 幂等方法或可重放 body 才允许整体重试。
idempotent := method == http.MethodGet || method == http.MethodHead ||
method == "PROPFIND" || method == "MKCOL" || method == http.MethodDelete ||
method == http.MethodOptions
retryable := idempotent || opts.retryBody
const maxAuthRetries = 2
budget := w.maxRetries + maxAuthRetries // 认证挑战重试不消耗退避预算
authRetries := 0
for attempt := 0; attempt < budget; attempt++ {
var body io.Reader
var length int64 = -1
if opts.body != nil {
var err error
body, length, err = opts.body()
if err != nil {
return fail(fmt.Errorf("%w: 构造请求体失败: %v", ErrUnavailable, err))
}
}
req, err := http.NewRequestWithContext(ctx, method, rawURL, body)
if err != nil {
return fail(fmt.Errorf("%w: 构造请求失败: %v", ErrInvalidPath, err))
}
if length >= 0 {
req.ContentLength = length
}
for k, v := range opts.headers {
req.Header.Set(k, v)
}
w.auth.apply(req)
resp, err := w.client.Do(req)
if err != nil {
if ctx.Err() != nil { // 调用方取消/超时优先
return fail(ctx.Err())
}
if retryable && attempt+1 < budget {
if sleepErr := w.backoff(ctx, attempt, 0); sleepErr != nil {
return fail(sleepErr)
}
continue
}
return fail(fmt.Errorf("%w: %s %s: %v", ErrUnavailable, method, rawURL, err))
}
// 401 认证挑战:切换 Basic/Digest 后立即重试(不退避、不额外计数)。
if resp.StatusCode == http.StatusUnauthorized && authRetries < maxAuthRetries {
challenge := resp.Header.Get("WWW-Authenticate")
drainClose(&davResponse{Response: resp})
if challenge != "" && w.auth.challenge(challenge) {
authRetries++
continue
}
return fail(fmt.Errorf("%w: WebDAV 认证失败(401%s", ErrUnavailable, rawURL))
}
// 5xx/429/408:幂等或可重放 body 时指数退避重试。
if retryable && isRetryStatus(resp.StatusCode) && attempt+1 < budget {
retryAfter := retryAfterSeconds(resp.Header.Get("Retry-After"))
drainClose(&davResponse{Response: resp})
if sleepErr := w.backoff(ctx, attempt, retryAfter); sleepErr != nil {
return fail(sleepErr)
}
continue
}
return &davResponse{Response: resp, cancel: cancel}, nil
}
return fail(fmt.Errorf("%w: WebDAV 重试耗尽(%s %s", ErrUnavailable, method, rawURL))
}
// davResponse WebDAV 响应 + 关联的超时取消函数。
// 非流式操作读完响应体后必须经 drainClose/readErrorBody 释放(触发 cancel)。
type davResponse struct {
*http.Response
cancel context.CancelFunc
}
// isRetryStatus 判断状态码是否值得重试。
func isRetryStatus(code int) bool {
switch code {
case http.StatusRequestTimeout, http.StatusTooManyRequests,
http.StatusInternalServerError, http.StatusBadGateway,
http.StatusServiceUnavailable, http.StatusGatewayTimeout:
return true
}
return false
}
// retryAfterSeconds 解析 Retry-After(秒);非法或负值返回 0。
func retryAfterSeconds(v string) time.Duration {
if v == "" {
return 0
}
n, err := strconv.Atoi(strings.TrimSpace(v))
if err != nil || n <= 0 {
return 0
}
if n > 5 {
n = 5 // 上限 5s,避免异常服务端拖死请求
}
return time.Duration(n) * time.Second
}
// backoff 指数退避:base * 2^attempt,封顶 2s,带 ±20% 抖动;retryAfter 优先。
func (w *WebDAVStorage) backoff(ctx context.Context, attempt int, retryAfter time.Duration) error {
d := retryAfter
if d <= 0 {
d = w.baseBackoff << attempt
if d > 2*time.Second {
d = 2 * time.Second
}
// ±20% 抖动
jitter := time.Duration(int64(d) / 5)
if jitter > 0 {
d -= time.Duration(rand.Int63n(int64(jitter)))
}
}
if d <= 0 {
d = time.Millisecond
}
timer := time.NewTimer(d)
defer timer.Stop()
select {
case <-ctx.Done():
return ctx.Err()
case <-timer.C:
return nil
}
}
// drainClose 读取少量残余并关闭响应体,保证连接可复用;随后触发超时清理。
func drainClose(resp *davResponse) {
if resp == nil || resp.Body == nil {
return
}
_, _ = io.Copy(io.Discard, io.LimitReader(resp.Body, 8<<10))
_ = resp.Body.Close()
if resp.cancel != nil {
resp.cancel()
}
}
// joinRemote 校验 savePath 并拼接远端完整路径(含根目录前缀)。
func (w *WebDAVStorage) joinRemote(savePath string) (string, error) {
cleaned, ok := SanitizePath(savePath)
if !ok {
return "", fmt.Errorf("%w: %q", ErrInvalidPath, savePath)
}
return path.Join(w.root, cleaned), nil
}
// urlFor 将远端路径转为完整 URLURL.String 自动按段转义)。
func (w *WebDAVStorage) urlFor(remotePath string) string {
u := *w.base
p := strings.TrimSuffix(u.Path, "/")
remotePath = strings.Trim(remotePath, "/")
if remotePath != "" && remotePath != "." {
p += "/" + remotePath
}
u.Path = p
return u.String()
}
// propfindBody PROPFIND 请求体:只取需要的属性。
const propfindBody = `<?xml version="1.0" encoding="utf-8"?>` +
`<D:propfind xmlns:D="DAV:"><D:prop>` +
`<D:resourcetype/><D:getcontentlength/><D:getcontenttype/>` +
`</D:prop></D:propfind>`
// davMultistatus 207 Multi-Status XML 解析结构(标签名与命名空间无关匹配)。
type davMultistatus struct {
Responses []struct {
Href string `xml:"href"`
Propstat []struct {
Status string `xml:"status"`
Prop struct {
ContentLength int64 `xml:"getcontentlength"`
ContentType string `xml:"getcontenttype"`
ResourceType struct {
Collection *struct{} `xml:"collection"`
} `xml:"resourcetype"`
} `xml:"prop"`
} `xml:"propstat"`
} `xml:"response"`
}
// firstProp 取第一个 HTTP 2xx 状态的属性块。
func (m *davMultistatus) firstProp() (length int64, ctype string, isDir bool, ok bool) {
for _, r := range m.Responses {
for _, ps := range r.Propstat {
if !strings.Contains(ps.Status, " 200 ") {
continue
}
return ps.Prop.ContentLength, ps.Prop.ContentType, ps.Prop.ResourceType.Collection != nil, true
}
}
return 0, "", false, false
}
// propfind 执行 PROPFIND 并解析 207 响应;404 时返回 (nil, nil)。
func (w *WebDAVStorage) propfind(ctx context.Context, rawURL string, depth string) (*davMultistatus, error) {
resp, err := w.do(ctx, "PROPFIND", rawURL, requestOpts{
body: func() (io.Reader, int64, error) {
return strings.NewReader(propfindBody), int64(len(propfindBody)), nil
},
headers: map[string]string{"Depth": depth, "Content-Type": "application/xml"},
})
if err != nil {
return nil, err
}
defer drainClose(resp)
switch resp.StatusCode {
case http.StatusMultiStatus, http.StatusOK:
body, err := io.ReadAll(io.LimitReader(resp.Body, 4<<20))
if err != nil {
return nil, fmt.Errorf("%w: PROPFIND 读取失败: %v", ErrUnavailable, err)
}
var ms davMultistatus
if err := xml.Unmarshal(body, &ms); err != nil {
return nil, fmt.Errorf("%w: PROPFIND XML 解析失败: %v", ErrUnavailable, err)
}
return &ms, nil
case http.StatusNotFound:
return nil, nil
default:
return nil, fmt.Errorf("%w: PROPFIND %s → %d", ErrUnavailable, rawURL, resp.StatusCode)
}
}
// remoteExists PROPFIND 探测远端路径存在性。
func (w *WebDAVStorage) remoteExists(ctx context.Context, remotePath string) (bool, error) {
ms, err := w.propfind(ctx, w.urlFor(remotePath), "0")
if err != nil {
return false, err
}
return ms != nil, nil
}
// markDir 记录已确认存在的目录(避免重复 PROPFIND/MKCOL 往返)。
func (w *WebDAVStorage) markDir(remotePath string) {
w.dirMu.Lock()
defer w.dirMu.Unlock()
w.knownDirs[remotePath] = struct{}{}
}
// unmarkDir 目录被删除时移除缓存。
func (w *WebDAVStorage) unmarkDir(remotePath string) {
w.dirMu.Lock()
defer w.dirMu.Unlock()
delete(w.knownDirs, remotePath)
}
// isMarkedDir 查询目录缓存。
func (w *WebDAVStorage) isMarkedDir(remotePath string) bool {
w.dirMu.RLock()
defer w.dirMu.RUnlock()
_, ok := w.knownDirs[remotePath]
return ok
}
// ensureDirs 按需逐级创建远端目录(含根前缀;MKCOL 级联,成功后写缓存)。
func (w *WebDAVStorage) ensureDirs(ctx context.Context, remotePath string) error {
segments := splitRemoteSegments(remotePath)
cur := ""
for _, seg := range segments {
cur = path.Join(cur, seg)
if w.isMarkedDir(cur) {
continue
}
exists, err := w.remoteExists(ctx, cur)
if err == nil && exists {
w.markDir(cur)
continue
}
if err != nil && !errors.Is(err, ErrNotFound) {
return err
}
resp, err := w.do(ctx, "MKCOL", w.urlFor(cur), requestOpts{})
if err != nil {
return err
}
status := resp.StatusCode
drainClose(resp)
// 201 创建成功;405 已存在;其余视为失败(409 通常因父目录缺失,理论上不会出现)。
if status == http.StatusCreated || status == http.StatusOK ||
status == http.StatusNoContent || status == http.StatusMethodNotAllowed {
w.markDir(cur)
continue
}
return fmt.Errorf("%w: MKCOL %s → %d", ErrUnavailable, cur, status)
}
return nil
}
// splitRemoteSegments 拆分远端路径段。
func splitRemoteSegments(p string) []string {
p = strings.Trim(strings.ReplaceAll(p, "\\", "/"), "/")
if p == "" {
return nil
}
return strings.Split(p, "/")
}
// deleteEmptyParents 删除空父目录(含根前缀,但不删根目录本身);尽力而为。
func (w *WebDAVStorage) deleteEmptyParents(ctx context.Context, remotePath string) {
dir := path.Dir(remotePath)
for dir != "" && dir != "." && dir != w.root && strings.HasPrefix(dir+"/", w.root+"/") {
ms, err := w.propfind(ctx, w.urlFor(dir), "1")
if err != nil || ms == nil {
return
}
if len(ms.Responses) > 1 { // 非空(自身 + 子项)
return
}
resp, err := w.do(ctx, http.MethodDelete, w.urlFor(dir), requestOpts{})
if err != nil {
return
}
ok := resp.StatusCode == http.StatusOK || resp.StatusCode == http.StatusNoContent
drainClose(resp)
if !ok {
return
}
w.unmarkDir(dir)
dir = path.Dir(dir)
}
}
// putFile PUT 上传:body 工厂每次尝试返回可重放的读取器。
func (w *WebDAVStorage) putFile(ctx context.Context, rawURL string, body func() (io.Reader, int64, error), retryBody bool) (*davResponse, error) {
return w.do(ctx, http.MethodPut, rawURL, requestOpts{
body: body,
retryBody: retryBody,
headers: map[string]string{"Content-Type": "application/octet-stream"},
streaming: true,
})
}
// checkPutStatus 校验 PUT 响应状态。
func checkPutStatus(resp *davResponse, op string) error {
switch resp.StatusCode {
case http.StatusOK, http.StatusCreated, http.StatusNoContent:
drainClose(resp)
return nil
default:
msg := readErrorBody(resp)
return fmt.Errorf("%w: %s → %d %s", ErrUnavailable, op, resp.StatusCode, msg)
}
}
// readErrorBody 读取错误响应前 200 字节并释放连接。
func readErrorBody(resp *davResponse) string {
if resp == nil || resp.Body == nil {
return ""
}
b, _ := io.ReadAll(io.LimitReader(resp.Body, 200))
_ = resp.Body.Close()
if resp.cancel != nil {
resp.cancel()
}
return strings.TrimSpace(string(b))
}
// SaveFile 流式保存(PUT):按需建目录,seekable 源可安全重试。
func (w *WebDAVStorage) SaveFile(ctx context.Context, r io.Reader, savePath string) (int64, error) {
remote, err := w.joinRemote(savePath)
if err != nil {
return 0, err
}
if err := w.ensureDirs(ctx, path.Dir(remote)); err != nil {
return 0, err
}
// 可重放判定:seekable 源失败后可从头重传(PUT 覆盖语义保证最终一致)。
seeker, seekable := r.(io.Seeker)
var knownLen int64 = -1
if seekable {
if cur, err := seeker.Seek(0, io.SeekCurrent); err == nil {
if end, err := seeker.Seek(0, io.SeekEnd); err == nil {
knownLen = end - cur
_, _ = seeker.Seek(cur, io.SeekStart)
}
}
}
src := &countingReader{r: r}
body := func() (io.Reader, int64, error) {
if seekable {
if _, err := seeker.Seek(0, io.SeekStart); err != nil {
return nil, -1, err
}
src.reset()
}
return src, knownLen, nil
}
resp, err := w.putFile(ctx, w.urlFor(remote), body, seekable)
if err != nil {
return src.count(), err
}
if err := checkPutStatus(resp, fmt.Sprintf("PUT %s", remote)); err != nil {
return src.count(), err
}
return src.count(), nil
}
// DeleteFile DELETE 文件 + 尽力清理空父目录。
func (w *WebDAVStorage) DeleteFile(ctx context.Context, savePath string) error {
remote, err := w.joinRemote(savePath)
if err != nil {
return err
}
resp, err := w.do(ctx, http.MethodDelete, w.urlFor(remote), requestOpts{})
if err != nil {
return err
}
switch resp.StatusCode {
case http.StatusOK, http.StatusNoContent, http.StatusNotFound:
drainClose(resp)
default:
msg := readErrorBody(resp)
return fmt.Errorf("%w: DELETE %s → %d %s", ErrUnavailable, remote, resp.StatusCode, msg)
}
w.deleteEmptyParents(ctx, remote)
return nil
}
// Open 打开下载流:Range 透传,io.Pipe 流式转发不落盘,ctx 可取消。
func (w *WebDAVStorage) Open(ctx context.Context, savePath string, rng *Range) (*Download, error) {
remote, err := w.joinRemote(savePath)
if err != nil {
return nil, err
}
opts := requestOpts{streaming: true}
if rng != nil {
opts.headers = map[string]string{"Range": rangeHeaderValue(rng)}
}
resp, err := w.do(ctx, http.MethodGet, w.urlFor(remote), opts)
if err != nil {
return nil, err
}
switch resp.StatusCode {
case http.StatusOK, http.StatusPartialContent:
// 正常,继续
case http.StatusNotFound:
drainClose(resp)
return nil, ErrNotFound
case http.StatusRequestedRangeNotSatisfiable:
drainClose(resp)
return nil, ErrRangeNotSatisfiable
default:
msg := readErrorBody(resp)
return nil, fmt.Errorf("%w: GET %s → %d %s", ErrUnavailable, remote, resp.StatusCode, msg)
}
total := resp.ContentLength
start, end := int64(0), total-1
if resp.StatusCode == http.StatusPartialContent {
if cr := resp.Header.Get("Content-Range"); cr != "" {
if s0, e0, t0, ok := parseContentRange(cr); ok {
start, end = s0, e0
if t0 >= 0 {
total = t0
}
}
}
}
if total < 0 { // 服务端未给出长度(chunked):按未知大小处理
start, end, total = 0, -1, -1
}
if rng == nil { // 对齐契约:完整文件 Start=0、End=Total-1
start, end = 0, total-1
}
if end < 0 { // 空文件或未知大小:End 未知语义
end = -1
}
// io.Pipe 流式桥接:HTTP 响应体 → 管道 → 调用方,全程不落盘;
// 调用方提前 Close 或 ctx 取消都会终止拷贝并释放连接。
body := resp.Body
pr, pw := io.Pipe()
go func() {
bufp, _ := w.spacesPool.Get().(*[]byte)
_, copyErr := io.CopyBuffer(pw, body, *bufp)
w.spacesPool.Put(bufp)
_ = body.Close()
pw.CloseWithError(copyErr) // copyErr 为 nil 时写入 EOF
}()
context.AfterFunc(ctx, func() {
_ = pw.CloseWithError(ctx.Err())
})
contentType := resp.Header.Get("Content-Type")
return &Download{
ReadCloser: pr,
Start: start,
End: end,
Total: total,
Meta: FileMeta{
Size: total,
ContentType: contentType,
AcceptRanges: true,
},
}, nil
}
// Stat PROPFIND Depth 0 获取元信息。
func (w *WebDAVStorage) Stat(ctx context.Context, savePath string) (*FileMeta, error) {
remote, err := w.joinRemote(savePath)
if err != nil {
return nil, err
}
ms, err := w.propfind(ctx, w.urlFor(remote), "0")
if err != nil {
if errors.Is(err, ErrNotFound) {
return nil, ErrNotFound
}
return nil, err
}
if ms == nil {
return nil, ErrNotFound
}
length, ctype, _, ok := ms.firstProp()
if !ok {
return nil, ErrNotFound
}
return &FileMeta{Size: length, ContentType: ctype, AcceptRanges: true}, nil
}
// HeadMeta 读取文件元信息与前 n 字节(WebDAV 实现:PROPFIND + Range GET)。
func (w *WebDAVStorage) HeadMeta(ctx context.Context, savePath string, headBytes int64) (*FileMeta, []byte, error) {
meta, err := w.Stat(ctx, savePath)
if err != nil {
return nil, nil, err
}
if headBytes <= 0 {
return meta, nil, nil
}
dl, err := w.Open(ctx, savePath, &Range{Start: 0, End: headBytes - 1})
if err != nil {
if errors.Is(err, ErrRangeNotSatisfiable) { // 空文件等边界:返回空头
return meta, nil, nil
}
return nil, nil, err
}
defer func() { _ = dl.Close() }()
head := make([]byte, headBytes)
n, _ := io.ReadFull(dl.ReadCloser, head)
return meta, head[:n], nil
}
// SaveChunk 保存分片:落临时文件获得精确长度与可重放 body,PUT 到分片路径。
func (w *WebDAVStorage) SaveChunk(ctx context.Context, uploadID string, chunkIndex int, r io.Reader, savePath string) (int64, error) {
if _, err := w.joinRemote(savePath); err != nil {
return 0, err
}
chunkRel, ok := SanitizePath(ChunkPartPath(savePath, uploadID, chunkIndex))
if !ok {
return 0, fmt.Errorf("%w: 分片路径非法 %q", ErrInvalidPath, savePath)
}
remote := path.Join(w.root, chunkRel)
if err := w.ensureDirs(ctx, path.Dir(remote)); err != nil {
return 0, err
}
// 分片体积有限(默认 ≤8MB):落临时文件换取精确 Content-Length 与可重试性。
tmp, err := os.CreateTemp("", "fcb-webdav-chunk-*")
if err != nil {
return 0, fmt.Errorf("storage/webdav: 创建临时文件失败: %w", err)
}
tmpName := tmp.Name()
defer func() { _ = tmp.Close(); _ = os.Remove(tmpName) }()
size, err := io.CopyBuffer(tmp, r, make([]byte, localChunkSize))
if err != nil {
return 0, fmt.Errorf("storage/webdav: 缓存分片失败: %w", err)
}
if _, err := tmp.Seek(0, io.SeekStart); err != nil {
return 0, fmt.Errorf("storage/webdav: 回卷分片失败: %w", err)
}
body := func() (io.Reader, int64, error) {
if _, err := tmp.Seek(0, io.SeekStart); err != nil {
return nil, -1, err
}
return tmp, size, nil
}
resp, err := w.putFile(ctx, w.urlFor(remote), body, true)
if err != nil {
return size, err
}
if err := checkPutStatus(resp, fmt.Sprintf("PUT 分片 %s", remote)); err != nil {
return size, err
}
return size, nil
}
// MergeChunks 合并 WebDAV 分片:
// 逐分片 GET 流式拼入本地临时文件(边拷贝边校验哈希)→ PUT 上传目标 → 清理远端分片与本地临时文件。
// 说明:WebDAV 无服务端聚合能力,合并必须经服务端中转;临时文件仅用于拼接与重试,最终 PUT 可重放。
func (w *WebDAVStorage) MergeChunks(ctx context.Context, uploadID string, total int, verifyHash func(index int) (string, error), savePath string) (int64, string, error) {
if total <= 0 {
return 0, "", fmt.Errorf("storage/webdav: 非法分片总数 %d", total)
}
remote, err := w.joinRemote(savePath)
if err != nil {
return 0, "", err
}
if err := w.ensureDirs(ctx, path.Dir(remote)); err != nil {
return 0, "", err
}
tmp, err := os.CreateTemp("", "fcb-webdav-merge-*")
if err != nil {
return 0, "", fmt.Errorf("storage/webdav: 创建合并临时文件失败: %w", err)
}
tmpName := tmp.Name()
defer func() { _ = tmp.Close(); _ = os.Remove(tmpName) }()
totalHash := sha256.New()
var size int64
for i := 0; i < total; i++ {
if err := ctx.Err(); err != nil {
return 0, "", err
}
chunkRel, ok := SanitizePath(ChunkPartPath(savePath, uploadID, i))
if !ok {
return 0, "", fmt.Errorf("%w: 分片 %d 路径非法", ErrInvalidPath, i)
}
resp, err := w.do(ctx, http.MethodGet, w.urlFor(path.Join(w.root, chunkRel)), requestOpts{streaming: true})
if err != nil {
return 0, "", fmt.Errorf("storage/webdav: 读取分片 %d 失败: %w", i, err)
}
if resp.StatusCode == http.StatusNotFound {
drainClose(resp)
return 0, "", fmt.Errorf("storage/webdav: 分片 %d 不存在", i)
}
if resp.StatusCode != http.StatusOK && resp.StatusCode != http.StatusPartialContent {
msg := readErrorBody(resp)
return 0, "", fmt.Errorf("storage/webdav: 读取分片 %d → %d %s", i, resp.StatusCode, msg)
}
chunkHash := sha256.New()
n, err := io.CopyBuffer(io.MultiWriter(tmp, totalHash, chunkHash), resp.Body, make([]byte, localChunkSize))
_ = resp.Body.Close()
if err != nil {
return 0, "", fmt.Errorf("storage/webdav: 拼接分片 %d 失败: %w", i, err)
}
if verifyHash != nil {
expected, err := verifyHash(i)
if err != nil {
return 0, "", err
}
if expected != "" && expected != hex.EncodeToString(chunkHash.Sum(nil)) {
return 0, "", fmt.Errorf("%w: 分片 %d 期望 %s", ErrHashMismatch, i, expected)
}
}
size += n
}
if _, err := tmp.Seek(0, io.SeekStart); err != nil {
return 0, "", fmt.Errorf("storage/webdav: 回卷合并文件失败: %w", err)
}
body := func() (io.Reader, int64, error) {
if _, err := tmp.Seek(0, io.SeekStart); err != nil {
return nil, -1, err
}
return tmp, size, nil
}
resp, err := w.putFile(ctx, w.urlFor(remote), body, true)
if err != nil {
return size, "", err
}
if err := checkPutStatus(resp, fmt.Sprintf("PUT 合并 %s", remote)); err != nil {
return size, "", err
}
// 合并成功后清理远端分片目录与本地临时文件(defer 兜底删除本地文件)。
_ = w.CleanChunks(ctx, uploadID, savePath)
return size, hex.EncodeToString(totalHash.Sum(nil)), nil
}
// CleanChunks 递归删除远端分片目录(RFC 4918 DELETE 对 collection 递归)。
func (w *WebDAVStorage) CleanChunks(ctx context.Context, uploadID string, savePath string) error {
dirRel, ok := SanitizePath(chunkDirOf(savePath, uploadID))
if !ok {
return fmt.Errorf("%w: 分片目录非法 %q", ErrInvalidPath, savePath)
}
remote := path.Join(w.root, dirRel)
resp, err := w.do(ctx, http.MethodDelete, w.urlFor(remote), requestOpts{})
if err != nil {
return err
}
switch resp.StatusCode {
case http.StatusOK, http.StatusNoContent, http.StatusNotFound:
drainClose(resp)
w.unmarkDir(remote)
default:
msg := readErrorBody(resp)
return fmt.Errorf("%w: 清理分片目录 %s → %d %s", ErrUnavailable, remote, resp.StatusCode, msg)
}
return nil
}
// FileExists PROPFIND 探测存在性;非法路径按不存在处理。
func (w *WebDAVStorage) FileExists(ctx context.Context, savePath string) (bool, error) {
remote, err := w.joinRemote(savePath)
if err != nil {
return false, nil
}
return w.remoteExists(ctx, remote)
}
// PresignGetURL WebDAV 无预签名直链能力。
func (w *WebDAVStorage) PresignGetURL(ctx context.Context, savePath string, expires int64) (string, error) {
return "", ErrNotSupported
}
// PresignPutURL WebDAV 无预签名直传能力。
func (w *WebDAVStorage) PresignPutURL(ctx context.Context, savePath string, expires int64) (string, error) {
return "", ErrNotSupported
}
// HealthCheck 健康检查:PROPFIND 根目录;不存在时 MKCOL 创建(启动自愈)。
// 同时完成凭据与连通性验证(do 内 401 协商)。
func (w *WebDAVStorage) HealthCheck(ctx context.Context) error {
exists, err := w.remoteExists(ctx, w.root)
if err == nil && exists {
w.markDir(w.root)
return nil
}
if err != nil && !errors.Is(err, ErrNotFound) {
return fmt.Errorf("%w: WebDAV 健康检查失败: %v", ErrUnavailable, err)
}
resp, err := w.do(ctx, "MKCOL", w.urlFor(w.root), requestOpts{})
if err != nil {
return err
}
switch resp.StatusCode {
case http.StatusCreated, http.StatusOK, http.StatusNoContent, http.StatusMethodNotAllowed:
drainClose(resp)
w.markDir(w.root)
return nil
default:
msg := readErrorBody(resp)
return fmt.Errorf("%w: WebDAV 根目录创建失败 → %d %s", ErrUnavailable, resp.StatusCode, msg)
}
}
// 接口编译期断言。
var _ Storage = (*WebDAVStorage)(nil)
+244
View File
@@ -0,0 +1,244 @@
package storage
import (
"crypto/md5"
"crypto/rand"
"crypto/sha256"
"encoding/hex"
"fmt"
"net/http"
"strings"
"sync"
)
// authMode 认证模式(WebDAV 服务端挑战后自动协商)。
type authMode int
const (
authModeUnknown authMode = iota // 未定:先发 Basic 探测
authModeBasic
authModeDigest
)
// authState WebDAV Basic/Digest 认证状态。
//
// 策略:
// - 首个请求预置 Basic;若服务端 401 且挑战为 Digest,则解析挑战参数切换为 Digest;
// - Digest 按 RFC 2617/7616 实现 qop=authMD5 / SHA-256,含 -sess 变体);
// qop 缺失时回退 RFC 2069 旧式响应;
// - nonce 变更时重置 nc 计数;nc/cnonce 在互斥锁内生成保证并发唯一。
type authState struct {
mu sync.Mutex
username string
password string
mode authMode
realm string
nonce string
qop string // 选定的 qop"auth" 或空 = RFC2069
opaque string
algorithm string // MD5 | MD5-sess | SHA-256 | SHA-256-sess
nc uint32
knownBasicOK bool // 已确认 Basic 可用
}
// newAuthState 构造认证状态(默认以 Basic 起步)。
func newAuthState(username, password string) *authState {
return &authState{username: username, password: password}
}
// apply 为请求设置 Authorization 头(每次请求调用,Digest 时消耗一个 nc)。
func (a *authState) apply(req *http.Request) {
a.mu.Lock()
defer a.mu.Unlock()
switch {
case a.mode == authModeDigest && a.nonce != "":
req.Header.Set("Authorization", a.digestHeader(req))
default:
req.SetBasicAuth(a.username, a.password)
}
}
// digestHeader 依据缓存的挑战参数计算 Digest Authorization 头(调用方需持锁)。
func (a *authState) digestHeader(req *http.Request) string {
uri := req.URL.RequestURI()
method := strings.ToUpper(req.Method)
ncStr := fmt.Sprintf("%08x", a.nc+1)
a.nc++
cnonce := randomHex(8)
var ha1 string
switch strings.ToLower(a.algorithm) {
case "md5-sess":
ha1 = hashHex("md5", hashHex("md5", a.username+":"+a.realm+":"+a.password)+":"+a.nonce+":"+cnonce)
case "sha-256-sess":
ha1 = hashHex("sha256", hashHex("sha256", a.username+":"+a.realm+":"+a.password)+":"+a.nonce+":"+cnonce)
case "sha-256":
ha1 = hashHex("sha256", a.username+":"+a.realm+":"+a.password)
default: // md5
ha1 = hashHex("md5", a.username+":"+a.realm+":"+a.password)
}
ha2 := hashHex(algoName(a.algorithm), method+":"+uri)
var response string
var fields []string
esc := escapeDigestValue(a.username)
if a.qop == "" { // RFC 2069
response = hashHex(algoName(a.algorithm), ha1+":"+a.nonce+":"+ha2)
fields = append(fields,
`Digest username="`+esc+`"`,
`realm="`+escapeDigestValue(a.realm)+`"`,
`nonce="`+escapeDigestValue(a.nonce)+`"`,
`uri="`+escapeDigestValue(uri)+`"`,
`response="`+response+`"`)
} else {
response = hashHex(algoName(a.algorithm), ha1+":"+a.nonce+":"+ncStr+":"+cnonce+":"+a.qop+":"+ha2)
fields = append(fields,
`Digest username="`+esc+`"`,
`realm="`+escapeDigestValue(a.realm)+`"`,
`nonce="`+escapeDigestValue(a.nonce)+`"`,
`uri="`+escapeDigestValue(uri)+`"`,
`cnonce="`+cnonce+`"`,
`nc=`+ncStr,
`qop=`+a.qop,
`response="`+response+`"`,
`algorithm=`+a.algorithm)
}
if a.opaque != "" {
fields = append(fields, `opaque="`+escapeDigestValue(a.opaque)+`"`)
}
return strings.Join(fields, ", ")
}
// challenge 处理 401 的 WWW-Authenticate 挑战;返回是否已切换认证方式可重试。
// 返回 false 表示凭据错误或算法不受支持,调用方应直接报错。
func (a *authState) challenge(header string) bool {
a.mu.Lock()
defer a.mu.Unlock()
h := strings.TrimSpace(header)
lower := strings.ToLower(h)
switch {
case strings.HasPrefix(lower, "digest"):
params := parseChallengeParams(strings.TrimPrefix(h[len("Digest"):], " "))
algo := strings.ToUpper(strings.TrimSpace(params["algorithm"]))
if algo == "" {
algo = "MD5"
}
switch algo {
case "MD5", "MD5-SESS", "SHA-256", "SHA-256-SESS":
default:
return false // 不支持的摘要算法
}
if params["nonce"] == "" || params["realm"] == "" {
return false
}
qop := ""
if raw := strings.TrimSpace(params["qop"]); raw != "" {
for _, candidate := range strings.Split(raw, ",") {
if strings.EqualFold(strings.TrimSpace(candidate), "auth") {
qop = "auth"
break
}
}
if qop == "" {
return false // 仅支持 auth-int 等需要 body 哈希的模式
}
}
if a.nonce != params["nonce"] {
a.nc = 0
}
a.realm, a.nonce, a.qop = params["realm"], params["nonce"], qop
a.opaque, a.algorithm = params["opaque"], strings.ToLower(algo)
a.mode = authModeDigest
a.knownBasicOK = false
return true
case strings.HasPrefix(lower, "basic"):
if a.knownBasicOK || a.mode == authModeBasic {
return false // 已用 Basic 仍 401:凭据错误
}
a.mode = authModeBasic
a.knownBasicOK = true
return true
}
return false
}
// parseChallengeParams 解析 "realm=\"x\", nonce=\"y\"" 形式的挑战参数(引号内逗号不切分)。
func parseChallengeParams(s string) map[string]string {
out := map[string]string{}
for _, item := range splitAuthParams(s) {
kv := strings.SplitN(item, "=", 2)
if len(kv) != 2 {
continue
}
k := strings.ToLower(strings.TrimSpace(kv[0]))
v := strings.TrimSpace(kv[1])
if len(v) >= 2 && strings.HasPrefix(v, `"`) && strings.HasSuffix(v, `"`) {
v = v[1 : len(v)-1]
}
out[k] = v
}
return out
}
// splitAuthParams 逗号切分但忽略引号内的逗号。
func splitAuthParams(s string) []string {
var parts []string
var b strings.Builder
inQuote := false
for i := 0; i < len(s); i++ {
c := s[i]
switch {
case c == '"':
inQuote = !inQuote
b.WriteByte(c)
case c == ',' && !inQuote:
if t := strings.TrimSpace(b.String()); t != "" {
parts = append(parts, t)
}
b.Reset()
default:
b.WriteByte(c)
}
}
if t := strings.TrimSpace(b.String()); t != "" {
parts = append(parts, t)
}
return parts
}
// algoName 映射哈希函数名。
func algoName(algorithm string) string {
switch strings.ToLower(algorithm) {
case "sha-256", "sha-256-sess":
return "sha256"
default:
return "md5"
}
}
// hashHex 通用哈希摘要(algo: md5|sha256)。
func hashHex(algo, s string) string {
if algo == "sha256" {
sum := sha256.Sum256([]byte(s))
return hex.EncodeToString(sum[:])
}
sum := md5.Sum([]byte(s))
return hex.EncodeToString(sum[:])
}
// escapeDigestValue 转义引号。
func escapeDigestValue(s string) string {
return strings.ReplaceAll(s, `"`, `\"`)
}
// randomHex 生成 n 字节随机 hex。
func randomHex(n int) string {
b := make([]byte, n)
if _, err := rand.Read(b); err != nil {
// crypto/rand 失败极其罕见;退化为全零仍保持协议可用。
for i := range b {
b[i] = 0
}
}
return hex.EncodeToString(b)
}
+824
View File
@@ -0,0 +1,824 @@
package storage
import (
"bytes"
"context"
"encoding/base64"
"fmt"
"io"
"net"
"net/http"
"net/http/httptest"
"strconv"
"strings"
"sync"
"sync/atomic"
"testing"
"time"
)
// ---- 最小 WebDAV 假服务:PUT/GET/HEAD/PROPFIND/MKCOL/DELETE + Basic/Digest 认证 ----
type davLog struct {
Method string
Path string
Status int
}
type fakeDav struct {
mu sync.Mutex
dirs map[string]bool
files map[string][]byte
// 认证配置:mode = none|basic|digestdigest 配合 algo = MD5|SHA-256。
mode string
username string
password string
realm string
nonce string
opaque string
algo string
failNext map[string]int // method → 剩余 503 次数
logs []davLog
}
func newFakeDav(mode string) *fakeDav {
return &fakeDav{
dirs: map[string]bool{},
files: map[string][]byte{},
mode: mode,
username: "fcb",
password: "fcb-pass",
realm: "test-realm",
nonce: "dcd98b7102dd2f0e8b11d0f600bfb0c0",
opaque: "5ccc069c403ebaf9f0171e9517f40e41",
algo: "MD5",
failNext: map[string]int{},
}
}
// auth 校验请求凭据;失败时写出 401 与对应挑战。
func (f *fakeDav) auth(w http.ResponseWriter, r *http.Request) bool {
if f.mode == "none" {
return true
}
h := r.Header.Get("Authorization")
ok := false
switch f.mode {
case "basic":
ok = h == "Basic "+basicAuth(f.username, f.password)
case "digest":
ok = f.checkDigest(r)
}
if ok {
return true
}
switch f.mode {
case "basic":
w.Header().Set("WWW-Authenticate", `Basic realm="`+f.realm+`"`)
case "digest":
w.Header().Set("WWW-Authenticate", fmt.Sprintf(
`Digest realm="%s", qop="auth", nonce="%s", opaque="%s", algorithm=%s, stale=false`,
f.realm, f.nonce, f.opaque, f.algo))
}
w.WriteHeader(http.StatusUnauthorized)
return false
}
// checkDigest 服务端重算 RFC 2617 摘要响应。
func (f *fakeDav) checkDigest(r *http.Request) bool {
h := r.Header.Get("Authorization")
if !strings.HasPrefix(h, "Digest ") {
return false
}
p := parseChallengeParams(strings.TrimSpace(h[len("Digest "):]))
ha1 := hashHex(algoName(f.algo), f.username+":"+f.realm+":"+f.password)
ha2 := hashHex(algoName(f.algo), strings.ToUpper(r.Method)+":"+r.URL.RequestURI())
got := hashHex(algoName(f.algo), ha1+":"+f.nonce+":"+p["nc"]+":"+p["cnonce"]+":"+p["qop"]+":"+ha2)
return p["username"] == f.username && p["response"] == got
}
func basicAuth(user, pass string) string {
return base64.StdEncoding.EncodeToString([]byte(user + ":" + pass))
}
func (f *fakeDav) record(method, path string, status int) {
f.logs = append(f.logs, davLog{Method: method, Path: path, Status: status})
}
// maybeFail 命中失败注入时返回 true(已写出 503)。
func (f *fakeDav) maybeFail(w http.ResponseWriter, method string) bool {
if f.failNext[method] > 0 {
f.failNext[method]--
w.WriteHeader(http.StatusServiceUnavailable)
return true
}
return false
}
func (f *fakeDav) ServeHTTP(w http.ResponseWriter, r *http.Request) {
f.mu.Lock()
defer f.mu.Unlock()
if !f.auth(w, r) {
f.record(r.Method, r.URL.Path, http.StatusUnauthorized)
return
}
p := strings.Trim(r.URL.Path, "/")
switch r.Method {
case http.MethodPut:
if f.maybeFail(w, "PUT") {
f.record(r.Method, p, 503)
return
}
if parent := parentOf(p); parent != "" && !f.dirs[parent] {
w.WriteHeader(http.StatusConflict) // 强制客户端先建目录
f.record(r.Method, p, 409)
return
}
body, _ := io.ReadAll(r.Body)
f.files[p] = body
f.record(r.Method, p, 201)
w.WriteHeader(http.StatusCreated)
case http.MethodGet:
if f.maybeFail(w, "GET") {
f.record(r.Method, p, 503)
return
}
data, ok := f.files[p]
if !ok {
w.WriteHeader(http.StatusNotFound)
f.record(r.Method, p, 404)
return
}
w.Header().Set("Content-Type", "application/octet-stream")
w.Header().Set("Accept-Ranges", "bytes")
if rng := r.Header.Get("Range"); rng != "" {
start, end := int64(0), int64(len(data))-1
spec := strings.TrimPrefix(rng, "bytes=")
if strings.HasSuffix(spec, "-") { // bytes=N- → 到文件尾
if s, err := strconv.ParseInt(strings.TrimSuffix(spec, "-"), 10, 64); err == nil {
start = s
}
} else if _, err := fmt.Sscanf(spec, "%d-%d", &start, &end); err != nil {
w.WriteHeader(http.StatusBadRequest)
f.record(r.Method, p, 400)
return
}
if start < 0 || start >= int64(len(data)) {
w.WriteHeader(http.StatusRequestedRangeNotSatisfiable)
f.record(r.Method, p, 416)
return
}
if end >= int64(len(data)) {
end = int64(len(data)) - 1
}
w.Header().Set("Content-Range", fmt.Sprintf("bytes %d-%d/%d", start, end, len(data)))
w.WriteHeader(http.StatusPartialContent)
_, _ = w.Write(data[start : end+1])
f.record(r.Method, p, 206)
return
}
w.WriteHeader(http.StatusOK)
_, _ = w.Write(data)
f.record(r.Method, p, 200)
case http.MethodHead:
data, ok := f.files[p]
if !ok {
w.WriteHeader(http.StatusNotFound)
f.record(r.Method, p, 404)
return
}
w.Header().Set("Content-Length", strconv.Itoa(len(data)))
w.WriteHeader(http.StatusOK)
f.record(r.Method, p, 200)
case "PROPFIND":
depth := r.Header.Get("Depth")
self, isDirSelf := f.stat(p)
if !isDirSelf {
w.WriteHeader(http.StatusNotFound)
f.record(r.Method, p, 404)
return
}
var b strings.Builder
b.WriteString(`<?xml version="1.0" encoding="utf-8"?>` +
`<D:multistatus xmlns:D="DAV:">`)
f.writeResponse(&b, p, self)
if depth == "1" && self.isDir {
for _, name := range f.children(p) {
child := name
cs, cd := f.stat(child)
f.writeResponse(&b, child, davStat{isDir: cd, size: cs.size})
}
}
b.WriteString(`</D:multistatus>`)
w.Header().Set("Content-Type", "application/xml; charset=utf-8")
w.WriteHeader(http.StatusMultiStatus)
_, _ = w.Write([]byte(b.String()))
f.record(r.Method, p, 207)
case "MKCOL":
if f.dirs[p] || f.files[p] != nil {
w.WriteHeader(http.StatusMethodNotAllowed) // 已存在
f.record(r.Method, p, 405)
return
}
if parent := parentOf(p); parent != "" && !f.dirs[parent] {
w.WriteHeader(http.StatusConflict)
f.record(r.Method, p, 409)
return
}
f.dirs[p] = true
w.WriteHeader(http.StatusCreated)
f.record(r.Method, p, 201)
case http.MethodDelete:
if _, ok := f.files[p]; ok {
delete(f.files, p)
w.WriteHeader(http.StatusNoContent)
f.record(r.Method, p, 204)
return
}
if f.dirs[p] {
// 递归删除目录
prefix := p + "/"
for name := range f.files {
if strings.HasPrefix(name, prefix) {
delete(f.files, name)
}
}
for name := range f.dirs {
if name == p || strings.HasPrefix(name+"/", prefix) {
delete(f.dirs, name)
}
}
w.WriteHeader(http.StatusNoContent)
f.record(r.Method, p, 204)
return
}
w.WriteHeader(http.StatusNotFound)
f.record(r.Method, p, 404)
default:
w.WriteHeader(http.StatusMethodNotAllowed)
f.record(r.Method, p, 405)
}
}
type davStat struct {
isDir bool
size int
}
func (f *fakeDav) stat(p string) (davStat, bool) {
if data, ok := f.files[p]; ok {
return davStat{size: len(data)}, true
}
if f.dirs[p] {
return davStat{isDir: true}, true
}
return davStat{}, false
}
func (f *fakeDav) children(p string) []string {
var out []string
prefix := p + "/"
for name := range f.files {
if strings.HasPrefix(name, prefix) && !strings.Contains(strings.TrimPrefix(name, prefix), "/") {
out = append(out, name)
}
}
for name := range f.dirs {
if strings.HasPrefix(name, prefix) && !strings.Contains(strings.TrimPrefix(name, prefix), "/") {
out = append(out, name)
}
}
return out
}
func (f *fakeDav) writeResponse(b *strings.Builder, href string, st davStat) {
b.WriteString(`<D:response><D:href>/` + href + `</D:href><D:propstat><D:prop><D:resourcetype>`)
if st.isDir {
b.WriteString(`<D:collection/>`)
}
b.WriteString(`</D:resourcetype><D:getcontentlength>` + strconv.Itoa(st.size) +
`</D:getcontentlength><D:getcontenttype>application/octet-stream</D:getcontenttype>` +
`</D:prop><D:status>HTTP/1.1 200 OK</D:status></D:propstat></D:response>`)
}
func parentOf(p string) string {
if i := strings.LastIndex(p, "/"); i > 0 {
return p[:i]
}
return ""
}
// newTestDav 构造 WebDAV 引擎 + 假服务。
func newTestDav(t *testing.T, mode string, tweak func(o *WebDAVOptions)) (*WebDAVStorage, *fakeDav, *int32) {
t.Helper()
f := newFakeDav(mode)
var conns int32
srv := httptest.NewUnstartedServer(f)
srv.Config.ConnState = func(c net.Conn, cs http.ConnState) {
if cs == http.StateNew {
atomic.AddInt32(&conns, 1)
}
}
srv.Start()
t.Cleanup(srv.Close)
opts := WebDAVOptions{
BaseURL: srv.URL,
Username: f.username,
Password: f.password,
RootPath: "fcb_root",
MaxRetries: 3,
}
if tweak != nil {
tweak(&opts)
}
st, err := NewWebDAVStorage(opts)
if err != nil {
t.Fatalf("NewWebDAVStorage: %v", err)
}
return st, f, &conns
}
// TestWebDAVBasicCRUD Basic 认证下的完整 CRUD 与 Range。
func TestWebDAVBasicCRUD(t *testing.T) {
st, f, _ := newTestDav(t, "basic", nil)
ctx := context.Background()
// 健康检查:根目录 404 → MKCOL 自建
if err := st.HealthCheck(ctx); err != nil {
t.Fatalf("HealthCheck: %v", err)
}
if !f.dirs["fcb_root"] {
t.Fatalf("根目录应被自动创建")
}
data := []byte("WebDAV 引擎数据 0123456789 ABCDEF")
n, err := st.SaveFile(ctx, bytes.NewReader(data), "2025/08/w.bin")
if err != nil {
t.Fatalf("SaveFile: %v", err)
}
if n != int64(len(data)) {
t.Fatalf("n = %d", n)
}
if string(f.files["fcb_root/2025/08/w.bin"]) != string(data) {
t.Fatalf("PUT 内容不匹配")
}
// 按需建目录:两级目录都应已创建
if !f.dirs["fcb_root/2025"] || !f.dirs["fcb_root/2025/08"] {
t.Fatalf("目录未按需创建: %v %v", f.dirs["fcb_root/2025"], f.dirs["fcb_root/2025/08"])
}
meta, err := st.Stat(ctx, "2025/08/w.bin")
if err != nil {
t.Fatalf("Stat: %v", err)
}
if meta.Size != int64(len(data)) || !meta.AcceptRanges {
t.Fatalf("Stat = %+v", meta)
}
if ok, _ := st.FileExists(ctx, "2025/08/w.bin"); !ok {
t.Fatalf("FileExists 应为 true")
}
// 完整下载(对齐 go-api 约定)
dl, err := st.Open(ctx, "2025/08/w.bin", nil)
if err != nil {
t.Fatalf("Open: %v", err)
}
got, err := io.ReadAll(dl)
_ = dl.Close()
if err != nil {
t.Fatalf("read: %v", err)
}
if !bytes.Equal(got, data) {
t.Fatalf("full mismatch")
}
if dl.Start != 0 || dl.End != int64(len(data))-1 || dl.Total != int64(len(data)) {
t.Fatalf("full offsets = %d,%d,%d", dl.Start, dl.End, dl.Total)
}
// Range 下载
dl, err = st.Open(ctx, "2025/08/w.bin", &Range{Start: 2, End: 7})
if err != nil {
t.Fatalf("Open range: %v", err)
}
got, err = io.ReadAll(dl)
_ = dl.Close()
if err != nil {
t.Fatalf("read range: %v", err)
}
if !bytes.Equal(got, data[2:8]) {
t.Fatalf("range mismatch")
}
if dl.Start != 2 || dl.End != 7 || dl.Total != int64(len(data)) {
t.Fatalf("range offsets = %d,%d,%d", dl.Start, dl.End, dl.Total)
}
// 416(起点越界)/ 404
if _, err := st.Open(ctx, "2025/08/w.bin", &Range{Start: int64(len(data)) + 9, End: -1}); err == nil ||
!strings.Contains(err.Error(), ErrRangeNotSatisfiable.Error()) {
t.Fatalf("want ErrRangeNotSatisfiable, got %v", err)
}
if _, err := st.Open(ctx, "no/such.bin", nil); err == nil ||
!strings.Contains(err.Error(), ErrNotFound.Error()) {
t.Fatalf("want ErrNotFound, got %v", err)
}
// 删除 + 空父目录清理
if err := st.DeleteFile(ctx, "2025/08/w.bin"); err != nil {
t.Fatalf("DeleteFile: %v", err)
}
if ok, _ := st.FileExists(ctx, "2025/08/w.bin"); ok {
t.Fatalf("删除后仍存在")
}
if _, err := st.Stat(ctx, "2025/08/w.bin"); err == nil ||
!strings.Contains(err.Error(), ErrNotFound.Error()) {
t.Fatalf("Stat 应 ErrNotFound, got %v", err)
}
}
// TestWebDAVDigestAuth DigestMD5)认证协商。
func TestWebDAVDigestAuth(t *testing.T) {
st, f, _ := newTestDav(t, "digest", nil)
ctx := context.Background()
if err := st.HealthCheck(ctx); err != nil {
t.Fatalf("HealthCheck(digest): %v", err)
}
// HealthCheck 流程应观察到 401 挑战(客户端先 Basic 探测 → 401 → Digest 重试)
saw401 := false
for _, l := range f.logs {
if l.Status == 401 {
saw401 = true
}
}
if !saw401 {
t.Fatalf("未观察到 401 挑战: %+v", f.logs)
}
// 认证后的 PROPFINDStat 已有目录)应得到 207
if _, err := st.Stat(ctx, ""); err == nil {
// Stat("") 非法路径属预期;这里换用 FileExists 对已有根目录探测
_ = err
}
data := []byte("digest 内容")
if _, err := st.SaveFile(ctx, bytes.NewReader(data), "d.bin"); err != nil {
t.Fatalf("SaveFile(digest): %v", err)
}
dl, err := st.Open(ctx, "d.bin", nil)
if err != nil {
t.Fatalf("Open(digest): %v", err)
}
got, _ := io.ReadAll(dl)
_ = dl.Close()
if !bytes.Equal(got, data) {
t.Fatalf("digest 下载内容不匹配")
}
// 全链路完成:确认存在成功的 2xx/207 请求
saw2xx := false
for _, l := range f.logs {
if l.Status == 207 || l.Status == 201 || l.Status == 200 {
saw2xx = true
}
}
if !saw2xx {
t.Fatalf("认证后应有成功请求: %+v", f.logs)
}
}
// TestWebDAVDigestSHA256 DigestSHA-256)算法。
func TestWebDAVDigestSHA256(t *testing.T) {
st, _, _ := newTestDav(t, "digest", nil)
st.auth.mu.Lock()
st.auth.algorithm = "sha-256"
st.auth.mu.Unlock()
// 服务端也切换到 SHA-256 重算摘要
st2, f, _ := newTestDav(t, "digest", nil)
f.algo = "SHA-256"
// 先让客户端完成一次 MD5 协商拿到挑战参数,再切 SHA-256 会 401 失败——
// 因此这里直接对 SHA-256 服务端做完整链路(client 首次探测 Basic→401→Digest)。
_ = st
ctx := context.Background()
if err := st2.HealthCheck(ctx); err != nil {
t.Fatalf("HealthCheck(SHA-256): %v", err)
}
}
// TestWebDAVDigestWrongPassword 凭据错误 → 明确报错而非重试风暴。
func TestWebDAVDigestWrongPassword(t *testing.T) {
f := newFakeDav("digest")
srv := httptest.NewServer(f)
defer srv.Close()
st, err := NewWebDAVStorage(WebDAVOptions{
BaseURL: srv.URL, Username: f.username, Password: "WRONG",
RootPath: "r", MaxRetries: 1, BaseBackoff: 1,
})
if err != nil {
t.Fatal(err)
}
if err := st.HealthCheck(context.Background()); err == nil ||
!strings.Contains(err.Error(), "401") {
t.Fatalf("错误凭据应报 401 相关错误, got %v", err)
}
}
// TestWebDAVRetryGet 5xx 指数退避重试(GET 幂等)。
func TestWebDAVRetryGet(t *testing.T) {
st, f, _ := newTestDav(t, "basic", func(o *WebDAVOptions) { o.BaseBackoff = 5 })
ctx := context.Background()
data := []byte("retry target")
if err := st.HealthCheck(ctx); err != nil {
t.Fatal(err)
}
if _, err := st.SaveFile(ctx, bytes.NewReader(data), "r.bin"); err != nil {
t.Fatal(err)
}
f.mu.Lock()
f.failNext["GET"] = 2
f.mu.Unlock()
dl, err := st.Open(ctx, "r.bin", nil)
if err != nil {
t.Fatalf("503×2 后应重试成功: %v", err)
}
got, _ := io.ReadAll(dl)
_ = dl.Close()
if !bytes.Equal(got, data) {
t.Fatalf("content mismatch")
}
// 验证确实发了 3 次 GET
gets := 0
for _, l := range f.logs {
if l.Method == "GET" && strings.HasSuffix(l.Path, "r.bin") {
gets++
}
}
if gets != 3 {
t.Fatalf("GET 次数 = %d, want 3", gets)
}
}
// TestWebDAVRetryPut 可重放 bodyseekable)PUT 失败重试;不可重放不重试。
func TestWebDAVRetryPut(t *testing.T) {
st, f, _ := newTestDav(t, "basic", func(o *WebDAVOptions) { o.BaseBackoff = 5 })
ctx := context.Background()
if err := st.HealthCheck(ctx); err != nil {
t.Fatal(err)
}
// seekable:重试成功
f.mu.Lock()
f.failNext["PUT"] = 1
f.mu.Unlock()
data := []byte("put with retry")
if _, err := st.SaveFile(ctx, bytes.NewReader(data), "pr.bin"); err != nil {
t.Fatalf("PUT 重试应成功: %v", err)
}
puts := 0
for _, l := range f.logs {
if l.Method == "PUT" && strings.HasSuffix(l.Path, "pr.bin") {
puts++
}
}
if puts != 2 {
t.Fatalf("PUT 次数 = %d, want 2", puts)
}
// 非 seekableio.Pipe):不重试,直接失败
f.mu.Lock()
f.failNext["PUT"] = 1
f.mu.Unlock()
pr, pw := io.Pipe()
go func() {
_, _ = pw.Write([]byte("non-seekable"))
_ = pw.Close()
}()
if _, err := st.SaveFile(ctx, pr, "ns.bin"); err == nil {
t.Fatalf("非重放 PUT 注入 503 应失败")
}
}
// TestWebDAVConnectionReuse 连接复用:多次请求不应各建一条 TCP 连接。
func TestWebDAVConnectionReuse(t *testing.T) {
st, _, conns := newTestDav(t, "basic", nil)
ctx := context.Background()
if err := st.HealthCheck(ctx); err != nil {
t.Fatal(err)
}
for i := 0; i < 12; i++ {
if _, err := st.SaveFile(ctx, bytes.NewReader([]byte("x")), fmt.Sprintf("reuse/%d.bin", i)); err != nil {
t.Fatal(err)
}
if _, err := st.Stat(ctx, fmt.Sprintf("reuse/%d.bin", i)); err != nil {
t.Fatal(err)
}
}
// 25 次请求(12 PUT + 12 PROPFIND + 1 HealthCheck 的 PROPFIND/MKCOL)只允许极少量新连接
if got := atomic.LoadInt32(conns); got > 4 {
t.Fatalf("新建 TCP 连接数 = %d,连接复用失效(应 ≤4)", got)
}
}
// TestWebDAVPipeStreaming io.Pipe 流式转发:完整读取 + 提前关闭。
func TestWebDAVPipeStreaming(t *testing.T) {
st, _, _ := newTestDav(t, "basic", nil)
ctx := context.Background()
if err := st.HealthCheck(ctx); err != nil {
t.Fatal(err)
}
big := bytes.Repeat([]byte("0123456789abcdef"), 4096) // 64KB
if _, err := st.SaveFile(ctx, bytes.NewReader(big), "big.bin"); err != nil {
t.Fatal(err)
}
dl, err := st.Open(ctx, "big.bin", nil)
if err != nil {
t.Fatal(err)
}
got, err := io.ReadAll(dl)
if err != nil {
t.Fatalf("read pipe: %v", err)
}
_ = dl.Close()
if !bytes.Equal(got, big) {
t.Fatalf("pipe content mismatch")
}
// 提前关闭:后续读取返回错误且不挂死
dl2, err := st.Open(ctx, "big.bin", nil)
if err != nil {
t.Fatal(err)
}
buf := make([]byte, 10)
if _, err := io.ReadFull(dl2, buf); err != nil {
t.Fatalf("read head: %v", err)
}
if err := dl2.Close(); err != nil {
t.Fatalf("early close: %v", err)
}
// ctx 取消同样会终止流
cctx, cancel := context.WithCancel(context.Background())
dl3, err := st.Open(cctx, "big.bin", nil)
if err != nil {
t.Fatal(err)
}
cancel()
time.Sleep(20 * time.Millisecond)
_, err = dl3.Read(buf)
if err == nil {
_ = dl3.Close()
t.Fatalf("ctx 取消后读取应报错")
}
_ = dl3.Close()
}
// TestWebDAVChunkMerge 分片保存/合并/清理。
func TestWebDAVChunkMerge(t *testing.T) {
st, f, _ := newTestDav(t, "basic", nil)
ctx := context.Background()
if err := st.HealthCheck(ctx); err != nil {
t.Fatal(err)
}
savePath := "2025/09/merged.bin"
uploadID := "uid-webdav"
chunks := [][]byte{[]byte("AAA"), []byte("BB"), []byte("CCCC")}
hashes := make([]string, len(chunks))
for i, c := range chunks {
n, err := st.SaveChunk(ctx, uploadID, i, bytes.NewReader(c), savePath)
if err != nil {
t.Fatalf("SaveChunk %d: %v", i, err)
}
if n != int64(len(c)) {
t.Fatalf("chunk %d size = %d", i, n)
}
hashes[i] = sha256Hex(c)
}
size, fileHash, err := st.MergeChunks(ctx, uploadID, len(chunks), func(i int) (string, error) {
return hashes[i], nil
}, savePath)
if err != nil {
t.Fatalf("MergeChunks: %v", err)
}
if size != 9 || fileHash != sha256Hex(bytes.Join(chunks, nil)) {
t.Fatalf("merge result = %d %s", size, fileHash)
}
if string(f.files["fcb_root/"+savePath]) != "AAABBCCCC" {
t.Fatalf("合并内容错误: %q", f.files["fcb_root/"+savePath])
}
// 分片目录已清理
for k := range f.files {
if strings.Contains(k, "chunks/"+uploadID) {
t.Fatalf("分片残留: %s", k)
}
}
if f.dirs["fcb_root/2025/09/chunks/"+uploadID] {
t.Fatalf("分片目录残留")
}
}
// TestWebDAVCleanChunks 清理与哈希失败路径。
func TestWebDAVCleanChunks(t *testing.T) {
st, f, _ := newTestDav(t, "basic", nil)
ctx := context.Background()
if err := st.HealthCheck(ctx); err != nil {
t.Fatal(err)
}
for i := 0; i < 2; i++ {
if _, err := st.SaveChunk(ctx, "uidc", i, strings.NewReader("z"), "c.bin"); err != nil {
t.Fatal(err)
}
}
if err := st.CleanChunks(ctx, "uidc", "c.bin"); err != nil {
t.Fatalf("CleanChunks: %v", err)
}
if len(f.files) != 0 {
t.Fatalf("分片未清理: %v", f.files)
}
// 幂等
if err := st.CleanChunks(ctx, "uidc", "c.bin"); err != nil {
t.Fatalf("CleanChunks idempotent: %v", err)
}
// 哈希不匹配
if _, err := st.SaveChunk(ctx, "uidm", 0, strings.NewReader("real"), "m.bin"); err != nil {
t.Fatal(err)
}
if _, _, err := st.MergeChunks(ctx, "uidm", 1, func(i int) (string, error) {
return sha256Hex([]byte("wrong")), nil
}, "m.bin"); err == nil || !strings.Contains(err.Error(), ErrHashMismatch.Error()) {
t.Fatalf("want ErrHashMismatch, got %v", err)
}
}
// TestWebDAVTimeout 非流式操作超时:PROPFIND 响应慢于 Timeout1s)→ context deadline exceeded。
func TestWebDAVTimeout(t *testing.T) {
f := newFakeDav("basic")
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method == "PROPFIND" {
time.Sleep(1500 * time.Millisecond) // > Timeout 1s
}
f.ServeHTTP(w, r)
}))
defer srv.Close()
st, err := NewWebDAVStorage(WebDAVOptions{
BaseURL: srv.URL, Username: f.username, Password: f.password,
RootPath: "r", MaxRetries: 0, Timeout: 1, BaseBackoff: 5,
})
if err != nil {
t.Fatal(err)
}
start := time.Now()
err = st.HealthCheck(context.Background())
if err == nil {
t.Fatalf("超时应报错")
}
if !strings.Contains(err.Error(), "context deadline exceeded") {
t.Fatalf("应为超时错误, got %v", err)
}
// 单次尝试 1s 超时 + 一次重试 ≈ 2s;若超时未生效会拖满 2×1.5s
if elapsed := time.Since(start); elapsed > 3500*time.Millisecond {
t.Fatalf("超时未生效(耗时 %v", elapsed)
}
}
// TestWebDAVPresignNotSupported 预签名 → ErrNotSupported。
func TestWebDAVPresignNotSupported(t *testing.T) {
st, _, _ := newTestDav(t, "basic", nil)
ctx := context.Background()
if _, err := st.PresignGetURL(ctx, "x.bin", 60); err == nil ||
!strings.Contains(err.Error(), ErrNotSupported.Error()) {
t.Fatalf("want ErrNotSupported, got %v", err)
}
if _, err := st.PresignPutURL(ctx, "x.bin", 60); err == nil ||
!strings.Contains(err.Error(), ErrNotSupported.Error()) {
t.Fatalf("want ErrNotSupported, got %v", err)
}
}
// TestWebDAVFactoryRegistry 工厂构造 + Digest 全链路。
func TestWebDAVFactoryRegistry(t *testing.T) {
f := newFakeDav("digest")
srv := httptest.NewServer(f)
defer srv.Close()
prev := engineOptions.WebDAV
engineOptions.WebDAV = WebDAVOptions{
BaseURL: srv.URL, Username: f.username, Password: f.password,
RootPath: "factory_root", MaxRetries: 3, BaseBackoff: 5,
}
defer func() { engineOptions.WebDAV = prev }()
st, err := NewEngine(context.Background(), "webdav")
if err != nil {
t.Fatalf("NewEngine(webdav): %v", err)
}
ctx := context.Background()
if err := st.HealthCheck(ctx); err != nil {
t.Fatalf("HealthCheck: %v", err)
}
if _, err := st.SaveFile(ctx, strings.NewReader("factory"), "f.txt"); err != nil {
t.Fatalf("SaveFile: %v", err)
}
dl, err := st.Open(ctx, "f.txt", nil)
if err != nil {
t.Fatalf("Open: %v", err)
}
got, _ := io.ReadAll(dl)
_ = dl.Close()
if string(got) != "factory" {
t.Fatalf("content = %q", got)
}
}
+1
View File
@@ -0,0 +1 @@
import{d as p,u as v,i as h,I as g,c as w,G as s,b as e,t as n,f as t,A as d,w as k,F as y,O as r,q as m,B as R,o as V}from"./index-DYsKpclu.js";import{_ as x}from"./SiteNav.vue_vue_type_script_setup_true_lang-CwaEJKIZ.js";import{u as A}from"./auth-B7MDTgxJ.js";import"./admin-DEOkRyTC.js";const B={class:"admin-shell"},C={class:"admin-aside"},b={class:"aside-title"},L=["aria-label"],N={class:"admin-main"},q=p({__name:"AdminLayout",setup(S){const{t:a}=v(),u=R(),o=A(),l=h();g(async()=>{o.isAuthed&&!o.checked&&await o.verify()});async function _(){await o.logout(),l.success(a("admin.nav.loggedOut")),u.replace({name:"admin-login"})}return(F,c)=>{const i=r("RouterLink"),f=r("RouterView");return V(),w(y,null,[s(x),e("div",B,[e("aside",C,[e("div",b,n(t(a)("admin.nav.title")),1),e("nav",{class:"aside-menu","aria-label":t(a)("admin.nav.menu")},[s(i,{to:{name:"admin-files"}},{default:d(()=>[m("📁 "+n(t(a)("admin.nav.files")),1)]),_:1}),s(i,{to:{name:"admin-audit"}},{default:d(()=>[m("🛡 "+n(t(a)("admin.nav.audit")),1)]),_:1}),s(i,{to:{name:"admin-settings"}},{default:d(()=>[m("⚙️ "+n(t(a)("admin.nav.settings")),1)]),_:1}),c[0]||(c[0]=e("div",{class:"aside-sep"},null,-1)),e("a",{href:"#",onClick:k(_,["prevent"])},"🚪 "+n(t(a)("admin.nav.logout")),1)],8,L)]),e("div",N,[s(f)])])],64)}}});export{q as default};
File diff suppressed because one or more lines are too long
+1
View File
@@ -0,0 +1 @@
.filter-bar[data-v-dcdef1a5]{display:flex;gap:10px;align-items:flex-end;flex-wrap:wrap;background:var(--glass-bg-soft);border:1px solid var(--glass-border);border-radius:var(--radius);padding:14px 16px;margin-bottom:16px}.filter-bar label[data-v-dcdef1a5]{display:flex;flex-direction:column;gap:4px;font-size:12.5px;color:var(--c-text-2);font-weight:600}.filter-bar .select[data-v-dcdef1a5],.filter-bar .input[data-v-dcdef1a5]{min-width:130px}.ip-cell[data-v-dcdef1a5]{font-family:var(--mono);font-size:12.5px}.ua-cell[data-v-dcdef1a5]{max-width:260px;overflow:hidden;text-overflow:ellipsis;color:var(--c-text-3)}
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
+1
View File
@@ -0,0 +1 @@
.field-sub[data-v-d049fcf9]{display:block;font-size:13px;font-weight:600;color:var(--c-text-2);margin-bottom:6px}.result-head[data-v-8bebf3d9]{display:flex;align-items:center;gap:10px;margin-bottom:14px}.result-name[data-v-8bebf3d9]{font-weight:600;word-break:break-all}.link-row[data-v-8bebf3d9]{display:flex;gap:8px}.link-row .input[data-v-8bebf3d9]{flex:1}.result-meta[data-v-8bebf3d9]{display:flex;align-items:center;gap:6px;color:var(--c-text-2);font-size:13px;flex-wrap:wrap}.code-copy[data-v-8bebf3d9]{display:block;width:100%;cursor:pointer;appearance:none;-webkit-appearance:none;transition:filter .15s}.code-copy[data-v-8bebf3d9]:hover{filter:brightness(.96)}.code-copy[data-v-8bebf3d9]:active{filter:brightness(.92)}.dropzone.disabled[data-v-1b57fe5c]{opacity:.55;cursor:not-allowed}
+1
View File
@@ -0,0 +1 @@
.login-head[data-v-26e7f2a4]{text-align:center;margin-bottom:18px}.login-logo[data-v-26e7f2a4]{width:48px;height:48px;object-fit:contain;margin-bottom:8px}
+1
View File
@@ -0,0 +1 @@
import{d as w,u as b,a as y,z as x,A as k,b as e,f as t,t as a,w as V,C as B,D as C,c as m,g as _,q as N,j as r,N as S,B as q,o as d,H as L,_ as P}from"./index-DYsKpclu.js";import{P as A}from"./PageShell-D87JkPkG.js";import{u as D}from"./auth-B7MDTgxJ.js";import"./SiteNav.vue_vue_type_script_setup_true_lang-CwaEJKIZ.js";import"./admin-DEOkRyTC.js";const E={class:"card",style:{"max-width":"380px",margin:"8vh auto 0"}},I={class:"login-head"},M=["src"],R={class:"card-title"},T={class:"card-sub"},U={class:"field"},j={for:"admin-password"},z=["placeholder"],H={key:0,class:"hint",style:{color:"var(--c-danger)","margin-bottom":"12px"}},F=["disabled"],G={key:0,class:"spin","aria-hidden":"true"},J={class:"hint",style:{"margin-top":"16px"}},K=w({__name:"LoginView",setup(O){const{t:s}=b(),c=S(),g=q(),h=D(),u=y(),i=r(""),l=r(!1),n=r("");async function f(){if(!i.value){n.value=s("admin.login.required");return}l.value=!0,n.value="";try{await h.login(i.value);const o=typeof c.query.redirect=="string"?c.query.redirect:"/admin/files";g.replace(o)}catch(o){n.value=o instanceof L?o.code===401?s("admin.login.wrongPassword"):o.msg:s("admin.login.failed"),i.value=""}finally{l.value=!1}}return(o,p)=>(d(),x(A,null,{default:k(()=>[e("section",E,[e("div",I,[e("img",{src:t(u).displayLogoUrl,alt:"Logo",class:"login-logo"},null,8,M),e("h1",R,a(t(s)("admin.login.title")),1),e("p",T,a(t(s)("admin.login.subtitle",{name:t(u).displayName})),1)]),e("form",{onSubmit:V(f,["prevent"])},[e("div",U,[e("label",j,a(t(s)("admin.login.password")),1),B(e("input",{id:"admin-password","onUpdate:modelValue":p[0]||(p[0]=v=>i.value=v),class:"input",type:"password",placeholder:t(s)("admin.login.passwordPlaceholder"),autocomplete:"current-password",autofocus:""},null,8,z),[[C,i.value]])]),n.value?(d(),m("p",H,a(n.value),1)):_("",!0),e("button",{class:"btn btn-block",type:"submit",disabled:l.value},[l.value?(d(),m("span",G)):_("",!0),N(" "+a(t(s)("admin.login.submit")),1)],8,F)],32),e("p",J,a(t(s)("admin.login.hint")),1)])]),_:1}))}}),$=P(K,[["__scopeId","data-v-26e7f2a4"]]);export{$ as default};
+1
View File
@@ -0,0 +1 @@
import{d as c,u as i,z as d,A as a,b as t,t as e,f as s,G as l,q as p,O as u,o as _}from"./index-DYsKpclu.js";import{P as m}from"./PageShell-D87JkPkG.js";import"./SiteNav.vue_vue_type_script_setup_true_lang-CwaEJKIZ.js";const f={class:"card empty",style:{"max-width":"480px",margin:"10vh auto 0"}},h={style:{"font-weight":"600",color:"var(--c-text)"}},x={class:"hint"},F=c({__name:"NotFoundView",setup(y){const{t:o}=i();return(g,n)=>{const r=u("RouterLink");return _(),d(m,null,{default:a(()=>[t("div",f,[n[0]||(n[0]=t("div",{class:"empty-icon"},"🧭",-1)),t("p",h,e(s(o)("notFound.title")),1),t("p",x,e(s(o)("notFound.desc")),1),l(r,{class:"btn",to:"/",style:{"margin-top":"10px"}},{default:a(()=>[p(e(s(o)("notFound.back")),1)]),_:1})])]),_:1})}}});export{F as default};
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
+1
View File
@@ -0,0 +1 @@
import{d as h,u as g,a as k,c as n,G as i,b as t,Q as y,e as v,f as e,t as a,g as c,q as r,A as d,F as x,h as B,O as C,o as l,_ as b}from"./index-DYsKpclu.js";import{_ as w}from"./SiteNav.vue_vue_type_script_setup_true_lang-CwaEJKIZ.js";const N={class:"site-footer"},S={class:"footer-left"},F={key:0,class:"footer-text"},P={key:1,class:"footer-beian"},T={key:2},V=["aria-label"],$=h({__name:"PageShell",setup(D){const{t:s}=g(),o=k(),m=new Date().getFullYear(),u=B(()=>!!(o.footerText.trim()||o.footerBeian.trim()));return(p,_)=>{const f=C("RouterLink");return l(),n(x,null,[i(w),t("main",{class:v(["page",{"page-wide":p.$route.meta.wide}])},[y(p.$slots,"default",{},void 0,!0)],2),t("footer",N,[t("div",S,[e(o).footerText.trim()?(l(),n("span",F,a(e(o).footerText),1)):c("",!0),e(o).footerBeian.trim()?(l(),n("span",P,a(e(o).footerBeian),1)):c("",!0),u.value?c("",!0):(l(),n("span",T,a(e(s)("footer.copyright",{year:e(m),name:e(o).displayName})),1))]),_[0]||(_[0]=t("span",{class:"footer-powered"},[r(" Powered by "),t("a",{href:"https://skymirror.top",target:"_blank",rel:"noopener noreferrer"},"SKYMirror"),r(" 26.9 ")],-1)),t("nav",{"aria-label":e(s)("footer.linkNav")},[i(f,{to:"/docs"},{default:d(()=>[r(a(e(s)("footer.docs")),1)]),_:1}),i(f,{to:"/openapi"},{default:d(()=>[r(a(e(s)("footer.openapi")),1)]),_:1}),i(f,{to:"/admin/files"},{default:d(()=>[r(a(e(s)("footer.admin")),1)]),_:1})],8,V)])],64)}}}),R=b($,[["__scopeId","data-v-80685d6f"]]);export{R as P};
+1
View File
@@ -0,0 +1 @@
.footer-left[data-v-80685d6f]{display:flex;flex-direction:column;gap:2px;min-width:0}.footer-powered[data-v-80685d6f]{font-size:13px;color:var(--c-text-2);white-space:nowrap}.footer-powered a[data-v-80685d6f]{color:var(--c-text-2);text-decoration:none;border-bottom:1px dotted currentColor}.footer-powered a[data-v-80685d6f]:hover{color:var(--c-text)}@media(min-width:720px){.site-footer[data-v-80685d6f]{display:flex;align-items:center;justify-content:space-between;gap:16px}.footer-left[data-v-80685d6f]{flex-direction:row;align-items:center;gap:16px}.footer-left .footer-beian[data-v-80685d6f]:before{content:"·";margin-right:16px;opacity:.7}}
@@ -0,0 +1 @@
import{d as b,u as p,c as r,b as o,t as i,f as c,h,o as f}from"./index-DYsKpclu.js";const v={class:"pager"},x={class:"pager-info"},k=["disabled"],M=["disabled"],y=b({__name:"Pager",props:{page:{},size:{},total:{}},emits:["change"],setup(t,{emit:m}){const a=t,d=m,{t:n}=p(),s=h(()=>Math.max(1,Math.ceil(a.total/a.size)));function g(l){const e=Math.min(Math.max(1,l),s.value);e!==a.page&&d("change",e,a.size)}return(l,e)=>(f(),r("div",v,[o("span",x,i(c(n)("common.pagerInfo",{total:t.total,page:t.page,pages:s.value})),1),o("button",{class:"btn btn-ghost btn-sm",type:"button",disabled:t.page<=1,onClick:e[0]||(e[0]=u=>g(t.page-1))},i(c(n)("common.previousPage")),9,k),o("button",{class:"btn btn-ghost btn-sm",type:"button",disabled:t.page>=s.value,onClick:e[1]||(e[1]=u=>g(t.page+1))},i(c(n)("common.nextPage")),9,M)]))}});export{y as _};
+1
View File
@@ -0,0 +1 @@
import{d as N,u as U,i as q,I,J as b,z as L,A as R,H as C,h as T,j as p,b as t,c as u,q as z,t as s,f as n,w as E,C as H,D as $,F as y,s as _,g as w,K as j,L as J,n as K,M as B,k as W,N as G,B as O,o as l}from"./index-DYsKpclu.js";import{P as Q}from"./PageShell-D87JkPkG.js";import{g as X,p as Y,b as Z}from"./share-CqDDLJWI.js";import"./SiteNav.vue_vue_type_script_setup_true_lang-CwaEJKIZ.js";const ee={class:"card",style:{"max-width":"640px",margin:"12px auto 0"}},te={key:0,class:"loading-block"},ae={class:"empty"},oe={style:{"font-weight":"600",color:"var(--c-text)"}},ne={class:"hint"},se=["placeholder"],ie={class:"btn",type:"submit"},le={style:{display:"flex","align-items":"center",gap:"10px","flex-wrap":"wrap","margin-bottom":"4px"}},ue={class:"badge"},re={style:{"font-size":"16px","word-break":"break-all"}},ce={key:0,class:"hint",style:{margin:"0"}},pe={class:"hint",style:{"margin-bottom":"18px"}},de={key:0,class:"loading-block"},me={class:"text-view"},ve={style:{display:"flex",gap:"10px","margin-top":"14px","flex-wrap":"wrap"}},ye={class:"file-summary"},ke={style:{"font-weight":"600","word-break":"break-all"}},fe={class:"hint",style:{margin:"0"}},ge={key:0,style:{margin:"16px 0 6px"}},he={class:"progress"},xe={class:"hint"},_e=["disabled"],Be=N({__name:"PickupView",setup(we){const{t:e}=U(),D=G(),F=O(),k=q(),d=T(()=>String(D.params.code??"").trim().split(/\s+/)[0]),r=p("loading"),m=p(""),o=p(null),v=p(""),f=p(!1),c=p(null);async function g(){if(!d.value){r.value="error",m.value=e("pickup.emptyCode");return}r.value="loading",m.value="",o.value=null,v.value="";try{const a=await X(d.value);if(o.value=a,r.value="ready",a.isText){f.value=!0;try{v.value=await Y(d.value)}finally{f.value=!1}}}catch(a){r.value="error",m.value=a instanceof C?a.code===404?e("pickup.notFound"):a.code===423?e("home.rateLimited"):a.code===428?e("home.notInitialized"):a.msg:e("pickup.failedDefault")}}I(g),b(d,()=>{g()}),b(()=>e("pickup.emptyCode"),()=>{r.value==="error"&&m.value&&g()});async function A(){if(o.value){c.value=0;try{const{blob:a,filename:i}=await Z(d.value,x=>c.value=x);B(a,i||o.value.name||"download"),k.success(e("pickup.downloaded"))}catch(a){k.error(a instanceof C?a.msg:e("pickup.downloadFailed"))}finally{c.value=null}}}async function M(){await W(v.value)?k.success(e("pickup.copied")):k.error(e("common.copyFailed"))}function P(){if(!o.value)return;const a=new Blob([v.value],{type:"text/plain;charset=utf-8"}),i=o.value.name?.includes(".")?o.value.name:`${o.value.name||"text"}.txt`;B(a,i)}const h=p("");function S(){const a=h.value.trim();a&&F.push({name:"pickup",params:{code:a}})}const V=T(()=>{const a=o.value;return a?a.remainingDownloads===null||a.remainingDownloads===void 0||a.remainingDownloads<0?e("pickup.remainingUnlimited"):e("pickup.remainingCount",{n:a.remainingDownloads}):""});return(a,i)=>(l(),L(Q,null,{default:R(()=>[t("section",ee,[r.value==="loading"?(l(),u("div",te,[i[1]||(i[1]=t("span",{class:"spin","aria-hidden":"true"},null,-1)),z(" "+s(n(e)("pickup.querying",{code:d.value})),1)])):r.value==="error"?(l(),u(y,{key:1},[t("div",ae,[i[2]||(i[2]=t("div",{class:"empty-icon"},"📮",-1)),t("p",oe,s(m.value||n(e)("pickup.failed")),1),t("p",ne,s(n(e)("pickup.confirmHint")),1)]),t("form",{class:"quick-pickup",style:{"margin-top":"6px","max-width":"none"},onSubmit:E(S,["prevent"])},[H(t("input",{"onUpdate:modelValue":i[0]||(i[0]=x=>h.value=x),class:"input input-mono",placeholder:n(e)("pickup.retryPlaceholder"),maxlength:"32"},null,8,se),[[$,h.value]]),t("button",ie,s(n(e)("pickup.retryButton")),1)],32)],64)):o.value?(l(),u(y,{key:2},[t("div",le,[t("span",ue,s(o.value.isText?n(e)("common.text"):n(e)("common.file")),1),t("strong",re,s(o.value.name),1),o.value.isText?w("",!0):(l(),u("span",ce,s(n(_)(o.value.size)),1))]),t("p",pe,s(V.value)+" · "+s(n(e)("pickup.expireAt",{time:o.value.expiredAt?n(j)(o.value.expiredAt):n(e)("time.permanent")}))+" · "+s(n(J)(o.value.expiredAt)),1),o.value.isText?(l(),u(y,{key:0},[f.value?(l(),u("div",de,[i[3]||(i[3]=t("span",{class:"spin","aria-hidden":"true"},null,-1)),z(" "+s(n(e)("pickup.loadingText")),1)])):(l(),u(y,{key:1},[t("pre",me,s(v.value),1),t("div",ve,[t("button",{class:"btn",type:"button",onClick:M},s(n(e)("pickup.copyContent")),1),t("button",{class:"btn btn-ghost",type:"button",onClick:P},s(n(e)("pickup.downloadTxt")),1)])],64))],64)):(l(),u(y,{key:1},[t("div",ye,[i[4]||(i[4]=t("span",{style:{"font-size":"30px"},"aria-hidden":"true"},"📄",-1)),t("div",null,[t("div",ke,s(o.value.name),1),t("div",fe,s(n(e)("pickup.sizeUsed",{size:n(_)(o.value.size),n:o.value.usedCount})),1)])]),c.value!==null?(l(),u("div",ge,[t("div",he,[t("i",{style:K({width:`${c.value}%`})},null,4)]),t("p",xe,s(n(e)("pickup.downloading",{percent:c.value})),1)])):w("",!0),t("button",{class:"btn btn-block",type:"button",disabled:c.value!==null,onClick:A}," ⬇ "+s(n(e)("pickup.downloadFile",{size:n(_)(o.value.size)})),9,_e)],64))],64)):w("",!0)])]),_:1}))}});export{Be as default};
File diff suppressed because one or more lines are too long
+1
View File
@@ -0,0 +1 @@
.settings-grid[data-v-03aa64c9]{display:grid;gap:16px}.save-row[data-v-03aa64c9]{display:flex;justify-content:flex-end}.notify-switch-row[data-v-03aa64c9]{display:flex;align-items:center;justify-content:space-between;gap:12px}.notify-switch-row label[data-v-03aa64c9]{margin-bottom:0}.unit-row[data-v-03aa64c9]{display:flex;gap:8px;align-items:center}.unit-row .input[data-v-03aa64c9]{flex:1}.unit-select[data-v-03aa64c9]{flex:0 0 110px!important}
@@ -0,0 +1 @@
import{d as z,u as E,a as x,c as i,G as B,f as t,A as m,b as a,F as _,r as v,az as S,t as r,h as d,aA as p,aB as T,O as A,o,z as M,q as V,e as b,aC as w}from"./index-DYsKpclu.js";const D={class:"site-nav"},F=["src"],I={class:"brand-name"},O=["aria-label"],R={class:"nav-controls"},U=["aria-label"],q=["aria-checked","title","onClick"],G={class:"nav-control-icon","aria-hidden":"true"},H=["title"],K=z({__name:"SiteNav",setup($){const{t:s}=E(),c=x(),{mode:u,setMode:k}=w(),g=d(()=>[{to:"/",label:s("nav.home"),match:n=>n==="/"}]);function y(n){return n.match(location.pathname)}const C={light:"☀️",dark:"🌙",system:"💻"},f=d(()=>({light:s("theme.light"),dark:s("theme.dark"),system:s("theme.system")})),L=d(()=>p()==="zh-CN"?"中文":"English");function N(){const n=p()==="zh-CN"?"en-US":"zh-CN";T(n)}return(n,l)=>{const h=A("RouterLink");return o(),i("header",D,[B(h,{class:"brand",to:"/",title:t(s)("nav.homeTitle",{name:t(c).displayName})},{default:m(()=>[a("img",{src:t(c).displayLogoUrl,alt:"Logo",onError:l[0]||(l[0]=e=>e.target.style.visibility="hidden")},null,40,F),a("span",I,r(t(c).displayName),1)]),_:1},8,["title"]),a("nav",{class:"nav-links","aria-label":t(s)("nav.mainNav")},[(o(!0),i(_,null,v(g.value,e=>(o(),M(h,{key:e.to,to:e.to,class:b({"router-link-active":y(e)})},{default:m(()=>[V(r(e.label),1)]),_:2},1032,["to","class"]))),128))],8,O),a("div",R,[a("div",{class:"theme-seg",role:"radiogroup","aria-label":t(s)("theme.label")},[(o(!0),i(_,null,v(t(S),e=>(o(),i("button",{key:e,type:"button",role:"radio","aria-checked":t(u)===e,class:b({active:t(u)===e}),title:f.value[e],onClick:j=>t(k)(e)},[a("span",G,r(C[e]),1)],10,q))),128))],8,U),a("button",{class:"nav-control",type:"button",title:t(s)("lang.label"),onClick:N},[l[1]||(l[1]=a("span",{class:"nav-control-icon","aria-hidden":"true"},"🌐",-1)),a("span",null,r(L.value),1)],8,H)])])}}});export{K as _};
+1
View File
@@ -0,0 +1 @@
import{v as n,S as e,x as s}from"./index-DYsKpclu.js";function u(t){return n(s.adminLogin,{method:"POST",json:{password:t}})}function m(){return n(s.adminVerify)}async function l(){await n(s.adminLogout,{method:"POST"})}async function g(t){const a=await n(s.adminFileList,{query:{page:t.page,size:t.size,keyword:t.keyword||void 0}}),o=e(a,["data","list","items","files"])??[],d=Number(e(a,["total","count"])??o.length);return{page:Number(e(a,["page"])??t.page),size:Number(e(a,["size"])??t.size),total:d,data:o.map(i=>({id:Number(e(i,["id"])??0),code:String(e(i,["code"])??""),name:String(e(i,["name"])??`${e(i,["prefix"])??""}${e(i,["suffix"])??""}`),suffix:String(e(i,["suffix"])??""),size:Number(e(i,["size"])??0),isText:!!(e(i,["isText","is_text"])??!1),expiredAt:e(i,["expiredAt","expired_at","expires_at"])??null,expiredCount:e(i,["expiredCount","expired_count"])??null,usedCount:Number(e(i,["usedCount","used_count"])??0),createdAt:e(i,["createdAt","created_at"])??null,isExpired:!!(e(i,["isExpired","is_expired"])??!1)}))}}async function p(t){await n(s.adminFileDelete,{method:"DELETE",json:{id:t}})}async function f(t){await n(s.adminFileBatchDelete,{method:"POST",json:{ids:t}})}async function y(t){await n(s.adminFileUpdate,{method:"PATCH",json:t})}async function w(){const t=await n(s.adminConfigGet);if(t&&typeof t=="object"&&!Array.isArray(t)){const a=t;return e(a,["config","data","settings"])??a}return{}}async function S(t){await n(s.adminConfigUpdate,{method:"PATCH",json:t})}async function _(t){const a=await n(s.adminStorageSwitch,{method:"POST",json:{engine:t}});return String(a?.engine??t)}async function x(t,a){await n(s.adminPasswordUpdate,{method:"PATCH",json:{old_password:t,new_password:a}})}async function b(t){const a=await n(s.adminAuditList,{query:{page:t.page,size:t.size,action:t.action||void 0,result:t.result||void 0,ip:t.ip||void 0,start_time:t.startTime||void 0,end_time:t.endTime||void 0}}),o=e(a,["data","list","items","logs"])??[],d=Number(e(a,["total","count"])??o.length);return{page:Number(e(a,["page"])??t.page),size:Number(e(a,["size"])??t.size),total:d,data:o.map((i,r)=>({id:Number(e(i,["id"])??r+1),action:String(e(i,["action"])??""),result:String(e(i,["result"])??""),fileCode:String(e(i,["file_code","fileCode","code"])??""),fileName:String(e(i,["file_name","fileName","name"])??""),sizeBytes:e(i,["size_bytes","sizeBytes","size"])??null,transferredBytes:e(i,["transferred_bytes","transferredBytes","bytes"])??null,ip:String(e(i,["ip","client_ip","clientIp"])??""),userAgent:String(e(i,["user_agent","userAgent"])??""),deviceOs:String(e(i,["device_os","deviceOs","os"])??""),deviceBrowser:String(e(i,["device_browser","deviceBrowser","browser"])??""),deviceType:String(e(i,["device_type","deviceType"])??""),actor:String(e(i,["actor"])??""),errorMsg:String(e(i,["error_msg","errorMsg","error"])??""),durationMs:e(i,["duration_ms","durationMs","duration"])??null,createdAt:e(i,["created_at","createdAt","time"])??null}))}}export{g as a,f as b,p as c,y as d,b as e,w as f,_ as g,S as h,x as i,l as j,m as k,u as l};
+1
View File
@@ -0,0 +1 @@
import{av as o,aw as a,W as r}from"./index-DYsKpclu.js";import{j as n,k as i,l as c}from"./admin-DEOkRyTC.js";const s="fcb_admin_flag",l=o("auth",{state:()=>({token:r(),checked:!1}),getters:{isAuthed:t=>!!t.token},actions:{async login(t){const e=(await c(t))?.token;if(!e)throw new Error("登录响应缺少 token");this.token=e,a(e);try{localStorage.setItem(s,"1")}catch{}this.checked=!0},async verify(){if(!this.token)return!1;try{return await i(),this.checked=!0,!0}catch{return this.reset(),!1}},async logout(){try{this.token&&await n()}catch{}this.reset()},reset(){this.token="",this.checked=!1,a("");try{localStorage.removeItem(s)}catch{}}}});export{l as u};
File diff suppressed because one or more lines are too long
Binary file not shown.

After

Width:  |  Height:  |  Size: 60 KiB

File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long

After

Width:  |  Height:  |  Size: 12 KiB

File diff suppressed because one or more lines are too long
+1
View File
@@ -0,0 +1 @@
import{v as x,S as s,U as m,V as g,x as u,W as w,H as i,X as f,Y as h}from"./index-DYsKpclu.js";function T(n,t,r,a=""){return x(u.shareText,{method:"POST",form:{text:n,expire_value:String(t),expire_style:r,code:a}})}function y(n,t,r,a,o=""){const e=new FormData;return e.append("file",n),e.append("expire_value",String(t)),e.append("expire_style",r),o&&e.append("code",o),h(u.shareFile,e,a)}async function S(n){const t=await x(u.shareMetadata,{query:{code:n}}),r=!!(s(t,["is_text","isText"])??s(t,["type"])==="text");return{code:String(s(t,["code"])??n),name:String(s(t,["name"])??"未命名"),size:Number(s(t,["size"])??0),type:r?"text":"file",isText:r,createdAt:s(t,["created_at","createdAt"])??null,expiredAt:s(t,["expired_at","expires_at","expiredAt","expiresAt"])??null,expiredCount:s(t,["expired_count","expiredCount","remaining_downloads","remainingDownloads"])??null,usedCount:Number(s(t,["used_count","usedCount"])??0),remainingDownloads:s(t,["remaining_downloads","remainingDownloads","expired_count","expiredCount"])??null}}async function b(n){const{blob:t}=await m(u.shareSelect,{query:{code:n}});return t.text()}function A(n,t){return new Promise((r,a)=>{const o=new URL(g(u.shareSelect),location.origin);o.searchParams.set("code",n);const e=new XMLHttpRequest;e.open("GET",o.toString()),e.timeout=6e5,e.responseType="blob";const l=w();l&&e.setRequestHeader("Authorization",`Bearer ${l}`),e.onprogress=p=>{p.lengthComputable&&t&&t(Math.round(p.loaded/p.total*100))},e.onload=()=>{if((e.getResponseHeader("content-type")??"").includes("application/json")){const d=new FileReader;d.onload=()=>{try{const c=JSON.parse(String(d.result));a(new i(c.code??e.status,c.msg||"取件失败",e.status))}catch{a(new i(e.status,"取件失败",e.status))}},d.readAsText(e.response);return}e.status>=200&&e.status<300?r({blob:e.response,filename:f(e.getResponseHeader("content-disposition"))}):a(new i(e.status,`取件失败(HTTP ${e.status}`,e.status))},e.onerror=()=>a(new i(0,"网络异常,取件失败")),e.ontimeout=()=>a(new i(0,"下载超时,请重试")),e.send()})}export{y as a,A as b,S as g,b as p,T as s};
+1
View File
@@ -0,0 +1 @@
function e(t){return t&&t.__esModule&&Object.prototype.hasOwnProperty.call(t,"default")?t.default:t}export{e as g};
+36
View File
@@ -0,0 +1,36 @@
<!doctype html>
<html lang="zh-CN">
<head>
<meta charset="UTF-8" />
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
<meta name="description" content="文件快传 - 开箱即用的文件快传系统" />
<!-- 需求②:首帧防闪烁——读 localStorage 主题模式并写入 html[data-theme]App 启动后由 theme 模块接管) -->
<script>
;(function () {
try {
var m = localStorage.getItem('fcb_theme_mode')
var dark = m === 'dark' || (m !== 'light' && window.matchMedia('(prefers-color-scheme: dark)').matches)
document.documentElement.dataset.theme = dark ? 'dark' : 'light'
} catch (e) {
document.documentElement.dataset.theme = 'light'
}
})()
</script>
<!-- 需求④:默认 favicon 使用本地打包资源;管理端可在系统设置中自定义 favicon_url 全站替换 -->
<link rel="icon" type="image/png" href="/assets/favicon-Dl6ZLL7S.png" />
<title>文件快传 · 文件快传</title>
<style>
html {
background: #eef1f6;
}
html[data-theme='dark'] {
background: #0b0f1a;
}
</style>
<script type="module" crossorigin src="/assets/index-DYsKpclu.js"></script>
<link rel="stylesheet" crossorigin href="/assets/index-CtyCxWf5.css">
</head>
<body>
<div id="app"></div>
</body>
</html>
+21
View File
@@ -0,0 +1,21 @@
// Package web 提供前端 SPA 的 go:embed 嵌入。
//
// dist/ 由前端任务(vue-dev)构建产出;deploy 构建脚本会把仓库根
// web/dist 拷贝到本目录后再编译(Dockerfile 已处理)。
// 本目录内的 index.html 是构建产物的一部分,随产物更新覆盖。
package web
import (
"embed"
"io/fs"
)
// distFS 嵌入 dist 全部产物。
//
//go:embed all:dist
var distFS embed.FS
// Dist 返回嵌入的前端文件系统(根为 dist/)。
func Dist() (fs.FS, error) {
return fs.Sub(distFS, "dist")
}