Merge pull request #142 from Dvorinka/fix/phase-c-bugs

Phase C bug fixes: SMTP auth selection, trash-walk OOM, WebDAV ranged PUT
pull/3582/head
Tomáš Dvořák 2 weeks ago committed by GitHub
commit a82c9f5277
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194

1
.gitignore vendored

@ -42,3 +42,4 @@ cloudreve
.playwright-mcp/ .playwright-mcp/
# Monorepo staging dir for asset packaging (transient) # Monorepo staging dir for asset packaging (transient)
/assets/ /assets/
pkg/util/test/

@ -136,7 +136,8 @@ Order = user-visible value first; each ships with backend + UI + tests.
## 5. Phase C — security + quality ## 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 - 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 - `desloppify` + `security-reviewer` passes; scorecard appended to README
## 6. Phase D — desktop, all platforms ## 6. Phase D — desktop, all platforms

@ -437,6 +437,18 @@
"replyToAddressDes": "The mailbox used to receive reply emails when users reply to emails sent by the system.", "replyToAddressDes": "The mailbox used to receive reply emails when users reply to emails sent by the system.",
"enforceSSL": "Enforce SSL connection", "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.", "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)", "smtpTTL": "SMTP connection TTL (seconds)",
"smtpTTLDes": "SMTP connections established during the TTL period will be reused by new mail delivery requests.", "smtpTTLDes": "SMTP connections established during the TTL period will be reused by new mail delivery requests.",
"emailTemplates": "Email Templates", "emailTemplates": "Email Templates",

@ -437,6 +437,18 @@
"replyToAddressDes": "用户回复系统发送的邮件时,用于接收回信的邮箱。", "replyToAddressDes": "用户回复系统发送的邮件时,用于接收回信的邮箱。",
"enforceSSL": "强制使用 SSL 连接", "enforceSSL": "强制使用 SSL 连接",
"enforceSSLDes": "是否强制使用 SSL 加密连接。如果无法发送邮件,可关闭此项,Cloudreve 会尝试使用 STARTTLS 并决定是否使用加密连接。", "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 连接有效期 (秒)", "smtpTTL": "SMTP 连接有效期 (秒)",
"smtpTTLDes": "有效期内建立的 SMTP 连接会被新邮件发送请求复用。", "smtpTTLDes": "有效期内建立的 SMTP 连接会被新邮件发送请求复用。",
"emailTemplates": "邮件模板", "emailTemplates": "邮件模板",

@ -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 { useSnackbar } from "notistack";
import { useContext, useState } from "react"; import { useContext, useState } from "react";
import { useTranslation } from "react-i18next"; import { useTranslation } from "react-i18next";
@ -6,9 +6,10 @@ import { sendTestSMTP } from "../../../../api/api.ts";
import { useAppDispatch } from "../../../../redux/hooks.ts"; import { useAppDispatch } from "../../../../redux/hooks.ts";
import { isTrueVal } from "../../../../session/utils.ts"; import { isTrueVal } from "../../../../session/utils.ts";
import { DefaultCloseAction } from "../../../Common/Snackbar/snackbar.tsx"; 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 DraggableDialog, { StyledDialogContentText } from "../../../Dialogs/DraggableDialog.tsx";
import MailOutlined from "../../../Icons/MailOutlined.tsx"; import MailOutlined from "../../../Icons/MailOutlined.tsx";
import { SquareMenuItem } from "../../../FileManager/ContextMenu/ContextMenu.tsx";
import SettingForm from "../../../Pages/Setting/SettingForm.tsx"; import SettingForm from "../../../Pages/Setting/SettingForm.tsx";
import { NoMarginHelperText, SettingSection, SettingSectionContent } from "../Settings.tsx"; import { NoMarginHelperText, SettingSection, SettingSectionContent } from "../Settings.tsx";
import { SettingContext } from "../SettingWrapper.tsx"; import { SettingContext } from "../SettingWrapper.tsx";
@ -174,6 +175,39 @@ const Email = () => {
</FormControl> </FormControl>
</SettingForm> </SettingForm>
<SettingForm title={t("settings.smtpAuthMethod")} lgWidth={5}>
<FormControl>
<DenseSelect
value={values.smtp_auth ?? "autodiscover"}
onChange={(e) => 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) => (
<SquareMenuItem key={v} value={v}>
<ListItemText
slotProps={{
primary: { variant: "body2" },
}}
>
{t(`settings.smtpAuth_${v.replace(/-/g, "_")}`)}
</ListItemText>
</SquareMenuItem>
))}
</DenseSelect>
<NoMarginHelperText>{t("settings.smtpAuthMethodDes")}</NoMarginHelperText>
</FormControl>
</SettingForm>
<SettingForm title={t("settings.smtpTTL")} lgWidth={5}> <SettingForm title={t("settings.smtpTTL")} lgWidth={5}>
<FormControl fullWidth> <FormControl fullWidth>
<DenseFilledTextField <DenseFilledTextField

