fix: harden request handling for files, search, proxy and SSO

Co-authored-by: PIKACHUIM <PIKACHUIM@users.noreply.github.com>
This commit is contained in:
ILoveScratch
2026-09-29 18:43:35 +08:00
parent 54ae9d7451
commit ea10624fb6
10 changed files with 490 additions and 46 deletions
+1
View File
@@ -211,6 +211,7 @@ func InitialSettings() []model.SettingItem {
{Key: conf.SSODefaultDir, Value: "/", Type: conf.TypeString, Group: model.SSO, Flag: model.PRIVATE}, {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.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.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 // ldap settings
{Key: conf.LdapLoginEnabled, Value: "false", Type: conf.TypeBool, Group: model.LDAP, Flag: model.PUBLIC}, {Key: conf.LdapLoginEnabled, Value: "false", Type: conf.TypeBool, Group: model.LDAP, Flag: model.PUBLIC},
+1
View File
@@ -117,6 +117,7 @@ const (
SSODefaultDir = "sso_default_dir" SSODefaultDir = "sso_default_dir"
SSODefaultPermission = "sso_default_permission" SSODefaultPermission = "sso_default_permission"
SSOCompatibilityMode = "sso_compatibility_mode" SSOCompatibilityMode = "sso_compatibility_mode"
SSOPostMessageOrigin = "sso_postmessage_origin"
// ldap // ldap
LdapLoginEnabled = "ldap_login_enabled" LdapLoginEnabled = "ldap_login_enabled"
+35 -2
View File
@@ -240,16 +240,49 @@ func closeWithError(err error, closer io.Closer) error {
} }
return stderrors.Join(err, closeErr) 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 { func ProcessHeader(origin, override http.Header) http.Header {
result := http.Header{} result := http.Header{}
// client header // client header
for h, val := range origin { 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 continue
} }
result[h] = val result[h] = val
} }
// needed header // needed header, produced by the storage driver rather than the client
for h, val := range override { for h, val := range override {
result[h] = val result[h] = val
} }
+67
View File
@@ -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)
}
}
+12
View File
@@ -113,6 +113,10 @@ func FsMove(c *gin.Context) {
srcDir += "/" srcDir += "/"
} }
for i, name := range req.Names { 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 // ensure req.Names is not a relative path
srcPath := stdpath.Join(srcDir, name) srcPath := stdpath.Join(srcDir, name)
if !strings.HasPrefix(srcPath+"/", srcDir) { if !strings.HasPrefix(srcPath+"/", srcDir) {
@@ -216,6 +220,10 @@ func FsCopy(c *gin.Context) {
srcDir += "/" srcDir += "/"
} }
for i, name := range req.Names { 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 // ensure req.Names is not a relative path
srcPath := stdpath.Join(srcDir, name) srcPath := stdpath.Join(srcDir, name)
if !strings.HasPrefix(srcPath+"/", srcDir) { if !strings.HasPrefix(srcPath+"/", srcDir) {
@@ -373,6 +381,10 @@ func FsRemove(c *gin.Context) {
reqPath += "/" reqPath += "/"
} }
for i, name := range req.Names { for i, name := range req.Names {
if err := checkRelativePath(name); err != nil {
common.ErrorResp(c, err, 403)
return
}
fullPath := stdpath.Join(reqPath, name) fullPath := stdpath.Join(reqPath, name)
if !strings.HasPrefix(fullPath+"/", reqPath) { if !strings.HasPrefix(fullPath+"/", reqPath) {
req.Names[i] = "" req.Names[i] = ""
+126
View File
@@ -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)
}
}
+17 -8
View File
@@ -44,14 +44,7 @@ func Search(c *gin.Context) {
return return
} }
nodes, total, err := search.SearchFiltered(c, req.SearchReq, func(node model.SearchNode) bool { nodes, total, err := search.SearchFiltered(c, req.SearchReq, func(node model.SearchNode) bool {
if !utils.IsSubPath(user.BasePath, node.Parent) { return isSearchNodeAccessible(user, node, req.Password, op.GetNearestMeta)
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)
}) })
if err != nil { if err != nil {
common.ErrorResp(c, err, 500) 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 { func nodeToSearchResp(node model.SearchNode) SearchResp {
return SearchResp{ return SearchResp{
SearchNode: node, SearchNode: node,
+78
View File
@@ -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)
}
})
}
}
+51 -36
View File
@@ -122,6 +122,53 @@ func generateSSOBindingToken(c *gin.Context, purpose, ssoID string) (string, err
return jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString(common.SecretKey) 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(`<!DOCTYPE html>
<head></head>
<body>
<script>
if (window.opener) { window.opener.postMessage(%s, %s) }
window.close()
</script>
</body>`, data, origin)
c.Data(200, "text/html; charset=utf-8", []byte(html))
}
func ssoRedirectUri(c *gin.Context, useCompatibility bool, method string) string { func ssoRedirectUri(c *gin.Context, useCompatibility bool, method string) string {
if useCompatibility { if useCompatibility {
return common.GetApiUrl(c) + "/api/auth/" + method 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) c.Redirect(302, common.GetApiUrl(c)+"/@manage?sso_id="+bindingProof)
return return
} }
html := fmt.Sprintf(`<!DOCTYPE html> ssoPostMessage(c, map[string]string{"sso_id": bindingProof})
<head></head>
<body>
<script>
window.opener.postMessage({"sso_id": "%s"}, "*")
window.close()
</script>
</body>`, bindingProof)
c.Data(200, "text/html; charset=utf-8", []byte(html))
return return
} }
if method == "sso_get_token" { if method == "sso_get_token" {
@@ -367,15 +406,7 @@ func OIDCLoginCallback(c *gin.Context) {
c.Redirect(302, common.GetApiUrl(c)+"/@login?token="+token) c.Redirect(302, common.GetApiUrl(c)+"/@login?token="+token)
return return
} }
html := fmt.Sprintf(`<!DOCTYPE html> ssoPostMessage(c, map[string]string{"token": token})
<head></head>
<body>
<script>
window.opener.postMessage({"token":"%s"}, "*")
window.close()
</script>
</body>`, token)
c.Data(200, "text/html; charset=utf-8", []byte(html))
return return
} }
} }
@@ -516,15 +547,7 @@ func SSOLoginCallback(c *gin.Context) {
c.Redirect(302, common.GetApiUrl(c)+"/@manage?sso_id="+bindingProof) c.Redirect(302, common.GetApiUrl(c)+"/@manage?sso_id="+bindingProof)
return return
} }
html := fmt.Sprintf(`<!DOCTYPE html> ssoPostMessage(c, map[string]string{"sso_id": bindingProof})
<head></head>
<body>
<script>
window.opener.postMessage({"sso_id": "%s"}, "*")
window.close()
</script>
</body>`, bindingProof)
c.Data(200, "text/html; charset=utf-8", []byte(html))
return return
} }
username := utils.Json.Get(resp.Body(), usernameField).ToString() 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) c.Redirect(302, common.GetApiUrl(c)+"/@login?token="+token)
return return
} }
html := fmt.Sprintf(`<!DOCTYPE html> ssoPostMessage(c, map[string]string{"token": token})
<head></head>
<body>
<script>
window.opener.postMessage({"token":"%s"}, "*")
window.close()
</script>
</body>`, token)
c.Data(200, "text/html; charset=utf-8", []byte(html))
} }
+102
View File
@@ -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)
}
}