mirror of
https://github.com/OpenListTeam/OpenList.git
synced 2026-10-10 04:53:09 +08:00
refactor(context): centralize request origin access (#3100)
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
@@ -0,0 +1,8 @@
|
||||
package conf
|
||||
|
||||
import "context"
|
||||
|
||||
func GetApiUrl(ctx context.Context) string {
|
||||
api, _ := ctx.Value(ApiUrlKey).(string)
|
||||
return api
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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",
|
||||
)
|
||||
}
|
||||
Reference in New Issue
Block a user