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:
@@ -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、全部业务查询方言无关;
|
||||
唯一原生 DDL(migrates 台账表)在 `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→failed,401/403/423/429/428→denied)
|
||||
```
|
||||
|
||||
下载响应字节数由中间件自动统计;上传字节数由 handler 填 `TransferredBytes`。
|
||||
|
||||
## 限流语义(对齐参考实现)
|
||||
|
||||
- `error`(取件错误)/`login`(登录失败):**仅在失败时计数**,handler 调用 `limiter.Add(c, kind)`
|
||||
- `upload`:**成功上传才计数**(先 `Check` 放行,成功后 `Add`)
|
||||
- `metadata`:每次访问即计数,可用 `RequireRateLimit` 中间件
|
||||
- 超限返回 HTTP 423;规则来自 settings KV(errorCount/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 的变量
|
||||
```
|
||||
@@ -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())
|
||||
}
|
||||
// v3:Manager 包装——保存/读取委托当前引擎;管理端可热切换
|
||||
//(构建闭包在每次切换前用最新 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、/setup:1MiB(文本内容本身限 222KB);
|
||||
// - 其余(上传类):max_file_size(0=回落 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
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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
@@ -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 单分片大小上限 32MB(M3:限制 io.ReadAll 内存占用)。
|
||||
const maxChunkSizeBytes = 32 * 1024 * 1024
|
||||
|
||||
// ============ POST /chunk/upload/init 初始化分片会话 ============
|
||||
|
||||
// requireChunkEnabled L4:enableChunk 开关后端强制(此前仅前端隐藏入口,
|
||||
// 开关关闭后 /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_size,0=回落 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/upload,upload_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": "上传已取消"})
|
||||
}
|
||||
@@ -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_type(secret/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 生成下载令牌(L2:HMAC-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_count(0=不限制,超限 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())
|
||||
}
|
||||
// 原子条件插入:已用 + 生效预留 + 本次 <= 上限。
|
||||
// Info5:Postgres 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 兼容归一化:旧前端 bundle(fetch 字符串 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]
|
||||
}
|
||||
@@ -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 应报错")
|
||||
}
|
||||
// day:7 天内合法
|
||||
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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,45 @@
|
||||
// policy.go — v2 上传策略统一读取与校验(需求 ④⑩)。
|
||||
//
|
||||
// 管理端在后台设置页修改策略(settings KV,t1 schema)后,上传链路
|
||||
// (share/file、chunk、presign)每次请求实时读取当前策略并动态校验:
|
||||
// - 大小上限:max_file_size(0=回落 uploadSize,语义见 config.MaxFileSize);
|
||||
// - 类型白名单:allowed_file_types("*" 不限制),由 validateFileMagic 统一执行;
|
||||
// - 保存策略:expire_style 白名单 + max_save_seconds 时间上限 + max_save_count
|
||||
// 次数上限,统一在 resolveExpire(helpers.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
|
||||
}
|
||||
@@ -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 构造带真实依赖的 Deps:sqlite 文件库(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:包装为 Manager(build 直接返回 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'}
|
||||
|
||||
// ============ ① 公开 config:v2 展示与策略字段下发 ============
|
||||
|
||||
// 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/update:v2 键全链路 + 类型范围校验 ============
|
||||
|
||||
// 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
|
||||
}
|
||||
@@ -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_size,0=回落 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))
|
||||
}
|
||||
@@ -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)
|
||||
// Info3:robotsText 配置键此前无路由承接,补上(内容可在管理端自定义)
|
||||
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_secret,settings.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 token(403)。
|
||||
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
|
||||
}
|
||||
@@ -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 L4:enableChunk=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 M3:chunk_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 L3:4 位自定义码拒绝、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("错误密码不应通过")
|
||||
}
|
||||
}
|
||||
@@ -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("&", "&", "<", "<", ">", ">", `"`, "'")
|
||||
return r.Replace(s)
|
||||
}
|
||||
@@ -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(配合全局 BodyLimit;441KB 为 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_size,0=回落 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
|
||||
}
|
||||
@@ -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 表单调用 handler(v3.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 body(fetch 字符串 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())
|
||||
}
|
||||
}
|
||||
@@ -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.html(SPA 回退用)。
|
||||
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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
Vendored
+42
@@ -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-15(cluster 模式忽略)
|
||||
}
|
||||
|
||||
// New 按配置构造缓存实现:redisAddr 为空 → 内存实现。
|
||||
func New(ctx context.Context, opt RedisOptions) (Cache, error) {
|
||||
if opt.Addr == "" {
|
||||
return NewMemory(), nil
|
||||
}
|
||||
return NewRedis(ctx, opt.Addr, opt.DB)
|
||||
}
|
||||
Vendored
+81
@@ -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)
|
||||
}
|
||||
}
|
||||
Vendored
+159
@@ -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:])
|
||||
}
|
||||
Vendored
+117
@@ -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
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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_DRIVER,sqlite|postgres,默认 sqlite(需求 ⑧)
|
||||
DBDSN string // FCB_DB_DSN;postgres 必需;sqlite 为空时用 DefaultSQLitePath
|
||||
RedisAddr string // FCB_REDIS_ADDR,可选;为空时缓存降级为内存实现
|
||||
RedisDB int // FCB_REDIS_DB,Redis 逻辑库号 0-15,默认 0(URL 形式地址以 URL 内库号优先)
|
||||
Listen string // FCB_LISTEN,监听地址,默认 :8466
|
||||
StorageEngine string // FCB_STORAGE_ENGINE,local|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_url;background 为参考实现既有键,保留兼容)
|
||||
"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_DB,0-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_DSN(Postgres 连接串)")
|
||||
}
|
||||
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.StorageEngine(env 校验过的 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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,97 @@
|
||||
// Package config — schema.go 定义 v2 新增配置键(KV)schema:
|
||||
// 键名常量、类型、默认值与取值边界。管理与 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 GiB,0 表示回落 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],不带路径;空=分享链接用当前访问地址)"},
|
||||
}
|
||||
}
|
||||
@@ -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("页脚默认应为空")
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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 为 *string(JSON 文本),双方言 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 + LIKE(admin 列表检索路径:真实代码先对关键词小写化再拼 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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
// 未命中审计动作的请求直接放行,不产生审计记录。
|
||||
// L5:admin 类动作同样需要建 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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,30 @@
|
||||
// Package middleware — bodylimit.go:全局请求体大小限制。
|
||||
//
|
||||
// 安全审计 M3:此前服务对请求体无任何上限,text/plain 形态的绑定会先
|
||||
// io.ReadAll 全量进内存、multipart 大文件先完整落盘后才被策略拒绝,
|
||||
// 未认证攻击者可借大 body 消耗内存/磁盘/带宽。
|
||||
// 本中间件用 http.MaxBytesReader 按「路径类别」给出上限:
|
||||
// - 管理端(/admin/*):1MiB;
|
||||
// - 文本/元信息/取件类(/share/text|metadata|select):1MiB;
|
||||
// - 其余(含上传):maxFileSize(0=回落 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()
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
@@ -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_secret;expires 为会话有效期。
|
||||
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_secret(settings 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()
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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"` // 归属存储引擎(v3:local|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()...)
|
||||
}
|
||||
@@ -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})
|
||||
}
|
||||
@@ -0,0 +1,150 @@
|
||||
// settings 包双方言测试:Manager 全流程(ensure 行、KV 读写合并、Reload、
|
||||
// UpdateKV 屏蔽内部键、SystemStart)分别在 sqlite(默认)与 postgres(FCB_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. SystemStart:sys_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)
|
||||
}
|
||||
@@ -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 工作因子:12(2026 年桌面 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 与明文一律 true;bcrypt 成本低于当前 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)
|
||||
}
|
||||
@@ -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("密钥未随机化")
|
||||
}
|
||||
}
|
||||
@@ -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("&", "&", "<", "<", ">", ">", `"`, """, "'", "'")
|
||||
return r.Replace(s)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,87 @@
|
||||
// Package settings — schema.go:v2 配置键 schema 常量与元数据表。
|
||||
//
|
||||
// 键名常量的单一事实来源在 internal/config/schema.go(defaults() 需引用);
|
||||
// 本文件 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
|
||||
}
|
||||
@@ -0,0 +1,172 @@
|
||||
// Package settings 提供数据库 settings KV 的运行时读写:
|
||||
// env(FCB_*)提供基线,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 }
|
||||
@@ -0,0 +1,10 @@
|
||||
package storage
|
||||
|
||||
import "errors"
|
||||
|
||||
// ErrRangeNotSatisfiable 请求的字节范围超出文件大小(对应 HTTP 416)。
|
||||
// 属于增量错误定义,不改动 interface.go 的既有签名。
|
||||
var ErrRangeNotSatisfiable = errors.New("storage: 请求范围超出文件大小")
|
||||
|
||||
// ErrHashMismatch 分片哈希校验失败(合并分片时检测到数据损坏)。
|
||||
var ErrHashMismatch = errors.New("storage: 分片哈希校验失败")
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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=0,End=Total-1);
|
||||
// - rng 非 nil:返回 [Start, End] 区间流。
|
||||
// 引擎应尽量透传 Range(WebDAV/S3)或按块 seek(local)。
|
||||
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
|
||||
}
|
||||
@@ -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)
|
||||
@@ -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("未知引擎应报错")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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 构造测试用 Manager:local 健康引擎起步;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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,70 @@
|
||||
package storage
|
||||
|
||||
// EngineOptions 引擎构造选项:由 main.go(API 层任务)从 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_url(MinIO 等;空则 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
|
||||
}
|
||||
}
|
||||
@@ -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:])
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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 等)。
|
||||
//
|
||||
// 相比参考实现(S3FileStorage,aioboto3)的改进:
|
||||
// - 单例客户端 + 自定义连接池 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 需要可重放流
|
||||
// 或 TLS,MinIO/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 // 5MB,S3 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 会话 ID(Abort 时复用 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 健康检查:列举 bucket(MaxKeys=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)
|
||||
@@ -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
|
||||
}
|
||||
@@ -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 引擎(本次重写的重点优化对象)。
|
||||
//
|
||||
// 相比参考实现(WebDAVFileStorage,aiohttp)的改进:
|
||||
// - 单例 http.Client + 连接池化 Transport(参考实现每个操作新建 ClientSession,无复用);
|
||||
// - Basic 与 Digest(RFC 2617,qop=auth,MD5/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 连接池化 Transport:Keep-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 将远端路径转为完整 URL(URL.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)
|
||||
@@ -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=auth(MD5 / 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)
|
||||
}
|
||||
@@ -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|digest;digest 配合 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 Digest(MD5)认证协商。
|
||||
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)
|
||||
}
|
||||
|
||||
// 认证后的 PROPFIND(Stat 已有目录)应得到 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 Digest(SHA-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 可重放 body(seekable)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)
|
||||
}
|
||||
// 非 seekable(io.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 响应慢于 Timeout(1s)→ 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)
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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)}
|
||||
+1
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
+1
File diff suppressed because one or more lines are too long
@@ -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}
|
||||
@@ -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}
|
||||
@@ -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};
|
||||
@@ -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};
|
||||
+267
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
@@ -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};
|
||||
@@ -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 _};
|
||||
@@ -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};
|
||||
+129
File diff suppressed because one or more lines are too long
@@ -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}
|
||||
+1
@@ -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
@@ -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
@@ -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};
|
||||
+2517
File diff suppressed because one or more lines are too long
BIN
Binary file not shown.
|
After Width: | Height: | Size: 60 KiB |
+1
File diff suppressed because one or more lines are too long
+28
File diff suppressed because one or more lines are too long
+3
File diff suppressed because one or more lines are too long
|
After Width: | Height: | Size: 12 KiB |
+60
File diff suppressed because one or more lines are too long
+1
@@ -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
@@ -0,0 +1 @@
|
||||
function e(t){return t&&t.__esModule&&Object.prototype.hasOwnProperty.call(t,"default")?t.default:t}export{e as g};
|
||||
Vendored
+36
@@ -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>
|
||||
@@ -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")
|
||||
}
|
||||
Reference in New Issue
Block a user