diff --git a/.gitignore b/.gitignore index ecd454ca..9ebd0f56 100644 --- a/.gitignore +++ b/.gitignore @@ -42,3 +42,4 @@ cloudreve .playwright-mcp/ # Monorepo staging dir for asset packaging (transient) /assets/ +pkg/util/test/ diff --git a/ROADMAP.md b/ROADMAP.md index 5dd7ceb2..652df88b 100644 --- a/ROADMAP.md +++ b/ROADMAP.md @@ -136,7 +136,8 @@ Order = user-visible value first; each ships with backend + UI + tests. ## 5. Phase C — security + quality - Own security review on top of upstream fixes: session/token entropy audit, SSRF guard re-test (NAT64 class), rate limiting on auth endpoints -- Fix upstream bug backlog by impact: #3574 OOM (trash_bin_collect streaming), #3118/#3005 WebDAV large-file, #3454 PG FK, #3375 SMTP auth discovery +- Fix upstream bug backlog by impact: ~~#3574 OOM~~ (done — paged tree walk + batched delete), ~~#3118/#3005 WebDAV large-file~~ (done — Content-Range assembly into one session; non-local policies get honest 501; single-PUT giant-file 500s are proxy/client timeouts, not fixable server-side), ~~#3375 SMTP auth discovery~~ (done — `smtp_auth` setting) +- #3454 (PG FK on upload) is **Pro-only** — `audit_logs` doesn't exist in this codebase. When B.5 adds our own audit log: insert the audit row in the same tx *after* the file row, never before. - `desloppify` + `security-reviewer` passes; scorecard appended to README ## 6. Phase D — desktop, all platforms diff --git a/frontend/public/locales/en-US/dashboard.json b/frontend/public/locales/en-US/dashboard.json index bad878b6..5d8aa001 100644 --- a/frontend/public/locales/en-US/dashboard.json +++ b/frontend/public/locales/en-US/dashboard.json @@ -437,6 +437,18 @@ "replyToAddressDes": "The mailbox used to receive reply emails when users reply to emails sent by the system.", "enforceSSL": "Enforce SSL connection", "enforceSSLDes": "Whether to enforce an SSL encrypted connection. If you cannot send emails, you can turn this off and Cloudreve will try to use STARTTLS and decide whether to use encrypted connections.", + "smtpAuthMethod": "SMTP authentication method", + "smtpAuthMethodDes": "Auto-discovery only selects secure mechanisms and never sends credentials in plaintext on unencrypted connections. If your server only offers PLAIN/LOGIN without encryption, pick the corresponding \"-noenc\" method — credentials will then be transmitted unencrypted.", + "smtpAuth_autodiscover": "Auto discovery", + "smtpAuth_plain": "PLAIN", + "smtpAuth_plain_noenc": "PLAIN (allow unencrypted)", + "smtpAuth_login": "LOGIN", + "smtpAuth_login_noenc": "LOGIN (allow unencrypted)", + "smtpAuth_cram_md5": "CRAM-MD5", + "smtpAuth_scram_sha_1": "SCRAM-SHA-1", + "smtpAuth_scram_sha_256": "SCRAM-SHA-256", + "smtpAuth_xoauth2": "XOAUTH2", + "smtpAuth_noauth": "No authentication", "smtpTTL": "SMTP connection TTL (seconds)", "smtpTTLDes": "SMTP connections established during the TTL period will be reused by new mail delivery requests.", "emailTemplates": "Email Templates", diff --git a/frontend/public/locales/zh-CN/dashboard.json b/frontend/public/locales/zh-CN/dashboard.json index 15867876..c0d0f2ba 100644 --- a/frontend/public/locales/zh-CN/dashboard.json +++ b/frontend/public/locales/zh-CN/dashboard.json @@ -437,6 +437,18 @@ "replyToAddressDes": "用户回复系统发送的邮件时,用于接收回信的邮箱。", "enforceSSL": "强制使用 SSL 连接", "enforceSSLDes": "是否强制使用 SSL 加密连接。如果无法发送邮件,可关闭此项,Cloudreve 会尝试使用 STARTTLS 并决定是否使用加密连接。", + "smtpAuthMethod": "SMTP 认证方式", + "smtpAuthMethodDes": "自动发现只会选择安全的认证机制,不会在未加密的连接上以明文发送凭据。如果服务器仅提供 PLAIN/LOGIN 且不支持加密,请选择对应的 \"-noenc\" 方式 —— 凭据将以明文传输。", + "smtpAuth_autodiscover": "自动发现", + "smtpAuth_plain": "PLAIN", + "smtpAuth_plain_noenc": "PLAIN(允许明文传输)", + "smtpAuth_login": "LOGIN", + "smtpAuth_login_noenc": "LOGIN(允许明文传输)", + "smtpAuth_cram_md5": "CRAM-MD5", + "smtpAuth_scram_sha_1": "SCRAM-SHA-1", + "smtpAuth_scram_sha_256": "SCRAM-SHA-256", + "smtpAuth_xoauth2": "XOAUTH2", + "smtpAuth_noauth": "不使用认证", "smtpTTL": "SMTP 连接有效期 (秒)", "smtpTTLDes": "有效期内建立的 SMTP 连接会被新邮件发送请求复用。", "emailTemplates": "邮件模板", diff --git a/frontend/src/component/Admin/Settings/Email/Email.tsx b/frontend/src/component/Admin/Settings/Email/Email.tsx index 0dc38f42..a80909b5 100644 --- a/frontend/src/component/Admin/Settings/Email/Email.tsx +++ b/frontend/src/component/Admin/Settings/Email/Email.tsx @@ -1,4 +1,4 @@ -import { Box, DialogContent, FormControl, FormControlLabel, Stack, Switch, Typography } from "@mui/material"; +import { Box, DialogContent, FormControl, FormControlLabel, ListItemText, Stack, Switch, Typography } from "@mui/material"; import { useSnackbar } from "notistack"; import { useContext, useState } from "react"; import { useTranslation } from "react-i18next"; @@ -6,9 +6,10 @@ import { sendTestSMTP } from "../../../../api/api.ts"; import { useAppDispatch } from "../../../../redux/hooks.ts"; import { isTrueVal } from "../../../../session/utils.ts"; import { DefaultCloseAction } from "../../../Common/Snackbar/snackbar.tsx"; -import { DenseFilledTextField, SecondaryButton } from "../../../Common/StyledComponents.tsx"; +import { DenseFilledTextField, DenseSelect, SecondaryButton } from "../../../Common/StyledComponents.tsx"; import DraggableDialog, { StyledDialogContentText } from "../../../Dialogs/DraggableDialog.tsx"; import MailOutlined from "../../../Icons/MailOutlined.tsx"; +import { SquareMenuItem } from "../../../FileManager/ContextMenu/ContextMenu.tsx"; import SettingForm from "../../../Pages/Setting/SettingForm.tsx"; import { NoMarginHelperText, SettingSection, SettingSectionContent } from "../Settings.tsx"; import { SettingContext } from "../SettingWrapper.tsx"; @@ -174,6 +175,39 @@ const Email = () => { + + + setSettings({ smtp_auth: e.target.value as string })} + > + {[ + "autodiscover", + "plain", + "plain-noenc", + "login", + "login-noenc", + "cram-md5", + "scram-sha-1", + "scram-sha-256", + "xoauth2", + "noauth", + ].map((v) => ( + + + {t(`settings.smtpAuth_${v.replace(/-/g, "_")}`)} + + + ))} + + {t("settings.smtpAuthMethodDes")} + + + { "smtpUser", "smtpPass", "smtpEncryption", + "smtp_auth", "fromName", "mail_activation_template", "mail_reset_template", diff --git a/inventory/setting.go b/inventory/setting.go index e822d9ae..064a6d2d 100644 --- a/inventory/setting.go +++ b/inventory/setting.go @@ -561,6 +561,7 @@ var DefaultSettings = map[string]string{ "email_filter_mode": "0", "email_filter_list": "", "email_disable_subaddress": "0", + "smtp_auth": "autodiscover", "captcha_type": "normal", "captcha_height": "60", "captcha_width": "240", diff --git a/middleware/ratelimit.go b/middleware/ratelimit.go new file mode 100644 index 00000000..bcd824fd --- /dev/null +++ b/middleware/ratelimit.go @@ -0,0 +1,51 @@ +package middleware + +import ( + "strconv" + "time" + + "github.com/cloudreve/Cloudreve/v4/application/dependency" + "github.com/cloudreve/Cloudreve/v4/pkg/serializer" + "github.com/gin-gonic/gin" +) + +const rateLimitPrefix = "rate_limit:" + +// RateLimit applies a fixed-window rate limit to requests sharing the same +// bucket key. Once `limit` requests arrive within `window`, further requests +// are rejected until the window expires. State lives in the shared KV store, +// so limits hold across cluster nodes when Redis is configured. +// +// The counter is approximate (read-then-write), which is acceptable for +// abuse prevention — it bounds attempts to roughly `limit` per window. +func RateLimit(limit int, window time.Duration, key func(*gin.Context) string) gin.HandlerFunc { + return func(c *gin.Context) { + kv := dependency.FromContext(c).KV() + bucket := rateLimitPrefix + key(c) + + count := 0 + if raw, ok := kv.Get(bucket); ok { + if v, ok := raw.(int); ok { + count = v + } + } + + if count >= limit { + c.Header("Retry-After", strconv.Itoa(int(window.Seconds()))) + c.JSON(200, serializer.NewError(serializer.CodeRateLimited, "Too many requests, please try again later.", nil)) + c.Abort() + return + } + + _ = kv.Set(bucket, count+1, int(window.Seconds())) + c.Next() + } +} + +// RateLimitByIP keys the rate limit on the client IP and a static bucket +// name, e.g. RateLimitByIP("login", 10, time.Minute). +func RateLimitByIP(bucket string, limit int, window time.Duration) gin.HandlerFunc { + return RateLimit(limit, window, func(c *gin.Context) string { + return bucket + ":" + c.ClientIP() + }) +} diff --git a/middleware/ratelimit_test.go b/middleware/ratelimit_test.go new file mode 100644 index 00000000..1b9a72dc --- /dev/null +++ b/middleware/ratelimit_test.go @@ -0,0 +1,123 @@ +package middleware + +import ( + "context" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/cloudreve/Cloudreve/v4/application/dependency" + "github.com/cloudreve/Cloudreve/v4/pkg/cache" + "github.com/gin-gonic/gin" +) + +var testEngine = func() *gin.Engine { + e := gin.New() + e.ContextWithFallback = true + return e +}() + +func newRateLimitContext(t *testing.T, dep dependency.Dep, remoteAddr string) (*gin.Context, *httptest.ResponseRecorder) { + t.Helper() + w := httptest.NewRecorder() + c := gin.CreateTestContextOnly(w, testEngine) + req := httptest.NewRequest(http.MethodPost, "/session/token", nil) + req.RemoteAddr = remoteAddr + ":12345" + c.Request = req.WithContext(context.WithValue(req.Context(), dependency.DepCtx{}, dep)) + return c, w +} + +func TestRateLimitByIP(t *testing.T) { + gin.SetMode(gin.TestMode) + dep := dependency.NewDependency(dependency.WithKV(cache.NewMemoStore("", nil))) + + handler := RateLimitByIP("login", 3, time.Minute) + + for i := 0; i < 3; i++ { + c, _ := newRateLimitContext(t, dep, "192.0.2.1") + handler(c) + if c.IsAborted() { + t.Fatalf("request %d unexpectedly rejected", i+1) + } + } + + c, w := newRateLimitContext(t, dep, "192.0.2.1") + handler(c) + if !c.IsAborted() { + t.Fatal("request over limit was not rejected") + } + if w.Header().Get("Retry-After") != "60" { + t.Fatalf("missing/incorrect Retry-After: %q", w.Header().Get("Retry-After")) + } + + // A different IP has its own bucket. + c, _ = newRateLimitContext(t, dep, "192.0.2.2") + handler(c) + if c.IsAborted() { + t.Fatal("unrelated IP shared the bucket") + } + + // A different bucket name on the same IP has its own counter. + handler2 := RateLimitByIP("register", 1, time.Minute) + c, _ = newRateLimitContext(t, dep, "192.0.2.1") + handler2(c) + if c.IsAborted() { + t.Fatal("unrelated bucket shared the counter") + } +} + +func TestRateLimitExpiry(t *testing.T) { + gin.SetMode(gin.TestMode) + dep := dependency.NewDependency(dependency.WithKV(cache.NewMemoStore("", nil))) + + handler := RateLimitByIP("login", 1, time.Second) + + c, _ := newRateLimitContext(t, dep, "192.0.2.1") + handler(c) + if c.IsAborted() { + t.Fatal("first request rejected") + } + + c, _ = newRateLimitContext(t, dep, "192.0.2.1") + handler(c) + if !c.IsAborted() { + t.Fatal("second request not rejected") + } + + // Window expiry frees the bucket. MemoStore TTL is Unix-second granular + // (item valid while Expires >= now), so a 1s window can live ~2s. + time.Sleep(2100 * time.Millisecond) + + c, _ = newRateLimitContext(t, dep, "192.0.2.1") + handler(c) + if c.IsAborted() { + t.Fatal("request after window expiry still rejected") + } +} + +func TestRateLimitCountOverflow(t *testing.T) { + gin.SetMode(gin.TestMode) + kv := cache.NewMemoStore("", nil) + dep := dependency.NewDependency(dependency.WithKV(kv)) + + // A non-int value in the bucket is treated as a fresh window. + if err := kv.Set(rateLimitPrefix+"login:192.0.2.1", "garbage", 60); err != nil { + t.Fatal(err) + } + + handler := RateLimitByIP("login", 1, time.Minute) + c, _ := newRateLimitContext(t, dep, "192.0.2.1") + handler(c) + if c.IsAborted() { + t.Fatal("corrupt bucket state rejected request") + } + + raw, ok := kv.Get(rateLimitPrefix + "login:192.0.2.1") + if !ok { + t.Fatal("bucket not persisted") + } + if _, isInt := raw.(int); !isInt { + t.Fatalf("bucket value type drifted: %T", raw) + } +} diff --git a/pkg/email/smtp.go b/pkg/email/smtp.go index ebb5c74f..ad9b6468 100644 --- a/pkg/email/smtp.go +++ b/pkg/email/smtp.go @@ -102,6 +102,36 @@ func (client *SMTPPool) Send(ctx context.Context, to, title, body string) error return nil } +// SMTPAuthType maps the admin-configured auth method to a go-mail mechanism. +// The library's own auto-discovery never picks plaintext mechanisms +// (PLAIN/LOGIN) on unencrypted connections, so relays that only advertise +// them need the explicit *-noenc choices. Anything unrecognized falls back +// to auto-discovery. +func SMTPAuthType(configured string) mail.SMTPAuthType { + switch strings.ToUpper(strings.TrimSpace(configured)) { + case string(mail.SMTPAuthPlain): + return mail.SMTPAuthPlain + case string(mail.SMTPAuthPlainNoEnc): + return mail.SMTPAuthPlainNoEnc + case string(mail.SMTPAuthLogin): + return mail.SMTPAuthLogin + case string(mail.SMTPAuthLoginNoEnc): + return mail.SMTPAuthLoginNoEnc + case string(mail.SMTPAuthCramMD5): + return mail.SMTPAuthCramMD5 + case string(mail.SMTPAuthSCRAMSHA1): + return mail.SMTPAuthSCRAMSHA1 + case string(mail.SMTPAuthSCRAMSHA256): + return mail.SMTPAuthSCRAMSHA256 + case string(mail.SMTPAuthXOAUTH2): + return mail.SMTPAuthXOAUTH2 + case string(mail.SMTPAuthNoAuth): + return mail.SMTPAuthNoAuth + default: + return mail.SMTPAuthAutoDiscover + } +} + // Close 关闭发送队列 func (client *SMTPPool) Close() { if client.ch != nil { @@ -125,7 +155,7 @@ func (client *SMTPPool) Init() { opts := []mail.Option{ mail.WithPort(client.config.Port), mail.WithTimeout(time.Duration(client.config.Keepalive+5) * time.Second), - mail.WithSMTPAuth(mail.SMTPAuthAutoDiscover), mail.WithTLSPortPolicy(mail.TLSOpportunistic), + mail.WithSMTPAuth(SMTPAuthType(client.config.AuthType)), mail.WithTLSPortPolicy(mail.TLSOpportunistic), mail.WithUsername(client.config.User), mail.WithPassword(client.config.Password), } if client.config.ForceEncryption { diff --git a/pkg/email/smtp_test.go b/pkg/email/smtp_test.go new file mode 100644 index 00000000..7b270b33 --- /dev/null +++ b/pkg/email/smtp_test.go @@ -0,0 +1,27 @@ +package email + +import ( + "testing" + + "github.com/stretchr/testify/assert" + mail "github.com/wneessen/go-mail" +) + +func TestSMTPAuthType(t *testing.T) { + assert.Equal(t, mail.SMTPAuthAutoDiscover, SMTPAuthType("")) + assert.Equal(t, mail.SMTPAuthAutoDiscover, SMTPAuthType("autodiscover")) + assert.Equal(t, mail.SMTPAuthAutoDiscover, SMTPAuthType("bogus")) + + assert.Equal(t, mail.SMTPAuthPlain, SMTPAuthType("plain")) + assert.Equal(t, mail.SMTPAuthPlainNoEnc, SMTPAuthType("plain-noenc")) + assert.Equal(t, mail.SMTPAuthLogin, SMTPAuthType("login")) + assert.Equal(t, mail.SMTPAuthLoginNoEnc, SMTPAuthType("login-noenc")) + assert.Equal(t, mail.SMTPAuthCramMD5, SMTPAuthType("cram-md5")) + assert.Equal(t, mail.SMTPAuthSCRAMSHA1, SMTPAuthType("scram-sha-1")) + assert.Equal(t, mail.SMTPAuthSCRAMSHA256, SMTPAuthType("scram-sha-256")) + assert.Equal(t, mail.SMTPAuthXOAUTH2, SMTPAuthType("xoauth2")) + assert.Equal(t, mail.SMTPAuthNoAuth, SMTPAuthType("noauth")) + + // Values are normalized before matching. + assert.Equal(t, mail.SMTPAuthPlainNoEnc, SMTPAuthType(" Plain-NoEnc ")) +} diff --git a/pkg/filemanager/fs/dbfs/manage.go b/pkg/filemanager/fs/dbfs/manage.go index 5e870cb9..43940787 100644 --- a/pkg/filemanager/fs/dbfs/manage.go +++ b/pkg/filemanager/fs/dbfs/manage.go @@ -848,29 +848,53 @@ func (f *DBFS) deleteFiles(ctx context.Context, targets map[Navigator][]*File, f defer reset() - // List all files to be deleted - toBeDeletedFiles := make([]*File, 0, len(files)) + // Walk the tree and delete in bounded batches: accumulating the + // whole tree before deleting materializes every file model + + // entity edges at once and OOMs on large trees (upstream #3574). + // Folders are deferred to a single delete after the walk — child + // listing resolves via HasParentWith, so parent rows must stay + // alive until their level has been fetched. + folderModels := make([]*ent.File, 0, 64) if err := n.Walk(ctx, files, intsets.MaxInt, intsets.MaxInt, func(targets []*File, level int) error { - toBeDeletedFiles = append(toBeDeletedFiles, targets...) indexToDelete = append(indexToDelete, lo.Map(targets, func(item *File, index int) int { return item.ID() })...) + + fileModels := make([]*ent.File, 0, len(targets)) + for _, item := range targets { + if item.Model.Type == int(types.FileTypeFolder) && !item.IsSymbolic() { + folderModels = append(folderModels, item.Model) + continue + } + fileModels = append(fileModels, item.Model) + } + + if len(fileModels) == 0 { + return nil + } + staleEntities, diff, err := fc.Delete(ctx, fileModels, opt) + if err != nil { + return fmt.Errorf("failed to delete files: %w", err) + } + storageDiff.Merge(diff) + allStaleEntities = append(allStaleEntities, lo.Map(staleEntities, func(item *ent.Entity, index int) fs.Entity { + return fs.NewEntity(item) + })...) return nil }); err != nil { return nil, nil, nil, fmt.Errorf("failed to walk files: %w", err) } - // Delete files - staleEntities, diff, err := fc.Delete(ctx, lo.Map(toBeDeletedFiles, func(item *File, index int) *ent.File { - return item.Model - }), opt) - if err != nil { - return nil, nil, nil, fmt.Errorf("failed to delete files: %w", err) + if len(folderModels) > 0 { + staleEntities, diff, err := fc.Delete(ctx, folderModels, opt) + if err != nil { + return nil, nil, nil, fmt.Errorf("failed to delete folders: %w", err) + } + storageDiff.Merge(diff) + allStaleEntities = append(allStaleEntities, lo.Map(staleEntities, func(item *ent.Entity, index int) fs.Entity { + return fs.NewEntity(item) + })...) } - storageDiff.Merge(diff) - allStaleEntities = append(allStaleEntities, lo.Map(staleEntities, func(item *ent.Entity, index int) fs.Entity { - return fs.NewEntity(item) - })...) } return allStaleEntities, storageDiff, indexToDelete, nil @@ -894,22 +918,12 @@ func (f *DBFS) copyFiles(ctx context.Context, targets map[Navigator][]*File, des newTargetsMap := make(map[int]*ent.File) storageDiff := make(inventory.StorageDiff) indexToCopy := make([]fs.IndexDiffCopyDetails, 0) - var diff inventory.StorageDiff for n, files := range targets { initialDstMap := make(map[int][]*ent.File) for _, file := range files { initialDstMap[file.Model.FileChildren] = dstAncestors } - firstLayer := true - // Let navigator use tx - reset, err := n.FollowTx(ctx) - if err != nil { - return nil, nil, nil, err - } - - defer reset() - if err := n.Walk(ctx, files, limit, intsets.MaxInt, func(targets []*File, level int) error { // check capacity for each file sizeTotal := int64(0) @@ -922,7 +936,7 @@ func (f *DBFS) copyFiles(ctx context.Context, targets map[Navigator][]*File, des } limit -= len(targets) - initialDstMap, diff, err = fc.Copy(ctx, &inventory.CopyParameter{ + newDstMap, diff, err := fc.Copy(ctx, &inventory.CopyParameter{ Files: lo.Map(targets, func(item *File, index int) *ent.File { return item.Model }), @@ -937,16 +951,29 @@ func (f *DBFS) copyFiles(ctx context.Context, targets map[Navigator][]*File, des return serializer.NewError(serializer.CodeDBError, "Failed to copy files", err) } - storageDiff.Merge(diff) - if firstLayer { - for k, v := range initialDstMap { + // Walk emits each level in bounded pages, so dst mappings must + // accumulate across callbacks — children of a folder copied in + // an earlier page may arrive in any later page. + for k, v := range newDstMap { + initialDstMap[k] = v + } + if level == 0 { + for k, v := range newDstMap { newTargetsMap[k] = v[0] } } + storageDiff.Merge(diff) + for _, file := range targets { if _, ok := file.Metadata()[FullTextIndexKey]; ok { - copiedFile := newTargetsMap[file.ID()] + // initialDstMap holds src->dst entries for every file + // copied so far, across all levels and pages. + copiedChain, ok := initialDstMap[file.ID()] + if !ok || len(copiedChain) == 0 { + continue + } + copiedFile := copiedChain[0] indexToCopy = append(indexToCopy, fs.IndexDiffCopyDetails{ OriginalFileID: file.ID(), FileID: copiedFile.ID, @@ -958,7 +985,6 @@ func (f *DBFS) copyFiles(ctx context.Context, targets map[Navigator][]*File, des } capacity.Used += sizeTotal - firstLayer = false return nil }); err != nil { diff --git a/pkg/filemanager/fs/dbfs/navigator.go b/pkg/filemanager/fs/dbfs/navigator.go index d98735e4..c2ac0be1 100644 --- a/pkg/filemanager/fs/dbfs/navigator.go +++ b/pkg/filemanager/fs/dbfs/navigator.go @@ -274,6 +274,12 @@ func (b *baseNavigator) children(ctx context.Context, parent *File, args *ListAr }, nil } +// walkFetchPageSize caps how many child rows a single GetChildFiles query +// may return during walk. Bounding the page (instead of fetching a whole +// level at once) keeps peak memory proportional to page size + folder +// count rather than tree size. +const walkFetchPageSize = 2000 + func (b *baseNavigator) walk(ctx context.Context, levelFiles []*File, limit, depth int, f WalkFunc) error { walked := 0 if len(levelFiles) == 0 { @@ -281,79 +287,91 @@ func (b *baseNavigator) walk(ctx context.Context, levelFiles []*File, limit, dep } owner := levelFiles[0].Owner() - level := 0 - for walked <= limit && depth >= 0 { - if len(levelFiles) == 0 { - break - } - - stop := false - depth-- - if len(levelFiles) > limit-walked { - levelFiles = levelFiles[:limit-walked] - stop = true - } - if err := f(levelFiles, level); err != nil { - return err - } - - if stop { - return ErrFileCountLimitedReached - } - - walked += len(levelFiles) - folders := lo.Filter(levelFiles, func(f *File, index int) bool { - return f.Model.Type == int(types.FileTypeFolder) && !f.IsSymbolic() - }) - if walked >= limit || len(folders) == 0 { - break - } + // Files still to emit at the current level. For level 0 this is the + // caller-provided slice; deeper levels are paged from the DB below. + pending := levelFiles + var parentMap map[int]*File + var parentModels []*ent.File + token := "" - levelFiles = levelFiles[:0] - leftCredit := limit - walked - parents := lo.SliceToMap(folders, func(file *File) (int, *File) { - return file.Model.ID, file - }) - for leftCredit > 0 { - token := "" - res, err := b.fileClient.GetChildFiles(ctx, - &inventory.ListFileParameters{ - PaginationArgs: &inventory.PaginationArgs{ - UseCursorPagination: true, - PageToken: token, - PageSize: leftCredit, + for depth >= 0 { + depth-- + folders := make([]*File, 0, 64) + + // Emit this level in bounded batches, tracking folder nodes for + // the next level. + for { + var batch []*File + if parentModels == nil { + if len(pending) == 0 { + break + } + remaining := limit - walked + if remaining <= 0 { + return ErrFileCountLimitedReached + } + batch = pending[:min(len(pending), remaining)] + pending = pending[len(batch):] + } else { + remaining := limit - walked + if remaining <= 0 { + return ErrFileCountLimitedReached + } + res, err := b.fileClient.GetChildFiles(ctx, + &inventory.ListFileParameters{ + PaginationArgs: &inventory.PaginationArgs{ + UseCursorPagination: true, + PageToken: token, + PageSize: min(remaining, walkFetchPageSize), + }, + MixedType: true, }, - MixedType: true, - }, - owner.ID, - lo.Map(folders, func(item *File, index int) *ent.File { - return item.Model - })...) - if err != nil { - return serializer.NewError(serializer.CodeDBError, "Failed to list children", err) + owner.ID, + parentModels...) + if err != nil { + return serializer.NewError(serializer.CodeDBError, "Failed to list children", err) + } + if len(res.Files) == 0 { + break + } + batch = lo.Map(res.Files, func(model *ent.File, index int) *File { + return newFile(parentMap[model.FileChildren], model) + }) + token = res.NextPageToken } - leftCredit -= len(res.Files) - - levelFiles = append(levelFiles, lo.Map(res.Files, func(model *ent.File, index int) *File { - p := parents[model.FileChildren] - return newFile(p, model) - })...) + if err := f(batch, level); err != nil { + return err + } + walked += len(batch) + for _, file := range batch { + if file.Model.Type == int(types.FileTypeFolder) && !file.IsSymbolic() { + folders = append(folders, file) + } + } - // All files listed - if res.NextPageToken == "" { + if parentModels == nil { + continue + } + if token == "" { break } + } - token = res.NextPageToken + if len(folders) == 0 { + return nil } - level++ - } - if walked >= limit { - return ErrFileCountLimitedReached + level++ + parentMap = lo.SliceToMap(folders, func(file *File) (int, *File) { + return file.Model.ID, file + }) + parentModels = lo.Map(folders, func(item *File, index int) *ent.File { + return item.Model + }) + token = "" } return nil diff --git a/pkg/filemanager/fs/dbfs/walk_test.go b/pkg/filemanager/fs/dbfs/walk_test.go new file mode 100644 index 00000000..412c2fa5 --- /dev/null +++ b/pkg/filemanager/fs/dbfs/walk_test.go @@ -0,0 +1,195 @@ +package dbfs + +import ( + "context" + "fmt" + "testing" + + "github.com/cloudreve/Cloudreve/v4/ent" + "github.com/cloudreve/Cloudreve/v4/ent/enttest" + entfile "github.com/cloudreve/Cloudreve/v4/ent/file" + entuser "github.com/cloudreve/Cloudreve/v4/ent/user" + "github.com/cloudreve/Cloudreve/v4/inventory" + "github.com/cloudreve/Cloudreve/v4/inventory/types" + "github.com/cloudreve/Cloudreve/v4/pkg/boolset" + "github.com/cloudreve/Cloudreve/v4/pkg/conf" + "github.com/cloudreve/Cloudreve/v4/pkg/hashid" + "github.com/stretchr/testify/require" +) + +// buildWalkTree creates root -> folder children. Each folder gets +// filesPerFolder file children. Returns the root model wrapped as *File. +func buildWalkTree(t *testing.T, client *ent.Client, folders, filesPerFolder int) (*ent.User, *File) { + ctx := context.Background() + group := client.Group.Create().SetName("walkers").SetPermissions(&boolset.BooleanSet{}).SaveX(ctx) + user := client.User.Create().SetEmail("walk@example.com").SetNick("walk").SetGroup(group).SaveX(ctx) + root := client.File.Create().SetName(inventory.RootFolderName).SetType(int(types.FileTypeFolder)).SetOwner(user).SaveX(ctx) + + folderModels := make([]*ent.File, 0, folders) + for i := 0; i < folders; i++ { + folderModels = append(folderModels, + client.File.Create().SetName(fmt.Sprintf("dir-%04d", i)).SetType(int(types.FileTypeFolder)).SetOwner(user).SetParent(root).SaveX(ctx)) + } + for i, folder := range folderModels { + creates := make([]*ent.FileCreate, 0, filesPerFolder) + for j := 0; j < filesPerFolder; j++ { + creates = append(creates, + client.File.Create().SetName(fmt.Sprintf("f-%04d-%04d", i, j)).SetType(int(types.FileTypeFile)).SetOwner(user).SetParent(folder)) + } + client.File.CreateBulk(creates...).SaveX(ctx) + } + + rootFile := newFile(nil, root) + rootFile.OwnerModel = user + return user, rootFile +} + +func walkTestNavigator(t *testing.T, client *ent.Client, user *ent.User) *baseNavigator { + t.Helper() + hasher, err := hashid.New("walk-test-salt") + require.NoError(t, err) + return newBaseNavigator(inventory.NewFileClient(client, conf.SQLiteDB, hasher), defaultFilter, user, nil, nil) +} + +func TestWalkEmitsWholeTree(t *testing.T) { + client := enttest.Open(t, "sqlite3", "file:"+t.Name()+"?mode=memory&cache=shared") + t.Cleanup(func() { require.NoError(t, client.Close()) }) + user, rootFile := buildWalkTree(t, client, 3, 5) + + nav := walkTestNavigator(t, client, user) + + emitted := make(map[int]int) // level -> count + seen := make(map[string]bool) // dedup guard + err := nav.walk(context.Background(), []*File{rootFile}, 1000, 10, func(files []*File, level int) error { + emitted[level] += len(files) + for _, f := range files { + key := fmt.Sprintf("%d:%d", level, f.ID()) + require.False(t, seen[key], "file %d emitted twice", f.ID()) + seen[key] = true + } + return nil + }) + require.NoError(t, err) + // level 0: root. level 1: 3 dirs. level 2: 15 files. + require.Equal(t, 1, emitted[0]) + require.Equal(t, 3, emitted[1]) + require.Equal(t, 15, emitted[2]) + require.Equal(t, 19, len(seen)) +} + +func TestWalkBatchesWideLevels(t *testing.T) { + client := enttest.Open(t, "sqlite3", "file:"+t.Name()+"?mode=memory&cache=shared") + t.Cleanup(func() { require.NoError(t, client.Close()) }) + // One folder with more children than walkFetchPageSize forces multiple + // callback invocations for the same level — the pre-fix implementation + // fetched and emitted the entire level at once. + user, rootFile := buildWalkTree(t, client, 1, walkFetchPageSize+500) + + nav := walkTestNavigator(t, client, user) + + levelCalls := make(map[int]int) + total := 0 + err := nav.walk(context.Background(), []*File{rootFile}, 10000, 10, func(files []*File, level int) error { + levelCalls[level]++ + total += len(files) + require.LessOrEqual(t, len(files), walkFetchPageSize, "callback batch exceeded page bound") + return nil + }) + require.NoError(t, err) + // root + 1 dir + (walkFetchPageSize+500) files + require.Equal(t, 2+walkFetchPageSize+500, total) + require.Greater(t, levelCalls[2], 1, "wide level must be emitted in multiple batches") +} + +func TestWalkLimit(t *testing.T) { + client := enttest.Open(t, "sqlite3", "file:"+t.Name()+"?mode=memory&cache=shared") + t.Cleanup(func() { require.NoError(t, client.Close()) }) + user, rootFile := buildWalkTree(t, client, 2, 10) + + nav := walkTestNavigator(t, client, user) + + emitted := 0 + err := nav.walk(context.Background(), []*File{rootFile}, 5, 10, func(files []*File, level int) error { + emitted += len(files) + return nil + }) + require.ErrorIs(t, err, ErrFileCountLimitedReached) + require.LessOrEqual(t, emitted, 5) +} + +// TestDeleteFilesBatched walks a nested tree inside a tx and asserts the +// batched delete removes every row — including grandchildren of folders +// whose deletion is deferred until after the walk (HasParentWith relies +// on parent rows staying alive mid-walk). +func TestDeleteFilesBatched(t *testing.T) { + client := enttest.Open(t, "sqlite3", "file:"+t.Name()+"?mode=memory&cache=shared") + t.Cleanup(func() { require.NoError(t, client.Close()) }) + ctx := context.Background() + + group := client.Group.Create().SetName("deleters").SetPermissions(&boolset.BooleanSet{}).SaveX(ctx) + user := client.User.Create().SetEmail("del@example.com").SetNick("del").SetGroup(group).SaveX(ctx) + policy := client.StoragePolicy.Create().SetName("local").SetType("local").SaveX(ctx) + root := client.File.Create().SetName(inventory.RootFolderName).SetType(int(types.FileTypeFolder)).SetOwner(user).SaveX(ctx) + + // dir -> sub -> leaf; dir also holds two files. + dir := client.File.Create().SetName("dir").SetType(int(types.FileTypeFolder)).SetOwner(user).SetParent(root).SaveX(ctx) + sub := client.File.Create().SetName("sub").SetType(int(types.FileTypeFolder)).SetOwner(user).SetParent(dir).SaveX(ctx) + leaf := client.File.Create().SetName("leaf.txt").SetType(int(types.FileTypeFile)).SetOwner(user).SetParent(sub).SaveX(ctx) + f1 := client.File.Create().SetName("a.txt").SetType(int(types.FileTypeFile)).SetOwner(user).SetParent(dir).SaveX(ctx) + f2 := client.File.Create().SetName("b.txt").SetType(int(types.FileTypeFile)).SetOwner(user).SetParent(dir).SaveX(ctx) + + for _, fm := range []*ent.File{leaf, f1, f2} { + client.Entity.Create(). + SetType(1).SetSource("src").SetSize(10). + SetStoragePolicyEntities(policy.ID). + AddFileIDs(fm.ID). + SaveX(ctx) + } + + fc := inventory.NewFileClient(client, conf.SQLiteDB, nil) + uc := inventory.NewUserClient(client) + txFc, tx, txCtx, err := inventory.WithTx(ctx, fc) + require.NoError(t, err) + txCtx = context.WithValue(txCtx, inventory.LoadFileEntity{}, true) + + // Targets: dir wrapped as *File with entities edge loaded. + dirModel := client.File.Query().WithEntities().Where(entfile.IDEQ(dir.ID)).OnlyX(txCtx) + dirFile := newFile(nil, dirModel) + dirFile.OwnerModel = user + + user = client.User.Query().WithGroup().Where(entuser.IDEQ(user.ID)).OnlyX(txCtx) + f := &DBFS{user: user} + targets := map[Navigator][]*File{ + &myNavigator{baseNavigator: newBaseNavigator(txFc, defaultFilter, user, nil, nil), user: user, fileClient: txFc, userClient: uc}: {dirFile}, + } + + stale, diff, indexToDelete, err := f.deleteFiles(txCtx, targets, txFc, nil) + require.NoError(t, err) + require.NoError(t, inventory.Commit(tx)) + + // dir, sub, leaf, f1, f2 all deleted; root survives. + remaining := client.File.Query().AllX(ctx) + require.Len(t, remaining, 1) + require.Equal(t, root.ID, remaining[0].ID) + require.ElementsMatch(t, []int{dir.ID, sub.ID, leaf.ID, f1.ID, f2.ID}, indexToDelete) + require.Len(t, stale, 3) + require.Equal(t, int64(-30), diff[user.ID]) +} + +func TestWalkDepthLimit(t *testing.T) { + client := enttest.Open(t, "sqlite3", "file:"+t.Name()+"?mode=memory&cache=shared") + t.Cleanup(func() { require.NoError(t, client.Close()) }) + user, rootFile := buildWalkTree(t, client, 2, 4) + + nav := walkTestNavigator(t, client, user) + + maxLevel := -1 + err := nav.walk(context.Background(), []*File{rootFile}, 1000, 1, func(files []*File, level int) error { + if level > maxLevel { + maxLevel = level + } + return nil + }) + require.NoError(t, err) + require.Equal(t, 1, maxLevel) +} diff --git a/pkg/filemanager/fs/fs.go b/pkg/filemanager/fs/fs.go index 501aee31..c27a340e 100644 --- a/pkg/filemanager/fs/fs.go +++ b/pkg/filemanager/fs/fs.go @@ -282,6 +282,10 @@ type ( // uploaded. Used to safely trigger CompleteUpload only after every // chunk has been received when the client uploads chunks concurrently. ChunksReceived map[int]struct{} + // RangesReceived records merged byte intervals [start,end) written for + // arbitrary-range uploads (e.g. WebDAV Content-Range PUTs). Used to + // trigger CompleteUpload only after the whole file is covered. + RangesReceived [][2]int64 } // UploadProps properties of an upload session/request. diff --git a/pkg/filemanager/manager/upload.go b/pkg/filemanager/manager/upload.go index f42d2665..a9b29b9c 100644 --- a/pkg/filemanager/manager/upload.go +++ b/pkg/filemanager/manager/upload.go @@ -48,6 +48,12 @@ type ( // uploads chunks concurrently and the last-indexed chunk arrives before some // earlier chunks are still in flight. MarkChunkUploaded(ctx context.Context, session *fs.UploadSession, chunkIndex int) (allReceived bool, err error) + // MarkRangeUploaded atomically records the given byte range [offset, offset+length) + // as received on the shared upload session and reports whether the whole file + // is covered. Callers must invoke CompleteUpload only when allReceived is true. + // It is used for arbitrary-range uploads (e.g. WebDAV Content-Range PUTs) where + // chunks are addressed by byte offset rather than a fixed chunk index. + MarkRangeUploaded(ctx context.Context, session *fs.UploadSession, offset, length int64) (allReceived bool, err error) } ) @@ -310,6 +316,77 @@ func (m *manager) MarkChunkUploaded(ctx context.Context, session *fs.UploadSessi return false, nil } +// mergeByteRanges inserts [start,end) into the sorted interval list and +// coalesces overlapping or adjacent intervals. +func mergeByteRanges(ranges [][2]int64, start, end int64) [][2]int64 { + res := make([][2]int64, 0, len(ranges)+1) + inserted := false + for _, r := range ranges { + if r[1] < start { + res = append(res, r) + continue + } + if r[0] > end { + if !inserted { + res = append(res, [2]int64{start, end}) + inserted = true + } + res = append(res, r) + continue + } + // Overlapping or adjacent — extend the pending interval. + if r[0] < start { + start = r[0] + } + if r[1] > end { + end = r[1] + } + } + if !inserted { + res = append(res, [2]int64{start, end}) + } + return res +} + +// rangesCoverFull reports whether merged intervals cover [0,total). +func rangesCoverFull(ranges [][2]int64, total int64) bool { + return len(ranges) == 1 && ranges[0][0] <= 0 && ranges[0][1] >= total +} + +func (m *manager) MarkRangeUploaded(ctx context.Context, session *fs.UploadSession, offset, length int64) (bool, error) { + if session == nil || session.Props == nil || length < 0 || offset < 0 { + return false, fmt.Errorf("invalid upload session or range") + } + + sessionID := session.Props.UploadSessionID + mu := lockUploadSession(sessionID) + mu.Lock() + defer mu.Unlock() + + raw, ok := m.kv.Get(UploadSessionCachePrefix + sessionID) + if !ok { + // Session already completed or cancelled elsewhere. + return false, nil + } + fresh, ok := raw.(fs.UploadSession) + if !ok { + return false, fmt.Errorf("unexpected upload session type in KV") + } + + fresh.RangesReceived = mergeByteRanges(fresh.RangesReceived, offset, offset+length) + + ttl := max(1, int(time.Until(fresh.Props.ExpireAt).Seconds())) + if err := m.kv.Set(UploadSessionCachePrefix+sessionID, fresh, ttl); err != nil { + return false, fmt.Errorf("failed to persist upload session progress: %w", err) + } + + if rangesCoverFull(fresh.RangesReceived, fresh.Props.Size) { + session.RangesReceived = fresh.RangesReceived + return true, nil + } + return false, nil +} + func (m *manager) CancelUploadSession(ctx context.Context, path *fs.URI, sessionID string) error { // Get upload session var session *fs.UploadSession diff --git a/pkg/filemanager/manager/upload_test.go b/pkg/filemanager/manager/upload_test.go new file mode 100644 index 00000000..9d2fb040 --- /dev/null +++ b/pkg/filemanager/manager/upload_test.go @@ -0,0 +1,139 @@ +package manager + +import ( + "context" + "testing" + "time" + + "github.com/cloudreve/Cloudreve/v4/pkg/cache" + "github.com/cloudreve/Cloudreve/v4/pkg/filemanager/fs" +) + +func TestMergeByteRanges(t *testing.T) { + tests := []struct { + name string + ranges [][2]int64 + start int64 + end int64 + want [][2]int64 + }{ + {"empty", nil, 0, 10, [][2]int64{{0, 10}}}, + {"append gap", [][2]int64{{0, 10}}, 20, 30, [][2]int64{{0, 10}, {20, 30}}}, + {"prepend gap", [][2]int64{{20, 30}}, 0, 10, [][2]int64{{0, 10}, {20, 30}}}, + {"adjacent merge", [][2]int64{{0, 10}}, 10, 20, [][2]int64{{0, 20}}}, + {"bridge two", [][2]int64{{0, 10}, {20, 30}}, 10, 20, [][2]int64{{0, 30}}}, + {"contained", [][2]int64{{0, 30}}, 5, 15, [][2]int64{{0, 30}}}, + {"extend right", [][2]int64{{0, 10}}, 5, 20, [][2]int64{{0, 20}}}, + {"extend left", [][2]int64{{10, 20}}, 0, 15, [][2]int64{{0, 20}}}, + {"middle gap", [][2]int64{{0, 5}, {30, 40}}, 10, 20, [][2]int64{{0, 5}, {10, 20}, {30, 40}}}, + {"merge all", [][2]int64{{0, 5}, {10, 15}, {20, 25}}, 4, 21, [][2]int64{{0, 25}}}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := mergeByteRanges(tt.ranges, tt.start, tt.end) + if len(got) != len(tt.want) { + t.Fatalf("got %v, want %v", got, tt.want) + } + for i := range got { + if got[i] != tt.want[i] { + t.Fatalf("got %v, want %v", got, tt.want) + } + } + }) + } +} + +func TestRangesCoverFull(t *testing.T) { + if !rangesCoverFull([][2]int64{{0, 100}}, 100) { + t.Fatal("expected full coverage") + } + if rangesCoverFull([][2]int64{{0, 99}}, 100) { + t.Fatal("partial coverage reported as full") + } + if rangesCoverFull([][2]int64{{0, 50}, {50, 100}}, 100) { + t.Fatal("unmerged intervals should not report full coverage") + } + if rangesCoverFull(nil, 100) { + t.Fatal("empty intervals reported as full") + } +} + +func newRangedSession(id string, size int64) fs.UploadSession { + return fs.UploadSession{ + Props: &fs.UploadProps{ + UploadSessionID: id, + Size: size, + ExpireAt: time.Now().Add(time.Hour), + }, + } +} + +func TestMarkRangeUploaded(t *testing.T) { + ctx := context.Background() + m := &manager{kv: cache.NewMemoStore("", nil)} + session := newRangedSession("test-range-session", 100) + if err := m.kv.Set(UploadSessionCachePrefix+"test-range-session", session, 60); err != nil { + t.Fatal(err) + } + + // Out-of-order arrival: second half first. + all, err := m.MarkRangeUploaded(ctx, &session, 50, 50) + if err != nil { + t.Fatal(err) + } + if all { + t.Fatal("reported complete after only second half") + } + + all, err = m.MarkRangeUploaded(ctx, &session, 0, 50) + if err != nil { + t.Fatal(err) + } + if !all { + t.Fatal("expected complete after full coverage") + } + + // Session record carries merged coverage. + raw, ok := m.kv.Get(UploadSessionCachePrefix + "test-range-session") + if !ok { + t.Fatal("session missing from KV") + } + stored := raw.(fs.UploadSession) + if len(stored.RangesReceived) != 1 || stored.RangesReceived[0] != [2]int64{0, 100} { + t.Fatalf("unexpected ranges: %v", stored.RangesReceived) + } +} + +func TestMarkRangeUploadedGap(t *testing.T) { + ctx := context.Background() + m := &manager{kv: cache.NewMemoStore("", nil)} + session := newRangedSession("test-gap-session", 100) + if err := m.kv.Set(UploadSessionCachePrefix+"test-gap-session", session, 60); err != nil { + t.Fatal(err) + } + + for _, r := range [][2]int64{{0, 40}, {60, 100}} { + all, err := m.MarkRangeUploaded(ctx, &session, r[0], r[1]-r[0]) + if err != nil { + t.Fatal(err) + } + if all { + t.Fatal("reported complete with a gap in coverage") + } + } +} + +func TestMarkRangeUploadedMissingSession(t *testing.T) { + ctx := context.Background() + m := &manager{kv: cache.NewMemoStore("", nil)} + session := newRangedSession("gone", 100) + + all, err := m.MarkRangeUploaded(ctx, &session, 0, 100) + if err != nil { + t.Fatal(err) + } + if all { + t.Fatal("missing session reported complete") + } +} diff --git a/pkg/request/ssrftest/ssrf_test.go b/pkg/request/ssrftest/ssrf_test.go index 995f5dd4..c23957c7 100644 --- a/pkg/request/ssrftest/ssrf_test.go +++ b/pkg/request/ssrftest/ssrf_test.go @@ -59,6 +59,16 @@ func TestValidateExternalURL_IPLiterals(t *testing.T) { "http://100.64.0.1/cgnat", "http://[::ffff:127.0.0.1]/v4mapped", "http://[::ffff:10.0.0.1]/v4mappedpriv", + // IPv4-in-IPv6 transition forms wrapping internal targets. + "http://[64:ff9b::127.0.0.1]/nat64-loopback", + "http://[64:ff9b::a9fe:a9fe]/nat64-metadata", + "http://[64:ff9b::169.254.169.254]/nat64-metadata-dec", + "http://[64:ff9b::a00:1]/nat64-private", + "http://[2002:7f00:0001::]/6to4-loopback", + "http://[2002:a9fe:a9fe::]/6to4-metadata", + "http://[2002:0a00:0001::]/6to4-private", + "http://[2001:0000:4136:e378:8000:63bf:f5ff:fffe]/teredo-10.0.0.1", + "http://[::127.0.0.1]/v4compatible", "http://224.0.0.1/multicast", } for _, raw := range cases { @@ -73,6 +83,9 @@ func TestValidateExternalURL_PublicIP(t *testing.T) { "http://1.1.1.1/", "https://8.8.8.8/", "http://[2606:4700:4700::1111]/", + // Transition forms wrapping *public* IPv4 remain allowed. + "http://[64:ff9b::808:808]/nat64-public", + "http://[2002:0808:0808::]/6to4-public", } for _, raw := range cases { err := request.ValidateExternalURL(ctx, raw, request.SSRFOptions{}) diff --git a/pkg/serializer/error.go b/pkg/serializer/error.go index 28666d33..bb030a43 100644 --- a/pkg/serializer/error.go +++ b/pkg/serializer/error.go @@ -257,6 +257,8 @@ const ( CodeAnonymouseAccessDenied = 40088 // CodeInsufficientScope OAuth token scope insufficient CodeInsufficientScope = 40089 + // CodeRateLimited 请求频率超限 + CodeRateLimited = 40090 // CodeDBError 数据库操作失败 CodeDBError = 50001 // CodeEncryptError 加密失败 diff --git a/pkg/setting/provider.go b/pkg/setting/provider.go index 393927e9..d0878a94 100644 --- a/pkg/setting/provider.go +++ b/pkg/setting/provider.go @@ -807,6 +807,7 @@ func (s *settingProvider) SMTP(ctx context.Context) *SMTP { ForceEncryption: s.getBoolean(ctx, "smtpEncryption", false), Port: s.getInt(ctx, "smtpPort", 25), Keepalive: s.getInt(ctx, "mail_keepalive", 30), + AuthType: s.getString(ctx, "smtp_auth", "autodiscover"), } } diff --git a/pkg/setting/types.go b/pkg/setting/types.go index 339e533e..67d3aece 100644 --- a/pkg/setting/types.go +++ b/pkg/setting/types.go @@ -65,6 +65,7 @@ type SMTP struct { ForceEncryption bool Port int Keepalive int + AuthType string } type TokenAuth struct { diff --git a/pkg/util/common.go b/pkg/util/common.go index 1b5aa0c4..3791d7de 100644 --- a/pkg/util/common.go +++ b/pkg/util/common.go @@ -37,24 +37,24 @@ func RandStringRunes(n int) string { } func RandStringRunesCrypto(n int) string { - b := make([]rune, n) - for i := range b { - num, err := cryptoRand.Int(cryptoRand.Reader, big.NewInt(int64(len(RandomVariantAll)))) - if err != nil { - // fallback to math/rand on crypto failure - b[i] = RandomVariantAll[rand.Intn(len(RandomVariantAll))] - } else { - b[i] = RandomVariantAll[num.Int64()] - } - } - return string(b) + return randStringCrypto(n, RandomVariantAll) } // RandString returns random string in given length and variant func RandString(n int, variant []rune) string { + return randStringCrypto(n, variant) +} + +func randStringCrypto(n int, variant []rune) string { b := make([]rune, n) for i := range b { - b[i] = variant[rand.Intn(len(variant))] + num, err := cryptoRand.Int(cryptoRand.Reader, big.NewInt(int64(len(variant)))) + if err != nil { + // fallback to math/rand on crypto failure + b[i] = variant[rand.Intn(len(variant))] + } else { + b[i] = variant[num.Int64()] + } } return string(b) } diff --git a/pkg/webdav/webdav.go b/pkg/webdav/webdav.go index 744a7248..78ec8027 100644 --- a/pkg/webdav/webdav.go +++ b/pkg/webdav/webdav.go @@ -7,11 +7,13 @@ package webdav // import "golang.org/x/net/webdav" import ( "context" + "crypto/sha1" "errors" "fmt" "net/http" "net/url" "path" + "strconv" "strings" "time" @@ -256,6 +258,34 @@ func handlePut(c *gin.Context, user *ent.User, fm manager.FileManager) (status i return http.StatusBadRequest, err } + // A PUT with no length information at all (e.g. chunked transfer encoding) + // cannot be sized — previously this silently created an empty file. + if fileSize == 0 && c.Request.ContentLength < 0 && + c.Request.Header.Get("X-Expected-Entity-Length") == "" { + return http.StatusLengthRequired, nil + } + + // Ranged PUTs ("Content-Range: bytes start-end/total") are used by some + // clients (e.g. Mountain Duck) to upload large files in pieces. A partial + // range goes through the chunked-assembly path; a full-range PUT falls + // through to the regular overwrite path. + contentRange, err := parseContentRange(c.Request.Header.Get("Content-Range")) + if err != nil { + return http.StatusBadRequest, err + } + + m := manager.NewFileManager(dependency.FromContext(ctx), user) + defer m.Recycle() + + if contentRange != nil { + if contentRange.start != 0 || contentRange.end+1 != contentRange.total { + return handleRangedPut(ctx, c, user, m, fm, uri, rc, fileSize, contentRange) + } + if contentRange.total != fileSize { + return http.StatusBadRequest, errInvalidContentRange + } + } + fileData := &fs.UploadRequest{ Props: &fs.UploadProps{ Uri: uri, @@ -266,9 +296,6 @@ func handlePut(c *gin.Context, user *ent.User, fm manager.FileManager) (status i Mode: fs.ModeOverwrite, } - m := manager.NewFileManager(dependency.FromContext(ctx), user) - defer m.Recycle() - // Update file res, err := m.Update(ctx, fileData) if err != nil { @@ -284,6 +311,150 @@ func handlePut(c *gin.Context, user *ent.User, fm manager.FileManager) (status i return http.StatusCreated, nil } +var errInvalidContentRange = errors.New("invalid Content-Range") + +// contentRange describes a "bytes start-end/total" request range, with end +// inclusive per RFC 7233. +type contentRange struct { + start, end, total int64 +} + +// parseContentRange parses a Content-Range header of the form +// "bytes start-end/total". It returns (nil, nil) when the header is absent. +func parseContentRange(h string) (*contentRange, error) { + if h == "" { + return nil, nil + } + + h = strings.TrimSpace(h) + if !strings.HasPrefix(h, "bytes ") { + return nil, errInvalidContentRange + } + + rangePart, totalPart, ok := strings.Cut(h[len("bytes "):], "/") + if !ok || totalPart == "*" || totalPart == "" { + // An unknown total cannot be turned into a sized upload session. + return nil, errInvalidContentRange + } + + startPart, endPart, ok := strings.Cut(rangePart, "-") + if !ok { + return nil, errInvalidContentRange + } + + start, err := strconv.ParseInt(strings.TrimSpace(startPart), 10, 64) + if err != nil { + return nil, errInvalidContentRange + } + end, err := strconv.ParseInt(strings.TrimSpace(endPart), 10, 64) + if err != nil { + return nil, errInvalidContentRange + } + total, err := strconv.ParseInt(strings.TrimSpace(totalPart), 10, 64) + if err != nil { + return nil, errInvalidContentRange + } + + if start < 0 || end < start || total <= 0 || end >= total { + return nil, errInvalidContentRange + } + return &contentRange{start: start, end: end, total: total}, nil +} + +// handleRangedPut assembles a multi-request ranged PUT into a single upload +// session. Byte-range coverage is tracked on the session in KV; the upload is +// completed once [0,total) has been received. Only local storage policies are +// supported, as remote drivers cannot honor arbitrary write offsets. +func handleRangedPut(ctx context.Context, c *gin.Context, user *ent.User, m manager.FileManager, fm manager.FileManager, uri *fs.URI, rc request.LimitReaderCloser, fileSize int64, cr *contentRange) (status int, err error) { + if fileSize != cr.end-cr.start+1 { + return http.StatusBadRequest, errInvalidContentRange + } + + dep := dependency.FromContext(c) + kv := dep.KV() + + // One in-flight ranged upload per user+path, keyed deterministically so + // that subsequent chunk requests resume the same upload session. + sessionKey := fmt.Sprintf("dav-put-%d-%x", user.ID, sha1.Sum([]byte(uri.String()))) + + var session *fs.UploadSession + if raw, ok := kv.Get(manager.UploadSessionCachePrefix + sessionKey); ok { + s, ok := raw.(fs.UploadSession) + if !ok || s.Props == nil { + kv.Delete(manager.UploadSessionCachePrefix, sessionKey) + } else if s.Props.Size == cr.total { + session = &s + } else { + // A different upload to the same path — fail the stale session. + m.OnUploadFailed(ctx, &s) + session = nil + } + } + + if session == nil { + ttl := dep.SettingProvider().UploadSessionTTL(ctx) + if _, err := m.CreateUploadSession(ctx, &fs.UploadRequest{ + Props: &fs.UploadProps{ + Uri: uri, + Size: cr.total, + UploadSessionID: sessionKey, + ExpireAt: time.Now().Add(ttl), + }, + Mode: fs.ModeOverwrite, + }); err != nil { + return purposeStatusCodeFromError(err), err + } + + raw, ok := kv.Get(manager.UploadSessionCachePrefix + sessionKey) + if !ok { + return http.StatusInternalServerError, errors.New("upload session not persisted") + } + s := raw.(fs.UploadSession) + session = &s + + // Only the local driver honors arbitrary write offsets; for other + // policies a ranged PUT cannot be assembled safely. + if session.Policy == nil || session.Policy.Type != types.PolicyTypeLocal { + m.OnUploadFailed(ctx, session) + return http.StatusNotImplemented, errors.New("ranged PUT not supported by this storage policy") + } + } + + chunkReq := &fs.UploadRequest{ + File: rc, + Offset: cr.start, + Props: session.Props.Copy(), + Mode: fs.ModeOverwrite, + } + if err := m.Upload(ctx, chunkReq, session.Policy, session); err != nil { + return purposeStatusCodeFromError(err), err + } + + if lrc, ok := chunkReq.File.(request.LimitReaderCloser); ok && lrc.Count() != fileSize { + return http.StatusInternalServerError, fmt.Errorf("uploaded data(%d) does not match purposed size(%d)", lrc.Count(), fileSize) + } + + allReceived, err := m.MarkRangeUploaded(ctx, session, cr.start, fileSize) + if err != nil { + return http.StatusInternalServerError, err + } + if !allReceived { + return http.StatusCreated, nil + } + + res, err := m.CompleteUpload(ctx, session) + if err != nil { + return purposeStatusCodeFromError(err), err + } + + etag, err := findETag(ctx, fm, res) + if err != nil { + return http.StatusInternalServerError, err + } + c.Writer.Header().Set("ETag", etag) + return http.StatusCreated, nil +} + func handleOptions(c *gin.Context, user *ent.User, fm manager.FileManager) (status int, err error) { allow := []string{"OPTIONS", "LOCK", "PUT", "MKCOL"} diff --git a/pkg/webdav/webdav_test.go b/pkg/webdav/webdav_test.go index 5a644560..8c3682c6 100644 --- a/pkg/webdav/webdav_test.go +++ b/pkg/webdav/webdav_test.go @@ -111,3 +111,48 @@ func mustWebDAVTestURI(t *testing.T, raw string) *fs.URI { } return uri } + +func TestParseContentRange(t *testing.T) { + tests := []struct { + header string + want *contentRange + err bool + }{ + {"", nil, false}, + {"bytes 0-1048575/3145728", &contentRange{0, 1048575, 3145728}, false}, + {"bytes 1048576-2097151/3145728", &contentRange{1048576, 2097151, 3145728}, false}, + {"bytes 0-99/100", &contentRange{0, 99, 100}, false}, + {" bytes 0-9/10 ", &contentRange{0, 9, 10}, false}, + {"bytes 0-9/*", nil, true}, + {"items 0-9/10", nil, true}, + {"bytes 0-9", nil, true}, + {"bytes */10", nil, true}, + {"bytes 9-0/10", nil, true}, + {"bytes 0-10/10", nil, true}, + {"bytes -1-9/10", nil, true}, + {"bytes 0-9/0", nil, true}, + {"bytes a-b/c", nil, true}, + } + + for _, tt := range tests { + got, err := parseContentRange(tt.header) + if tt.err { + if err == nil { + t.Fatalf("header %q: expected error, got %+v", tt.header, got) + } + continue + } + if err != nil { + t.Fatalf("header %q: unexpected error: %v", tt.header, err) + } + if tt.want == nil { + if got != nil { + t.Fatalf("header %q: expected nil, got %+v", tt.header, got) + } + continue + } + if got == nil || *got != *tt.want { + t.Fatalf("header %q: got %+v, want %+v", tt.header, got, tt.want) + } + } +} diff --git a/routers/router.go b/routers/router.go index 24ad3606..0b124e06 100644 --- a/routers/router.go +++ b/routers/router.go @@ -2,6 +2,7 @@ package routers import ( "net/http" + "time" "github.com/cloudreve/Cloudreve/v4/application/constants" "github.com/cloudreve/Cloudreve/v4/application/dependency" @@ -289,6 +290,7 @@ func initMasterRouter(dep dependency.Dep) *gin.Engine { { // 用户登录 token.POST("", + middleware.RateLimitByIP("login", 10, time.Minute), middleware.CaptchaRequired(func(c *gin.Context) bool { return dep.SettingProvider().LoginCaptchaEnabled(c) }), @@ -298,11 +300,13 @@ func initMasterRouter(dep dependency.Dep) *gin.Engine { ) // 2-factor authentication token.POST("2fa", + middleware.RateLimitByIP("login_2fa", 10, time.Minute), controllers.FromJSON[usersvc.OtpValidationService](usersvc.OtpValidationParameterCtx{}), controllers.UserLogin2FAValidation, controllers.UserIssueToken, ) token.POST("refresh", + middleware.RateLimitByIP("token_refresh", 30, time.Minute), middleware.RequiredScopes(types.ScopeOfflineAccess), controllers.FromJSON[usersvc.RefreshTokenService](usersvc.RefreshTokenParameterCtx{}), controllers.UserRefreshToken, @@ -331,6 +335,7 @@ func initMasterRouter(dep dependency.Dep) *gin.Engine { controllers.UserSSOCallback, ) ssoRouter.POST("exchange", + middleware.RateLimitByIP("sso_exchange", 20, time.Minute), controllers.FromJSON[usersvc.SSOExchangeService](usersvc.SSOExchangeParameterCtx{}), controllers.UserSSOExchange, ) @@ -348,6 +353,7 @@ func initMasterRouter(dep dependency.Dep) *gin.Engine { controllers.GrantAppConsent, ) oauthRouter.POST("token", + middleware.RateLimitByIP("oauth_token", 20, time.Minute), controllers.FromForm[oauth.ExchangeTokenService](oauth.ExchangeTokenParamCtx{}), controllers.ExchangeToken, ) @@ -369,6 +375,7 @@ func initMasterRouter(dep dependency.Dep) *gin.Engine { { // WebAuthn login prepare authn.PUT("", + middleware.RateLimitByIP("authn", 20, time.Minute), middleware.IsFunctionEnabled(func(c *gin.Context) bool { return dep.SettingProvider().AuthnEnabled(c) }), @@ -376,6 +383,7 @@ func initMasterRouter(dep dependency.Dep) *gin.Engine { ) // WebAuthn finish login authn.POST("", + middleware.RateLimitByIP("authn", 20, time.Minute), middleware.IsFunctionEnabled(func(c *gin.Context) bool { return dep.SettingProvider().AuthnEnabled(c) }), @@ -391,6 +399,7 @@ func initMasterRouter(dep dependency.Dep) *gin.Engine { { // 用户注册 Done user.POST("", + middleware.RateLimitByIP("register", 5, time.Minute), middleware.IsFunctionEnabled(func(c *gin.Context) bool { return dep.SettingProvider().RegisterEnabled(c) }), @@ -402,12 +411,14 @@ func initMasterRouter(dep dependency.Dep) *gin.Engine { ) // 通过邮件里的链接重设密码 user.PATCH("reset/:id", + middleware.RateLimitByIP("reset_apply", 10, time.Minute), middleware.HashID(hashid.UserID), controllers.FromJSON[usersvc.UserResetService](usersvc.UserResetParameterCtx{}), controllers.UserReset, ) // 发送密码重设邮件 user.POST("reset", + middleware.RateLimitByIP("reset_mail", 5, time.Minute), middleware.CaptchaRequired(func(c *gin.Context) bool { return dep.SettingProvider().ForgotPasswordCaptchaEnabled(c) }), diff --git a/service/admin/tools.go b/service/admin/tools.go index ac701b44..be76a999 100644 --- a/service/admin/tools.go +++ b/service/admin/tools.go @@ -10,6 +10,7 @@ import ( "github.com/cloudreve/Cloudreve/v4/application/dependency" "github.com/cloudreve/Cloudreve/v4/pkg/boolset" + "github.com/cloudreve/Cloudreve/v4/pkg/email" "github.com/cloudreve/Cloudreve/v4/pkg/filemanager/manager" request2 "github.com/cloudreve/Cloudreve/v4/pkg/request" "github.com/cloudreve/Cloudreve/v4/pkg/serializer" @@ -142,7 +143,7 @@ func (s *TestSMTPService) Test(c *gin.Context) error { opts := []mail.Option{ mail.WithPort(port), - mail.WithSMTPAuth(mail.SMTPAuthAutoDiscover), mail.WithTLSPortPolicy(mail.TLSOpportunistic), + mail.WithSMTPAuth(email.SMTPAuthType(s.Settings["smtp_auth"])), mail.WithTLSPortPolicy(mail.TLSOpportunistic), mail.WithUsername(s.Settings["smtpUser"]), mail.WithPassword(s.Settings["smtpPass"]), } if setting.IsTrueValue(s.Settings["smtpEncryption"]) {