diff --git a/internal/bootstrap/data/setting.go b/internal/bootstrap/data/setting.go index b902aeef9..4d87f9dbf 100644 --- a/internal/bootstrap/data/setting.go +++ b/internal/bootstrap/data/setting.go @@ -211,6 +211,7 @@ func InitialSettings() []model.SettingItem { {Key: conf.SSODefaultDir, Value: "/", Type: conf.TypeString, Group: model.SSO, Flag: model.PRIVATE}, {Key: conf.SSODefaultPermission, Value: "0", Type: conf.TypeNumber, Group: model.SSO, Flag: model.PRIVATE}, {Key: conf.SSOCompatibilityMode, Value: "false", Type: conf.TypeBool, Group: model.SSO, Flag: model.PUBLIC}, + {Key: conf.SSOPostMessageOrigin, Value: "", Type: conf.TypeString, Group: model.SSO, Flag: model.PUBLIC}, // ldap settings {Key: conf.LdapLoginEnabled, Value: "false", Type: conf.TypeBool, Group: model.LDAP, Flag: model.PUBLIC}, diff --git a/internal/conf/const.go b/internal/conf/const.go index cc8a51d41..2193d3d9d 100644 --- a/internal/conf/const.go +++ b/internal/conf/const.go @@ -117,6 +117,7 @@ const ( SSODefaultDir = "sso_default_dir" SSODefaultPermission = "sso_default_permission" SSOCompatibilityMode = "sso_compatibility_mode" + SSOPostMessageOrigin = "sso_postmessage_origin" // ldap LdapLoginEnabled = "ldap_login_enabled" diff --git a/internal/net/serve.go b/internal/net/serve.go index 5af982bdf..16e9e3fc1 100644 --- a/internal/net/serve.go +++ b/internal/net/serve.go @@ -240,16 +240,49 @@ func closeWithError(err error, closer io.Closer) error { } return stderrors.Join(err, closeErr) } + +// unsafeProxyHeaders are never forwarded from the client request to the +// upstream storage, regardless of the proxy_ignore_headers setting. They either +// carry the caller's credentials, describe the hop to this server rather than +// the hop to upstream, or let the caller influence how upstream routes and +// authenticates the request. +var unsafeProxyHeaders = map[string]struct{}{ + "authorization": {}, + "cookie": {}, + "proxy-authorization": {}, + "www-authenticate": {}, + "host": {}, + "referer": {}, + "origin": {}, + "connection": {}, + "keep-alive": {}, + "proxy-connection": {}, + "te": {}, + "trailer": {}, + "transfer-encoding": {}, + "upgrade": {}, + "forwarded": {}, + "x-forwarded-for": {}, + "x-forwarded-host": {}, + "x-forwarded-proto": {}, + "x-real-ip": {}, +} + + func ProcessHeader(origin, override http.Header) http.Header { result := http.Header{} // client header for h, val := range origin { - if utils.SliceContains(conf.SlicesMap[conf.ProxyIgnoreHeaders], strings.ToLower(h)) { + lower := strings.ToLower(h) + if _, unsafe := unsafeProxyHeaders[lower]; unsafe { + continue + } + if utils.SliceContains(conf.SlicesMap[conf.ProxyIgnoreHeaders], lower) { continue } result[h] = val } - // needed header + // needed header, produced by the storage driver rather than the client for h, val := range override { result[h] = val } diff --git a/internal/net/serve_header_test.go b/internal/net/serve_header_test.go new file mode 100644 index 000000000..d1df43c14 --- /dev/null +++ b/internal/net/serve_header_test.go @@ -0,0 +1,67 @@ +package net + +import ( + "net/http" + "testing" + + "github.com/OpenListTeam/OpenList/v4/internal/conf" +) + +// The client must not be able to smuggle credential or routing headers into the +// request that this server makes to the upstream storage, even when the +// proxy_ignore_headers setting has been emptied. +func TestProcessHeaderDropsUnsafeClientHeaders(t *testing.T) { + conf.SlicesMap[conf.ProxyIgnoreHeaders] = nil + + origin := http.Header{} + origin.Set("Authorization", "Bearer victim-token") + origin.Set("Cookie", "session=victim") + origin.Set("X-Forwarded-For", "127.0.0.1") + origin.Set("Host", "internal.example") + origin.Set("Range", "bytes=0-1023") + + result := ProcessHeader(origin, nil) + + for _, h := range []string{"Authorization", "Cookie", "X-Forwarded-For", "Host"} { + if got := result.Get(h); got != "" { + t.Errorf("header %q must not be forwarded upstream, got %q", h, got) + } + } + if got := result.Get("Range"); got != "bytes=0-1023" { + t.Errorf("Range must be preserved, got %q", got) + } +} + +// Headers supplied by the storage driver still win, since they carry the +// credentials needed to reach upstream. +func TestProcessHeaderOverrideWins(t *testing.T) { + conf.SlicesMap[conf.ProxyIgnoreHeaders] = nil + + origin := http.Header{} + origin.Set("Authorization", "Bearer victim-token") + + override := http.Header{} + override.Set("Authorization", "Bearer driver-token") + + result := ProcessHeader(origin, override) + if got := result.Get("Authorization"); got != "Bearer driver-token" { + t.Errorf("driver header must be used, got %q", got) + } +} + +func TestProcessHeaderStillHonoursIgnoreSetting(t *testing.T) { + conf.SlicesMap[conf.ProxyIgnoreHeaders] = []string{"x-custom"} + t.Cleanup(func() { conf.SlicesMap[conf.ProxyIgnoreHeaders] = nil }) + + origin := http.Header{} + origin.Set("X-Custom", "drop-me") + origin.Set("X-Keep", "keep-me") + + result := ProcessHeader(origin, nil) + if got := result.Get("X-Custom"); got != "" { + t.Errorf("configured ignore header must be dropped, got %q", got) + } + if got := result.Get("X-Keep"); got != "keep-me" { + t.Errorf("unrelated header must be preserved, got %q", got) + } +} diff --git a/server/handles/fsmanage.go b/server/handles/fsmanage.go index d97d36d24..11e904bef 100644 --- a/server/handles/fsmanage.go +++ b/server/handles/fsmanage.go @@ -113,6 +113,10 @@ func FsMove(c *gin.Context) { srcDir += "/" } for i, name := range req.Names { + if err := checkRelativePath(name); err != nil { + common.ErrorResp(c, err, 403) + return + } // ensure req.Names is not a relative path srcPath := stdpath.Join(srcDir, name) if !strings.HasPrefix(srcPath+"/", srcDir) { @@ -216,6 +220,10 @@ func FsCopy(c *gin.Context) { srcDir += "/" } for i, name := range req.Names { + if err := checkRelativePath(name); err != nil { + common.ErrorResp(c, err, 403) + return + } // ensure req.Names is not a relative path srcPath := stdpath.Join(srcDir, name) if !strings.HasPrefix(srcPath+"/", srcDir) { @@ -373,6 +381,10 @@ func FsRemove(c *gin.Context) { reqPath += "/" } for i, name := range req.Names { + if err := checkRelativePath(name); err != nil { + common.ErrorResp(c, err, 403) + return + } fullPath := stdpath.Join(reqPath, name) if !strings.HasPrefix(fullPath+"/", reqPath) { req.Names[i] = "" diff --git a/server/handles/fsmanage_test.go b/server/handles/fsmanage_test.go new file mode 100644 index 000000000..37096e97a --- /dev/null +++ b/server/handles/fsmanage_test.go @@ -0,0 +1,126 @@ +package handles + +import ( + "bytes" + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "testing" + + _ "github.com/OpenListTeam/OpenList/v4/drivers/local" + "github.com/OpenListTeam/OpenList/v4/internal/conf" + "github.com/OpenListTeam/OpenList/v4/internal/db" + "github.com/OpenListTeam/OpenList/v4/internal/model" + "github.com/OpenListTeam/OpenList/v4/internal/op" + "github.com/OpenListTeam/OpenList/v4/pkg/utils" + "github.com/gin-gonic/gin" + "github.com/glebarez/sqlite" + "gorm.io/gorm" +) + +func setupBackslashTraversalTest(t *testing.T, root string, permission int32) *model.User { + t.Helper() + database, err := gorm.Open(sqlite.Open("file:"+t.Name()+"?mode=memory&cache=shared"), &gorm.Config{}) + if err != nil { + t.Fatal(err) + } + conf.Conf = conf.DefaultConfig(t.TempDir()) + db.Init(database) + addition, err := utils.Json.MarshalToString(map[string]string{"root_folder_path": root}) + if err != nil { + t.Fatal(err) + } + if _, err = op.CreateStorage(context.Background(), model.Storage{ + Driver: "Local", MountPath: "/", Addition: addition, + }); err != nil { + t.Fatal(err) + } + return &model.User{ + Username: "restricted-user", BasePath: "/team/a", Role: model.GENERAL, + Permission: permission, + } +} + +func prepareBackslashTraversalFs(t *testing.T) (root string, secretPath string) { + t.Helper() + root = t.TempDir() + if err := os.MkdirAll(filepath.Join(root, "team", "a", "writable"), 0o700); err != nil { + t.Fatal(err) + } + if err := os.MkdirAll(filepath.Join(root, "team", "ab"), 0o700); err != nil { + t.Fatal(err) + } + secretPath = filepath.Join(root, "team", "ab", "secret.txt") + if err := os.WriteFile(secretPath, []byte("synthetic-secret"), 0o600); err != nil { + t.Fatal(err) + } + return root, secretPath +} + +func invokeHandler(t *testing.T, user *model.User, payload any, handler gin.HandlerFunc) *httptest.ResponseRecorder { + t.Helper() + body, err := json.Marshal(payload) + if err != nil { + t.Fatal(err) + } + recorder := httptest.NewRecorder() + ctx, _ := gin.CreateTestContext(recorder) + req := httptest.NewRequest(http.MethodPost, "/api/fs/remove", bytes.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + req = req.WithContext(context.WithValue(req.Context(), conf.UserKey, user)) + ctx.Request = req + handler(ctx) + return recorder +} + +func TestFsRemoveRejectsBackslashTraversal(t *testing.T) { + gin.SetMode(gin.TestMode) + root, secretPath := prepareBackslashTraversalFs(t) + user := setupBackslashTraversalTest(t, root, 1<<3|1<<7) + + for _, name := range []string{"../../ab/secret.txt", `..\..\ab\secret.txt`} { + recorder := invokeHandler(t, user, map[string]any{"dir": "/writable", "names": []string{name}}, FsRemove) + if recorder.Code != http.StatusOK || !strings.Contains(recorder.Body.String(), `"code":403`) { + t.Fatalf("payload %q: got status=%d body=%s, want 403", name, recorder.Code, recorder.Body.String()) + } + if _, err := os.Stat(secretPath); err != nil { + t.Fatalf("payload %q deleted sibling file: %v", name, err) + } + } +} + +func TestFsMoveRejectsBackslashTraversal(t *testing.T) { + gin.SetMode(gin.TestMode) + root, secretPath := prepareBackslashTraversalFs(t) + user := setupBackslashTraversalTest(t, root, 1<<3|1<<5) + + recorder := invokeHandler(t, user, map[string]any{ + "src_dir": "/writable", "dst_dir": "/", "names": []string{`..\..\ab\secret.txt`}, + }, FsMove) + if recorder.Code != http.StatusOK || !strings.Contains(recorder.Body.String(), `"code":403`) { + t.Fatalf("got status=%d body=%s, want 403", recorder.Code, recorder.Body.String()) + } + if _, err := os.Stat(secretPath); err != nil { + t.Fatalf("backslash traversal moved sibling file: %v", err) + } +} + +func TestFsCopyRejectsBackslashTraversal(t *testing.T) { + gin.SetMode(gin.TestMode) + root, secretPath := prepareBackslashTraversalFs(t) + user := setupBackslashTraversalTest(t, root, 1<<3|1<<6) + + recorder := invokeHandler(t, user, map[string]any{ + "src_dir": "/writable", "dst_dir": "/", "names": []string{`..\..\ab\secret.txt`}, + }, FsCopy) + if recorder.Code != http.StatusOK || !strings.Contains(recorder.Body.String(), `"code":403`) { + t.Fatalf("got status=%d body=%s, want 403", recorder.Code, recorder.Body.String()) + } + if _, err := os.Stat(secretPath); err != nil { + t.Fatalf("backslash traversal affected sibling file: %v", err) + } +} diff --git a/server/handles/search.go b/server/handles/search.go index bbc18cae0..e59e25977 100644 --- a/server/handles/search.go +++ b/server/handles/search.go @@ -44,14 +44,7 @@ func Search(c *gin.Context) { return } nodes, total, err := search.SearchFiltered(c, req.SearchReq, func(node model.SearchNode) bool { - if !utils.IsSubPath(user.BasePath, node.Parent) { - return false - } - meta, err := op.GetNearestMeta(node.Parent) - if err != nil && !errors.Is(errors.Cause(err), errs.MetaNotFound) { - return false - } - return common.CanAccess(user, meta, path.Join(node.Parent, node.Name), req.Password) + return isSearchNodeAccessible(user, node, req.Password, op.GetNearestMeta) }) if err != nil { common.ErrorResp(c, err, 500) @@ -63,6 +56,22 @@ func Search(c *gin.Context) { }) } +func isSearchNodeAccessible(user *model.User, node model.SearchNode, password string, resolveMeta func(string) (*model.Meta, error)) bool { + if !utils.IsSubPath(user.BasePath, node.Parent) { + return false + } + nodePath := path.Join(node.Parent, node.Name) + metaPath := node.Parent + if node.IsDir { + metaPath = nodePath + } + meta, err := resolveMeta(metaPath) + if err != nil && !errors.Is(errors.Cause(err), errs.MetaNotFound) { + return false + } + return common.CanAccess(user, meta, nodePath, password) +} + func nodeToSearchResp(node model.SearchNode) SearchResp { return SearchResp{ SearchNode: node, diff --git a/server/handles/search_test.go b/server/handles/search_test.go new file mode 100644 index 000000000..3e22b3329 --- /dev/null +++ b/server/handles/search_test.go @@ -0,0 +1,78 @@ +package handles + +import ( + "path" + "testing" + + "github.com/OpenListTeam/OpenList/v4/internal/errs" + "github.com/OpenListTeam/OpenList/v4/internal/model" +) + +func fakeResolveMeta(metas map[string]*model.Meta) func(string) (*model.Meta, error) { + return func(p string) (*model.Meta, error) { + for { + if meta, ok := metas[p]; ok { + return meta, nil + } + if p == "/" { + return nil, errs.MetaNotFound + } + p = path.Dir(p) + } + } +} + +func TestIsSearchNodeAccessible(t *testing.T) { + tests := []struct { + name string + metas map[string]*model.Meta + node model.SearchNode + want bool + wantMetaPath string + }{ + { + name: "restricted directory", + metas: map[string]*model.Meta{"/private": {Path: "/private", ReadUsers: []uint{1}}}, + node: model.SearchNode{Parent: "/", Name: "private", IsDir: true}, + want: false, + wantMetaPath: "/private", + }, + { + name: "restricted sub directory", + metas: map[string]*model.Meta{"/private": {Path: "/private", ReadUsers: []uint{1}, ReadUsersSub: true}}, + node: model.SearchNode{Parent: "/private", Name: "sub", IsDir: true}, + want: false, + wantMetaPath: "/private/sub", + }, + { + name: "file keeps parent scope", + metas: map[string]*model.Meta{"/private": {Path: "/private", ReadUsers: []uint{1}}}, + node: model.SearchNode{Parent: "/private", Name: "a.txt", IsDir: false}, + want: true, + wantMetaPath: "/private", + }, + { + name: "outside base path", + node: model.SearchNode{Parent: "/other", Name: "private", IsDir: true}, + want: false, + wantMetaPath: "", + }, + } + user := &model.User{ID: 2, BasePath: "/"} + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + resolve := fakeResolveMeta(tt.metas) + var gotMetaPath string + spy := func(p string) (*model.Meta, error) { + gotMetaPath = p + return resolve(p) + } + if got := isSearchNodeAccessible(user, tt.node, "", spy); got != tt.want { + t.Fatalf("isSearchNodeAccessible() = %v, want %v", got, tt.want) + } + if tt.wantMetaPath != "" && gotMetaPath != tt.wantMetaPath { + t.Fatalf("meta resolved at %q, want %q", gotMetaPath, tt.wantMetaPath) + } + }) + } +} diff --git a/server/handles/ssologin.go b/server/handles/ssologin.go index c8096fd4f..3b152505c 100644 --- a/server/handles/ssologin.go +++ b/server/handles/ssologin.go @@ -122,6 +122,53 @@ func generateSSOBindingToken(c *gin.Context, purpose, ssoID string) (string, err return jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString(common.SecretKey) } +// ssoTargetOrigin returns the origin that is allowed to receive the SSO result +// via postMessage. It honours the operator-configured sso_postmessage_origin so +// a frontend served from a different origin than the API can still receive the +// result; otherwise it falls back to the API origin, or "/" to restrict +// delivery to same-origin openers when that cannot be resolved. +func ssoTargetOrigin(c *gin.Context) string { + if configured := setting.GetStr(conf.SSOPostMessageOrigin); configured != "" { + if u, err := url.Parse(configured); err == nil && + (u.Scheme == "http" || u.Scheme == "https") && + u.Host != "" && u.User == nil && + (u.Path == "" || u.Path == "/") && + u.RawQuery == "" && u.Fragment == "" { + return u.Scheme + "://" + u.Host + } + } + u, err := url.Parse(common.GetApiUrl(c)) + if err != nil || u.Scheme == "" || u.Host == "" { + return "/" + } + return u.Scheme + "://" + u.Host +} + +// ssoPostMessage hands the SSO result back to the window that started the login. +// The target origin is pinned so that an arbitrary page cannot open the SSO +// endpoint in a popup and read the payload out of the message event. +func ssoPostMessage(c *gin.Context, payload map[string]string) { + data, err := utils.Json.MarshalToString(payload) + if err != nil { + common.ErrorResp(c, err, 500) + return + } + origin, err := utils.Json.MarshalToString(ssoTargetOrigin(c)) + if err != nil { + common.ErrorResp(c, err, 500) + return + } + html := fmt.Sprintf(` +
+ + + `, data, origin) + c.Data(200, "text/html; charset=utf-8", []byte(html)) +} + func ssoRedirectUri(c *gin.Context, useCompatibility bool, method string) string { if useCompatibility { return common.GetApiUrl(c) + "/api/auth/" + method @@ -338,15 +385,7 @@ func OIDCLoginCallback(c *gin.Context) { c.Redirect(302, common.GetApiUrl(c)+"/@manage?sso_id="+bindingProof) return } - html := fmt.Sprintf(` - - - - `, bindingProof) - c.Data(200, "text/html; charset=utf-8", []byte(html)) + ssoPostMessage(c, map[string]string{"sso_id": bindingProof}) return } if method == "sso_get_token" { @@ -367,15 +406,7 @@ func OIDCLoginCallback(c *gin.Context) { c.Redirect(302, common.GetApiUrl(c)+"/@login?token="+token) return } - html := fmt.Sprintf(` - - - - `, token) - c.Data(200, "text/html; charset=utf-8", []byte(html)) + ssoPostMessage(c, map[string]string{"token": token}) return } } @@ -516,15 +547,7 @@ func SSOLoginCallback(c *gin.Context) { c.Redirect(302, common.GetApiUrl(c)+"/@manage?sso_id="+bindingProof) return } - html := fmt.Sprintf(` - - - - `, bindingProof) - c.Data(200, "text/html; charset=utf-8", []byte(html)) + ssoPostMessage(c, map[string]string{"sso_id": bindingProof}) return } username := utils.Json.Get(resp.Body(), usernameField).ToString() @@ -545,13 +568,5 @@ func SSOLoginCallback(c *gin.Context) { c.Redirect(302, common.GetApiUrl(c)+"/@login?token="+token) return } - html := fmt.Sprintf(` - - - - `, token) - c.Data(200, "text/html; charset=utf-8", []byte(html)) + ssoPostMessage(c, map[string]string{"token": token}) } diff --git a/server/handles/ssologin_origin_test.go b/server/handles/ssologin_origin_test.go new file mode 100644 index 000000000..388ebc746 --- /dev/null +++ b/server/handles/ssologin_origin_test.go @@ -0,0 +1,102 @@ +package handles + +import ( + "context" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/OpenListTeam/OpenList/v4/internal/conf" + "github.com/OpenListTeam/OpenList/v4/internal/model" + "github.com/OpenListTeam/OpenList/v4/internal/op" + "github.com/gin-gonic/gin" +) + +func ssoTestContext(apiUrl string) (*gin.Context, *httptest.ResponseRecorder) { + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, engine := gin.CreateTestContext(rec) + // Matches server.Init, which is what lets GetApiUrl reach the value the + // middleware stored on the request context. + engine.ContextWithFallback = true + req := httptest.NewRequest(http.MethodGet, "/api/auth/sso?method=sso_get_token", nil) + if apiUrl != "" { + req = req.WithContext(context.WithValue(req.Context(), conf.ApiUrlKey, apiUrl)) + } + c.Request = req + // Keep setting lookups off the (uninitialised) database: ssoTargetOrigin + // reads sso_postmessage_origin through the setting cache. + op.Cache.SetSetting(conf.SSOPostMessageOrigin, &model.SettingItem{ + Key: conf.SSOPostMessageOrigin, + Value: "", + }) + return c, rec +} + +// A page that opens the SSO endpoint in a popup must not be able to read the +// token: the postMessage target origin has to name the site, never "*". +func TestSSOPostMessagePinsTargetOrigin(t *testing.T) { + c, rec := ssoTestContext("https://openlist.example.com/base") + ssoPostMessage(c, map[string]string{"token": "secret-token"}) + + body := rec.Body.String() + if strings.Contains(body, `"*"`) || strings.Contains(body, `, '*'`) { + t.Fatalf("wildcard target origin present in response:\n%s", body) + } + if !strings.Contains(body, `"https://openlist.example.com"`) { + t.Errorf("expected the site origin as target, got:\n%s", body) + } + if !strings.Contains(body, "secret-token") { + t.Errorf("payload should still reach a legitimate opener, got:\n%s", body) + } +} + +// A frontend served from a different origin than the API needs the operator to +// be able to point the target at the frontend origin. The configured origin +// must win over the API origin. +func TestSSOPostMessageUsesConfiguredOrigin(t *testing.T) { + c, rec := ssoTestContext("https://api.example.com/base") + op.Cache.SetSetting(conf.SSOPostMessageOrigin, &model.SettingItem{ + Key: conf.SSOPostMessageOrigin, + Value: "https://frontend.example.com", + }) + defer op.Cache.ClearAll() + + ssoPostMessage(c, map[string]string{"token": "secret-token"}) + + body := rec.Body.String() + if !strings.Contains(body, `"https://frontend.example.com"`) { + t.Errorf("expected the configured origin as target, got:\n%s", body) + } +} + +// If the site URL cannot be resolved the fallback must tighten delivery to +// same-origin openers, not widen it back to every origin. +func TestSSOPostMessageFallsBackToSameOrigin(t *testing.T) { + c, rec := ssoTestContext("") + ssoPostMessage(c, map[string]string{"token": "secret-token"}) + + body := rec.Body.String() + if strings.Contains(body, `"*"`) { + t.Fatalf("fallback must not be a wildcard origin:\n%s", body) + } + if !strings.Contains(body, `"/"`) { + t.Errorf(`expected "/" fallback origin, got:\n%s`, body) + } +} + +// userID comes from the identity provider, so it must be encoded rather than +// interpolated into the JS string literal it used to land in. +func TestSSOPostMessageEscapesProviderControlledValue(t *testing.T) { + c, rec := ssoTestContext("https://openlist.example.com") + ssoPostMessage(c, map[string]string{"sso_id": `"});alert(document.domain);//`}) + + body := rec.Body.String() + if strings.Contains(body, `alert(document.domain)`) && !strings.Contains(body, `\"`) { + t.Fatalf("provider value was not escaped:\n%s", body) + } + if !strings.Contains(body, `\"});alert`) { + t.Errorf("expected the injected quote to be escaped, got:\n%s", body) + } +}