@ -302,6 +302,7 @@ const Settings = () => {
"smtpUser", "smtpUser",
"smtpPass", "smtpPass",
"smtpEncryption", "smtpEncryption",
"smtp_auth",
"fromName", "fromName",
"mail_activation_template", "mail_activation_template",
"mail_reset_template", "mail_reset_template",

@ -561,6 +561,7 @@ var DefaultSettings = map[string]string{
"email_filter_mode": "0", "email_filter_mode": "0",
"email_filter_list": "", "email_filter_list": "",
"email_disable_subaddress": "0", "email_disable_subaddress": "0",
"smtp_auth": "autodiscover",
"captcha_type": "normal", "captcha_type": "normal",
"captcha_height": "60", "captcha_height": "60",
"captcha_width": "240", "captcha_width": "240",

@ -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()
})
}

@ -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)
}
}

@ -102,6 +102,36 @@ func (client *SMTPPool) Send(ctx context.Context, to, title, body string) error
return nil 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 关闭发送队列 // Close 关闭发送队列
func (client *SMTPPool) Close() { func (client *SMTPPool) Close() {
if client.ch != nil { if client.ch != nil {
@ -125,7 +155,7 @@ func (client *SMTPPool) Init() {
opts := []mail.Option{ opts := []mail.Option{
mail.WithPort(client.config.Port), mail.WithPort(client.config.Port),
mail.WithTimeout(time.Duration(client.config.Keepalive+5) * time.Second), 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), mail.WithUsername(client.config.User), mail.WithPassword(client.config.Password),
} }
if client.config.ForceEncryption { if client.config.ForceEncryption {

@ -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 "))
}

@ -848,29 +848,53 @@ func (f *DBFS) deleteFiles(ctx context.Context, targets map[Navigator][]*File, f
defer reset() defer reset()
// List all files to be deleted // Walk the tree and delete in bounded batches: accumulating the
toBeDeletedFiles := make([]*File, 0, len(files)) // 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 { 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 { indexToDelete = append(indexToDelete, lo.Map(targets, func(item *File, index int) int {
return item.ID() 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 return nil
}); err != nil { }); err != nil {
return nil, nil, nil, fmt.Errorf("failed to walk files: %w", err) return nil, nil, nil, fmt.Errorf("failed to walk files: %w", err)
} }
// Delete files if len(folderModels) > 0 {
staleEntities, diff, err := fc.Delete(ctx, lo.Map(toBeDeletedFiles, func(item *File, index int) *ent.File { staleEntities, diff, err := fc.Delete(ctx, folderModels, opt)
return item.Model if err != nil {
}), opt) return nil, nil, nil, fmt.Errorf("failed to delete folders: %w", err)
if err != nil { }
return nil, nil, nil, 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)
})...)
} }
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 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) newTargetsMap := make(map[int]*ent.File)
storageDiff := make(inventory.StorageDiff) storageDiff := make(inventory.StorageDiff)
indexToCopy := make([]fs.IndexDiffCopyDetails, 0) indexToCopy := make([]fs.IndexDiffCopyDetails, 0)
var diff inventory.StorageDiff
for n, files := range targets { for n, files := range targets {
initialDstMap := make(map[int][]*ent.File) initialDstMap := make(map[int][]*ent.File)
for _, file := range files { for _, file := range files {
initialDstMap[file.Model.FileChildren] = dstAncestors 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 { if err := n.Walk(ctx, files, limit, intsets.MaxInt, func(targets []*File, level int) error {
// check capacity for each file // check capacity for each file
sizeTotal := int64(0) sizeTotal := int64(0)
@ -922,7 +936,7 @@ func (f *DBFS) copyFiles(ctx context.Context, targets map[Navigator][]*File, des
} }
limit -= len(targets) 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 { Files: lo.Map(targets, func(item *File, index int) *ent.File {
return item.Model 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) return serializer.NewError(serializer.CodeDBError, "Failed to copy files", err)
} }
storageDiff.Merge(diff) // Walk emits each level in bounded pages, so dst mappings must
if firstLayer { // accumulate across callbacks — children of a folder copied in
for k, v := range initialDstMap { // 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] newTargetsMap[k] = v[0]
} }
} }
storageDiff.Merge(diff)
for _, file := range targets { for _, file := range targets {
if _, ok := file.Metadata()[FullTextIndexKey]; ok { 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{ indexToCopy = append(indexToCopy, fs.IndexDiffCopyDetails{
OriginalFileID: file.ID(), OriginalFileID: file.ID(),
FileID: copiedFile.ID, FileID: copiedFile.ID,
@ -958,7 +985,6 @@ func (f *DBFS) copyFiles(ctx context.Context, targets map[Navigator][]*File, des
} }
capacity.Used += sizeTotal capacity.Used += sizeTotal
firstLayer = false
return nil return nil
}); err != nil { }); err != nil {

@ -274,6 +274,12 @@ func (b *baseNavigator) children(ctx context.Context, parent *File, args *ListAr
}, nil }, 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 { func (b *baseNavigator) walk(ctx context.Context, levelFiles []*File, limit, depth int, f WalkFunc) error {
walked := 0 walked := 0
if len(levelFiles) == 0 { if len(levelFiles) == 0 {
@ -281,79 +287,91 @@ func (b *baseNavigator) walk(ctx context.Context, levelFiles []*File, limit, dep
} }
owner := levelFiles[0].Owner() owner := levelFiles[0].Owner()
level := 0 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 { // Files still to emit at the current level. For level 0 this is the
break // 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] for depth >= 0 {
leftCredit := limit - walked depth--
parents := lo.SliceToMap(folders, func(file *File) (int, *File) { folders := make([]*File, 0, 64)
return file.Model.ID, file
}) // Emit this level in bounded batches, tracking folder nodes for
for leftCredit > 0 { // the next level.
token := "" for {
res, err := b.fileClient.GetChildFiles(ctx, var batch []*File
&inventory.ListFileParameters{ if parentModels == nil {
PaginationArgs: &inventory.PaginationArgs{ if len(pending) == 0 {
UseCursorPagination: true, break
PageToken: token, }
PageSize: leftCredit, 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,
}, parentModels...)
owner.ID, if err != nil {
lo.Map(folders, func(item *File, index int) *ent.File { return serializer.NewError(serializer.CodeDBError, "Failed to list children", err)
return item.Model }
})...) if len(res.Files) == 0 {
if err != nil { break
return serializer.NewError(serializer.CodeDBError, "Failed to list children", err) }
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) if err := f(batch, level); err != nil {
return err
levelFiles = append(levelFiles, lo.Map(res.Files, func(model *ent.File, index int) *File { }
p := parents[model.FileChildren] walked += len(batch)
return newFile(p, model) for _, file := range batch {
})...) if file.Model.Type == int(types.FileTypeFolder) && !file.IsSymbolic() {
folders = append(folders, file)
}
}
// All files listed if parentModels == nil {
if res.NextPageToken == "" { continue
}
if token == "" {
break break
} }
}
token = res.NextPageToken if len(folders) == 0 {
return nil
} }
level++
}
if walked >= limit { level++
return ErrFileCountLimitedReached 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 return nil

@ -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)
}

@ -282,6 +282,10 @@ type (
// uploaded. Used to safely trigger CompleteUpload only after every // uploaded. Used to safely trigger CompleteUpload only after every
// chunk has been received when the client uploads chunks concurrently. // chunk has been received when the client uploads chunks concurrently.
ChunksReceived map[int]struct{} 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. // UploadProps properties of an upload session/request.

@ -48,6 +48,12 @@ type (
// uploads chunks concurrently and the last-indexed chunk arrives before some // uploads chunks concurrently and the last-indexed chunk arrives before some
// earlier chunks are still in flight. // earlier chunks are still in flight.
MarkChunkUploaded(ctx context.Context, session *fs.UploadSession, chunkIndex int) (allReceived bool, err error) 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 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 { func (m *manager) CancelUploadSession(ctx context.Context, path *fs.URI, sessionID string) error {
// Get upload session // Get upload session
var session *fs.UploadSession var session *fs.UploadSession

@ -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")
}
}

@ -59,6 +59,16 @@ func TestValidateExternalURL_IPLiterals(t *testing.T) {
"http://100.64.0.1/cgnat", "http://100.64.0.1/cgnat",
"http://[::ffff:127.0.0.1]/v4mapped", "http://[::ffff:127.0.0.1]/v4mapped",
"http://[::ffff:10.0.0.1]/v4mappedpriv", "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", "http://224.0.0.1/multicast",
} }
for _, raw := range cases { for _, raw := range cases {
@ -73,6 +83,9 @@ func TestValidateExternalURL_PublicIP(t *testing.T) {
"http://1.1.1.1/", "http://1.1.1.1/",
"https://8.8.8.8/", "https://8.8.8.8/",
"http://[2606:4700:4700::1111]/", "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 { for _, raw := range cases {
err := request.ValidateExternalURL(ctx, raw, request.SSRFOptions{}) err := request.ValidateExternalURL(ctx, raw, request.SSRFOptions{})

@ -257,6 +257,8 @@ const (
CodeAnonymouseAccessDenied = 40088 CodeAnonymouseAccessDenied = 40088
// CodeInsufficientScope OAuth token scope insufficient // CodeInsufficientScope OAuth token scope insufficient
CodeInsufficientScope = 40089 CodeInsufficientScope = 40089
// CodeRateLimited 请求频率超限
CodeRateLimited = 40090
// CodeDBError 数据库操作失败 // CodeDBError 数据库操作失败
CodeDBError = 50001 CodeDBError = 50001
// CodeEncryptError 加密失败 // CodeEncryptError 加密失败

@ -807,6 +807,7 @@ func (s *settingProvider) SMTP(ctx context.Context) *SMTP {
ForceEncryption: s.getBoolean(ctx, "smtpEncryption", false), ForceEncryption: s.getBoolean(ctx, "smtpEncryption", false),
Port: s.getInt(ctx, "smtpPort", 25), Port: s.getInt(ctx, "smtpPort", 25),
Keepalive: s.getInt(ctx, "mail_keepalive", 30), Keepalive: s.getInt(ctx, "mail_keepalive", 30),
AuthType: s.getString(ctx, "smtp_auth", "autodiscover"),
} }
} }

@ -65,6 +65,7 @@ type SMTP struct {
ForceEncryption bool ForceEncryption bool
Port int Port int
Keepalive int Keepalive int
AuthType string
} }
type TokenAuth struct { type TokenAuth struct {

@ -37,24 +37,24 @@ func RandStringRunes(n int) string {
} }
func RandStringRunesCrypto(n int) string { func RandStringRunesCrypto(n int) string {
b := make([]rune, n) return randStringCrypto(n, RandomVariantAll)
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)
} }
// RandString returns random string in given length and variant // RandString returns random string in given length and variant
func RandString(n int, variant []rune) string { func RandString(n int, variant []rune) string {
return randStringCrypto(n, variant)
}
func randStringCrypto(n int, variant []rune) string {
b := make([]rune, n) b := make([]rune, n)
for i := range b { 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) return string(b)
} }

@ -7,11 +7,13 @@ package webdav // import "golang.org/x/net/webdav"
import ( import (
"context" "context"
"crypto/sha1"
"errors" "errors"
"fmt" "fmt"
"net/http" "net/http"
"net/url" "net/url"
"path" "path"
"strconv"
"strings" "strings"
"time" "time"
@ -256,6 +258,34 @@ func handlePut(c *gin.Context, user *ent.User, fm manager.FileManager) (status i
return http.StatusBadRequest, err 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{ fileData := &fs.UploadRequest{
Props: &fs.UploadProps{ Props: &fs.UploadProps{
Uri: uri, Uri: uri,
@ -266,9 +296,6 @@ func handlePut(c *gin.Context, user *ent.User, fm manager.FileManager) (status i
Mode: fs.ModeOverwrite, Mode: fs.ModeOverwrite,
} }
m := manager.NewFileManager(dependency.FromContext(ctx), user)
defer m.Recycle()
// Update file // Update file
res, err := m.Update(ctx, fileData) res, err := m.Update(ctx, fileData)
if err != nil { if err != nil {
@ -284,6 +311,150 @@ func handlePut(c *gin.Context, user *ent.User, fm manager.FileManager) (status i
return http.StatusCreated, nil 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) { func handleOptions(c *gin.Context, user *ent.User, fm manager.FileManager) (status int, err error) {
allow := []string{"OPTIONS", "LOCK", "PUT", "MKCOL"} allow := []string{"OPTIONS", "LOCK", "PUT", "MKCOL"}

@ -111,3 +111,48 @@ func mustWebDAVTestURI(t *testing.T, raw string) *fs.URI {
} }
return 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)
}
}
}

@ -2,6 +2,7 @@ package routers
import ( import (
"net/http" "net/http"
"time"
"github.com/cloudreve/Cloudreve/v4/application/constants" "github.com/cloudreve/Cloudreve/v4/application/constants"
"github.com/cloudreve/Cloudreve/v4/application/dependency" "github.com/cloudreve/Cloudreve/v4/application/dependency"
@ -289,6 +290,7 @@ func initMasterRouter(dep dependency.Dep) *gin.Engine {
{ {
// 用户登录 // 用户登录
token.POST("", token.POST("",
middleware.RateLimitByIP("login", 10, time.Minute),
middleware.CaptchaRequired(func(c *gin.Context) bool { middleware.CaptchaRequired(func(c *gin.Context) bool {
return dep.SettingProvider().LoginCaptchaEnabled(c) return dep.SettingProvider().LoginCaptchaEnabled(c)
}), }),
@ -298,11 +300,13 @@ func initMasterRouter(dep dependency.Dep) *gin.Engine {
) )
// 2-factor authentication // 2-factor authentication
token.POST("2fa", token.POST("2fa",
middleware.RateLimitByIP("login_2fa", 10, time.Minute),
controllers.FromJSON[usersvc.OtpValidationService](usersvc.OtpValidationParameterCtx{}), controllers.FromJSON[usersvc.OtpValidationService](usersvc.OtpValidationParameterCtx{}),
controllers.UserLogin2FAValidation, controllers.UserLogin2FAValidation,
controllers.UserIssueToken, controllers.UserIssueToken,
) )
token.POST("refresh", token.POST("refresh",
middleware.RateLimitByIP("token_refresh", 30, time.Minute),
middleware.RequiredScopes(types.ScopeOfflineAccess), middleware.RequiredScopes(types.ScopeOfflineAccess),
controllers.FromJSON[usersvc.RefreshTokenService](usersvc.RefreshTokenParameterCtx{}), controllers.FromJSON[usersvc.RefreshTokenService](usersvc.RefreshTokenParameterCtx{}),
controllers.UserRefreshToken, controllers.UserRefreshToken,
@ -331,6 +335,7 @@ func initMasterRouter(dep dependency.Dep) *gin.Engine {
controllers.UserSSOCallback, controllers.UserSSOCallback,
) )
ssoRouter.POST("exchange", ssoRouter.POST("exchange",
middleware.RateLimitByIP("sso_exchange", 20, time.Minute),
controllers.FromJSON[usersvc.SSOExchangeService](usersvc.SSOExchangeParameterCtx{}), controllers.FromJSON[usersvc.SSOExchangeService](usersvc.SSOExchangeParameterCtx{}),
controllers.UserSSOExchange, controllers.UserSSOExchange,
) )
@ -348,6 +353,7 @@ func initMasterRouter(dep dependency.Dep) *gin.Engine {
controllers.GrantAppConsent, controllers.GrantAppConsent,
) )
oauthRouter.POST("token", oauthRouter.POST("token",
middleware.RateLimitByIP("oauth_token", 20, time.Minute),
controllers.FromForm[oauth.ExchangeTokenService](oauth.ExchangeTokenParamCtx{}), controllers.FromForm[oauth.ExchangeTokenService](oauth.ExchangeTokenParamCtx{}),
controllers.ExchangeToken, controllers.ExchangeToken,
) )
@ -369,6 +375,7 @@ func initMasterRouter(dep dependency.Dep) *gin.Engine {
{ {
// WebAuthn login prepare // WebAuthn login prepare
authn.PUT("", authn.PUT("",
middleware.RateLimitByIP("authn", 20, time.Minute),
middleware.IsFunctionEnabled(func(c *gin.Context) bool { middleware.IsFunctionEnabled(func(c *gin.Context) bool {
return dep.SettingProvider().AuthnEnabled(c) return dep.SettingProvider().AuthnEnabled(c)
}), }),
@ -376,6 +383,7 @@ func initMasterRouter(dep dependency.Dep) *gin.Engine {
) )
// WebAuthn finish login // WebAuthn finish login
authn.POST("", authn.POST("",
middleware.RateLimitByIP("authn", 20, time.Minute),
middleware.IsFunctionEnabled(func(c *gin.Context) bool { middleware.IsFunctionEnabled(func(c *gin.Context) bool {
return dep.SettingProvider().AuthnEnabled(c) return dep.SettingProvider().AuthnEnabled(c)
}), }),
@ -391,6 +399,7 @@ func initMasterRouter(dep dependency.Dep) *gin.Engine {
{ {
// 用户注册 Done // 用户注册 Done
user.POST("", user.POST("",
middleware.RateLimitByIP("register", 5, time.Minute),
middleware.IsFunctionEnabled(func(c *gin.Context) bool { middleware.IsFunctionEnabled(func(c *gin.Context) bool {
return dep.SettingProvider().RegisterEnabled(c) return dep.SettingProvider().RegisterEnabled(c)
}), }),
@ -402,12 +411,14 @@ func initMasterRouter(dep dependency.Dep) *gin.Engine {
) )
// 通过邮件里的链接重设密码 // 通过邮件里的链接重设密码
user.PATCH("reset/:id", user.PATCH("reset/:id",
middleware.RateLimitByIP("reset_apply", 10, time.Minute),
middleware.HashID(hashid.UserID), middleware.HashID(hashid.UserID),
controllers.FromJSON[usersvc.UserResetService](usersvc.UserResetParameterCtx{}), controllers.FromJSON[usersvc.UserResetService](usersvc.UserResetParameterCtx{}),
controllers.UserReset, controllers.UserReset,
) )
// 发送密码重设邮件 // 发送密码重设邮件
user.POST("reset", user.POST("reset",
middleware.RateLimitByIP("reset_mail", 5, time.Minute),
middleware.CaptchaRequired(func(c *gin.Context) bool { middleware.CaptchaRequired(func(c *gin.Context) bool {
return dep.SettingProvider().ForgotPasswordCaptchaEnabled(c) return dep.SettingProvider().ForgotPasswordCaptchaEnabled(c)
}), }),

@ -10,6 +10,7 @@ import (
"github.com/cloudreve/Cloudreve/v4/application/dependency" "github.com/cloudreve/Cloudreve/v4/application/dependency"
"github.com/cloudreve/Cloudreve/v4/pkg/boolset" "github.com/cloudreve/Cloudreve/v4/pkg/boolset"
"github.com/cloudreve/Cloudreve/v4/pkg/email"
"github.com/cloudreve/Cloudreve/v4/pkg/filemanager/manager" "github.com/cloudreve/Cloudreve/v4/pkg/filemanager/manager"
request2 "github.com/cloudreve/Cloudreve/v4/pkg/request" request2 "github.com/cloudreve/Cloudreve/v4/pkg/request"
"github.com/cloudreve/Cloudreve/v4/pkg/serializer" "github.com/cloudreve/Cloudreve/v4/pkg/serializer"
@ -142,7 +143,7 @@ func (s *TestSMTPService) Test(c *gin.Context) error {
opts := []mail.Option{ opts := []mail.Option{
mail.WithPort(port), 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"]), mail.WithUsername(s.Settings["smtpUser"]), mail.WithPassword(s.Settings["smtpPass"]),
} }
if setting.IsTrueValue(s.Settings["smtpEncryption"]) { if setting.IsTrueValue(s.Settings["smtpEncryption"]) {

Loading…
Cancel
Save