From 90acfa18e461e052482c6f80528f464d15ed1365 Mon Sep 17 00:00:00 2001 From: Nostalgia Date: Thu, 24 Sep 2026 01:57:21 +0800 Subject: [PATCH] refactor(context): centralize request origin access (#3100) --- internal/authn/authn.go | 3 +- internal/conf/context.go | 8 +++ internal/conf/context_test.go | 29 ++++++++++ internal/fs/archive.go | 3 +- internal/fs/copy_move.go | 3 +- internal/fs/link.go | 4 +- internal/fs/put.go | 3 +- internal/offline_download/tool/add.go | 3 +- internal/offline_download/tool/transfer.go | 5 +- internal/sharing/link.go | 4 +- internal/task/base_test.go | 20 +++++++ server/common/base.go | 3 +- server/middlewares/check_test.go | 62 ++++++++++++++++++++++ 13 files changed, 131 insertions(+), 19 deletions(-) create mode 100644 internal/conf/context.go create mode 100644 internal/conf/context_test.go create mode 100644 internal/task/base_test.go create mode 100644 server/middlewares/check_test.go diff --git a/internal/authn/authn.go b/internal/authn/authn.go index a57823d15..b94454998 100644 --- a/internal/authn/authn.go +++ b/internal/authn/authn.go @@ -6,13 +6,12 @@ import ( "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/setting" - "github.com/OpenListTeam/OpenList/v4/server/common" "github.com/gin-gonic/gin" "github.com/go-webauthn/webauthn/webauthn" ) func NewAuthnInstance(c *gin.Context) (*webauthn.WebAuthn, error) { - siteUrl, err := url.Parse(common.GetApiUrl(c.Request.Context())) + siteUrl, err := url.Parse(conf.GetApiUrl(c.Request.Context())) if err != nil { return nil, err } diff --git a/internal/conf/context.go b/internal/conf/context.go new file mode 100644 index 000000000..5b921d90b --- /dev/null +++ b/internal/conf/context.go @@ -0,0 +1,8 @@ +package conf + +import "context" + +func GetApiUrl(ctx context.Context) string { + api, _ := ctx.Value(ApiUrlKey).(string) + return api +} diff --git a/internal/conf/context_test.go b/internal/conf/context_test.go new file mode 100644 index 000000000..9d92c3188 --- /dev/null +++ b/internal/conf/context_test.go @@ -0,0 +1,29 @@ +package conf_test + +import ( + "context" + "testing" + + "github.com/OpenListTeam/OpenList/v4/internal/conf" +) + +func TestGetApiUrl(t *testing.T) { + const want = "https://openlist.example" + tests := []struct { + name string + ctx context.Context + want string + }{ + {name: "present", ctx: context.WithValue(context.Background(), conf.ApiUrlKey, want), want: want}, + {name: "absent", ctx: context.Background()}, + {name: "wrong type", ctx: context.WithValue(context.Background(), conf.ApiUrlKey, 1)}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := conf.GetApiUrl(tt.ctx); got != tt.want { + t.Fatalf("origin = %q, want %q", got, tt.want) + } + }) + } +} diff --git a/internal/fs/archive.go b/internal/fs/archive.go index 784ba587a..021983562 100644 --- a/internal/fs/archive.go +++ b/internal/fs/archive.go @@ -21,7 +21,6 @@ import ( "github.com/OpenListTeam/OpenList/v4/internal/task" "github.com/OpenListTeam/OpenList/v4/internal/task_group" "github.com/OpenListTeam/OpenList/v4/pkg/utils" - "github.com/OpenListTeam/OpenList/v4/server/common" "github.com/OpenListTeam/tache" "github.com/pkg/errors" log "github.com/sirupsen/logrus" @@ -415,7 +414,7 @@ func archiveDecompress(ctx context.Context, srcObjPath, dstDirPath string, args return nil, err } else { tsk.Creator, _ = ctx.Value(conf.UserKey).(*model.User) - tsk.ApiUrl = common.GetApiUrl(ctx) + tsk.ApiUrl = conf.GetApiUrl(ctx) ArchiveDownloadTaskManager.Add(tsk) return tsk, nil } diff --git a/internal/fs/copy_move.go b/internal/fs/copy_move.go index e78fc9be8..1d171a9b9 100644 --- a/internal/fs/copy_move.go +++ b/internal/fs/copy_move.go @@ -14,7 +14,6 @@ import ( "github.com/OpenListTeam/OpenList/v4/internal/task" "github.com/OpenListTeam/OpenList/v4/internal/task_group" "github.com/OpenListTeam/OpenList/v4/pkg/utils" - "github.com/OpenListTeam/OpenList/v4/server/common" "github.com/OpenListTeam/tache" "github.com/pkg/errors" ) @@ -166,7 +165,7 @@ func transfer(ctx context.Context, taskType taskType, srcObjPath, dstDirPath str } t.Creator, _ = ctx.Value(conf.UserKey).(*model.User) - t.ApiUrl = common.GetApiUrl(ctx) + t.ApiUrl = conf.GetApiUrl(ctx) if taskType == copy || taskType == merge { CopyTaskManager.Add(t) } else { diff --git a/internal/fs/link.go b/internal/fs/link.go index d84090cfc..8e1dedfd1 100644 --- a/internal/fs/link.go +++ b/internal/fs/link.go @@ -4,9 +4,9 @@ import ( "context" "strings" + "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/model" "github.com/OpenListTeam/OpenList/v4/internal/op" - "github.com/OpenListTeam/OpenList/v4/server/common" "github.com/pkg/errors" ) @@ -20,7 +20,7 @@ func link(ctx context.Context, path string, args model.LinkArgs) (*model.Link, m return nil, nil, errors.WithMessage(err, "failed link") } if l.URL != "" && !strings.HasPrefix(l.URL, "http://") && !strings.HasPrefix(l.URL, "https://") { - l.URL = common.GetApiUrl(ctx) + l.URL + l.URL = conf.GetApiUrl(ctx) + l.URL } return l, obj, nil } diff --git a/internal/fs/put.go b/internal/fs/put.go index 0b905be08..9042ae88a 100644 --- a/internal/fs/put.go +++ b/internal/fs/put.go @@ -7,7 +7,6 @@ import ( "time" "github.com/OpenListTeam/OpenList/v4/pkg/utils" - "github.com/OpenListTeam/OpenList/v4/server/common" "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/driver" @@ -81,7 +80,7 @@ func putAsTask(ctx context.Context, dstDirPath string, file model.FileStreamer) t := &UploadTask{ TaskExtension: task.TaskExtension{ Creator: taskCreator, - ApiUrl: common.GetApiUrl(ctx), + ApiUrl: conf.GetApiUrl(ctx), }, storage: storage, dstDirActualPath: dstDirActualPath, diff --git a/internal/offline_download/tool/add.go b/internal/offline_download/tool/add.go index 4a963a51d..68a94f5f4 100644 --- a/internal/offline_download/tool/add.go +++ b/internal/offline_download/tool/add.go @@ -25,7 +25,6 @@ import ( "github.com/OpenListTeam/OpenList/v4/internal/op" "github.com/OpenListTeam/OpenList/v4/internal/setting" "github.com/OpenListTeam/OpenList/v4/internal/task" - "github.com/OpenListTeam/OpenList/v4/server/common" "github.com/google/uuid" "github.com/pkg/errors" ) @@ -184,7 +183,7 @@ func AddURL(ctx context.Context, args *AddURLArgs) (task.TaskExtensionInfo, erro t := &DownloadTask{ TaskExtension: task.TaskExtension{ Creator: taskCreator, - ApiUrl: common.GetApiUrl(ctx), + ApiUrl: conf.GetApiUrl(ctx), }, Url: args.URL, DstDirPath: args.DstDirPath, diff --git a/internal/offline_download/tool/transfer.go b/internal/offline_download/tool/transfer.go index 7109669ee..7c5bd164a 100644 --- a/internal/offline_download/tool/transfer.go +++ b/internal/offline_download/tool/transfer.go @@ -20,7 +20,6 @@ import ( "github.com/OpenListTeam/OpenList/v4/pkg/http_range" "github.com/OpenListTeam/OpenList/v4/pkg/torrent" "github.com/OpenListTeam/OpenList/v4/pkg/utils" - "github.com/OpenListTeam/OpenList/v4/server/common" "github.com/OpenListTeam/tache" "github.com/pkg/errors" log "github.com/sirupsen/logrus" @@ -140,7 +139,7 @@ func transferStd(ctx context.Context, tempDir, dstDirPath string, deletePolicy D TaskData: fs.TaskData{ TaskExtension: task.TaskExtension{ Creator: taskCreator, - ApiUrl: common.GetApiUrl(ctx), + ApiUrl: conf.GetApiUrl(ctx), }, SrcActualPath: stdpath.Join(tempDir, entry.Name()), DstActualPath: dstDirActualPath, @@ -276,7 +275,7 @@ func transferObj(ctx context.Context, tempDir, dstDirPath string, deletePolicy D TaskData: fs.TaskData{ TaskExtension: task.TaskExtension{ Creator: taskCreator, - ApiUrl: common.GetApiUrl(ctx), + ApiUrl: conf.GetApiUrl(ctx), }, SrcActualPath: stdpath.Join(srcObjActualPath, obj.GetName()), DstActualPath: dstDirActualPath, diff --git a/internal/sharing/link.go b/internal/sharing/link.go index 32ae6b836..e021f3c23 100644 --- a/internal/sharing/link.go +++ b/internal/sharing/link.go @@ -4,11 +4,11 @@ import ( "context" "strings" + "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/errs" "github.com/OpenListTeam/OpenList/v4/internal/model" "github.com/OpenListTeam/OpenList/v4/internal/op" "github.com/OpenListTeam/OpenList/v4/pkg/utils" - "github.com/OpenListTeam/OpenList/v4/server/common" "github.com/pkg/errors" ) @@ -38,7 +38,7 @@ func link(ctx context.Context, sid, path string, args *LinkArgs) (*model.Sharing return nil, nil, nil, errors.WithMessage(err, "failed get sharing link") } if l.URL != "" && !strings.HasPrefix(l.URL, "http://") && !strings.HasPrefix(l.URL, "https://") { - l.URL = common.GetApiUrl(ctx) + l.URL + l.URL = conf.GetApiUrl(ctx) + l.URL } return sharing, l, obj, nil } diff --git a/internal/task/base_test.go b/internal/task/base_test.go new file mode 100644 index 000000000..51ddc675b --- /dev/null +++ b/internal/task/base_test.go @@ -0,0 +1,20 @@ +package task_test + +import ( + "context" + "testing" + + "github.com/OpenListTeam/OpenList/v4/internal/conf" + "github.com/OpenListTeam/OpenList/v4/internal/task" +) + +func TestTaskExtensionRestoresAPIURL(t *testing.T) { + const want = "https://openlist.example" + extension := task.TaskExtension{ApiUrl: want} + + extension.SetCtx(context.Background()) + + if got := conf.GetApiUrl(extension.Ctx()); got != want { + t.Fatalf("restored origin = %q, want %q", got, want) + } +} diff --git a/server/common/base.go b/server/common/base.go index 8aa669e3a..53ef633b4 100644 --- a/server/common/base.go +++ b/server/common/base.go @@ -31,6 +31,5 @@ func GetApiUrlFromRequest(r *http.Request) string { } func GetApiUrl(ctx context.Context) string { - api, _ := ctx.Value(conf.ApiUrlKey).(string) - return api + return conf.GetApiUrl(ctx) } diff --git a/server/middlewares/check_test.go b/server/middlewares/check_test.go new file mode 100644 index 000000000..a92136378 --- /dev/null +++ b/server/middlewares/check_test.go @@ -0,0 +1,62 @@ +package middlewares + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/OpenListTeam/OpenList/v4/internal/conf" + "github.com/gin-gonic/gin" +) + +func TestStoragesLoadedAdmitsRequestOrigin(t *testing.T) { + originalMode := gin.Mode() + gin.SetMode(gin.TestMode) + originalConf := conf.Conf + originalLoaded := conf.StoragesLoaded + t.Cleanup(func() { + gin.SetMode(originalMode) + conf.Conf = originalConf + conf.StoragesLoaded = originalLoaded + }) + conf.StoragesLoaded = true + + router := gin.New() + router.Use(StoragesLoaded) + router.GET("/", func(c *gin.Context) { + c.String(http.StatusOK, conf.GetApiUrl(c.Request.Context())) + }) + + assertOrigin := func(name, siteURL, target string, header http.Header, want string) { + t.Run(name, func(t *testing.T) { + conf.Conf = &conf.Config{SiteURL: siteURL} + + req := httptest.NewRequest(http.MethodGet, target, nil) + req.Header = header + rec := httptest.NewRecorder() + router.ServeHTTP(rec, req) + + if got := rec.Body.String(); got != want { + t.Fatalf("origin = %q, want %q", got, want) + } + }) + } + + assertOrigin( + "configured site URL", + "https://openlist.example/base/", + "http://ignored.example/", + nil, + "https://openlist.example/base", + ) + assertOrigin( + "forwarded request", + "", + "http://internal.example/", + http.Header{ + "X-Forwarded-Proto": {"https"}, + "X-Forwarded-Host": {"public.example"}, + }, + "https://public.example", + ) +}