refactor(context): centralize request origin access (#3100)

This commit is contained in:
Nostalgia
2026-09-24 01:57:21 +08:00
committed by GitHub
parent 1462d63a48
commit 90acfa18e4
13 changed files with 131 additions and 19 deletions
+1 -2
View File
@@ -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
}
+8
View File
@@ -0,0 +1,8 @@
package conf
import "context"
func GetApiUrl(ctx context.Context) string {
api, _ := ctx.Value(ApiUrlKey).(string)
return api
}
+29
View File
@@ -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)
}
})
}
}
+1 -2
View File
@@ -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
}
+1 -2
View File
@@ -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 {
+2 -2
View File
@@ -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
}
+1 -2
View File
@@ -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,
+1 -2
View File
@@ -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,
+2 -3
View File
@@ -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,
+2 -2
View File
@@ -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
}
+20
View File
@@ -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)
}
}
+1 -2
View File
@@ -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)
}
+62
View File
@@ -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",
)
}