Generic OIDC consumer covering Keycloak, Authentik, Logto and other standards-compliant providers. Authorization-code flow with one-time KV state (replay-safe), RS256 ID-token validation against provider JWKS (issuer/audience/expiry/nonce enforced, key-rotation retry), userinfo fallback for the email claim, and automatic account provisioning honoring group defaults and the email filter. Callback hands the SPA a single-use 60s ticket exchanged via JSON API - tokens never touch URLs. All IdP endpoints SSRF-validated with only the issuer host allowlisted. Client secret is redacted in settings responses and kept when submitted empty. Signup email filtering (whitelist/blacklist domain list + sub-address block) now applies to both classic registration and SSO provisioning; the previously dead admin controls are wired to the real keys. Admin UI: OIDC settings accordion (enable toggle, display name, issuer, client id/secret, extra scopes, callback URL helper, auto-register toggle). Login page gains an SSO button and localized error mapping for every failure stage. /callback/sso exchanges the ticket through the existing multi-account session upsert. Covers upstream asks for SSO/third-party sign-in (B.3). Authored By: TDvorak <info@tdvorak.dev>pull/3582/head
parent
f62488a354
commit
aeed9904f9
@ -0,0 +1,64 @@
|
|||||||
|
import { Box, CircularProgress, Typography } from "@mui/material";
|
||||||
|
import { useEffect, useState } from "react";
|
||||||
|
import { useTranslation } from "react-i18next";
|
||||||
|
import { useNavigate } from "react-router-dom";
|
||||||
|
import { sendSSOExchange } from "../../../../api/api.ts";
|
||||||
|
import { setHeadlessFrameLoading } from "../../../../redux/globalStateSlice.ts";
|
||||||
|
import { useAppDispatch } from "../../../../redux/hooks.ts";
|
||||||
|
import { refreshUserSession } from "../../../../redux/thunks/session.ts";
|
||||||
|
import { useQuery } from "../../../../util";
|
||||||
|
import DismissCircleFilled from "../../../Icons/DismissCircleFilled.tsx";
|
||||||
|
|
||||||
|
const SSOCallback = () => {
|
||||||
|
const { t } = useTranslation();
|
||||||
|
const dispatch = useAppDispatch();
|
||||||
|
const navigate = useNavigate();
|
||||||
|
const query = useQuery();
|
||||||
|
const [missingTicket, setMissingTicket] = useState(false);
|
||||||
|
|
||||||
|
useEffect(() => {
|
||||||
|
dispatch(setHeadlessFrameLoading(true));
|
||||||
|
const ticket = query.get("ticket");
|
||||||
|
const redirect = query.get("redirect");
|
||||||
|
|
||||||
|
const finish = async () => {
|
||||||
|
if (!ticket) {
|
||||||
|
setMissingTicket(true);
|
||||||
|
dispatch(setHeadlessFrameLoading(false));
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
try {
|
||||||
|
const loginRes = await dispatch(sendSSOExchange(ticket));
|
||||||
|
dispatch(refreshUserSession(loginRes, redirect));
|
||||||
|
} catch {
|
||||||
|
navigate("/session?sso_error=sso_exchange_failed", { replace: true });
|
||||||
|
}
|
||||||
|
};
|
||||||
|
finish();
|
||||||
|
}, []);
|
||||||
|
|
||||||
|
return (
|
||||||
|
<Box
|
||||||
|
sx={{
|
||||||
|
display: "flex",
|
||||||
|
flexDirection: "column",
|
||||||
|
alignItems: "center",
|
||||||
|
pt: 7,
|
||||||
|
pb: 9,
|
||||||
|
}}
|
||||||
|
>
|
||||||
|
{missingTicket ? (
|
||||||
|
<>
|
||||||
|
<DismissCircleFilled fontSize="large" color="error" />
|
||||||
|
<Typography variant="body2" sx={{ color: (theme) => theme.palette.error.main, mt: 2 }}>
|
||||||
|
{t("login.ssoFailed")}
|
||||||
|
</Typography>
|
||||||
|
</>
|
||||||
|
) : (
|
||||||
|
<CircularProgress />
|
||||||
|
)}
|
||||||
|
</Box>
|
||||||
|
);
|
||||||
|
};
|
||||||
|
|
||||||
|
export default SSOCallback;
|
||||||
@ -0,0 +1,31 @@
|
|||||||
|
import { Button, ButtonProps } from "@mui/material";
|
||||||
|
import { useTranslation } from "react-i18next";
|
||||||
|
import { ApiPrefix } from "../../../../api/request.ts";
|
||||||
|
import { useAppSelector } from "../../../../redux/hooks.ts";
|
||||||
|
import { useQuery } from "../../../../util";
|
||||||
|
import Enter from "../../../Icons/Enter.tsx";
|
||||||
|
|
||||||
|
export default function SSOLoginButton(props: ButtonProps) {
|
||||||
|
const { t } = useTranslation();
|
||||||
|
const query = useQuery();
|
||||||
|
const { sso_enabled, sso_display_name } = useAppSelector((state) => state.siteConfig.login.config);
|
||||||
|
|
||||||
|
if (!sso_enabled) {
|
||||||
|
return null;
|
||||||
|
}
|
||||||
|
|
||||||
|
const startLogin = () => {
|
||||||
|
const redirect = query.get("redirect");
|
||||||
|
const target = new URL(ApiPrefix + "/session/sso", window.location.origin);
|
||||||
|
if (redirect) {
|
||||||
|
target.searchParams.set("redirect", redirect);
|
||||||
|
}
|
||||||
|
window.location.assign(target.toString());
|
||||||
|
};
|
||||||
|
|
||||||
|
return (
|
||||||
|
<Button fullWidth variant="outlined" startIcon={<Enter />} onClick={startLogin} {...props}>
|
||||||
|
{t("login.signInWith", { name: sso_display_name || "SSO" })}
|
||||||
|
</Button>
|
||||||
|
);
|
||||||
|
}
|
||||||
@ -0,0 +1,126 @@
|
|||||||
|
package auth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/rsa"
|
||||||
|
"encoding/base64"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"math/big"
|
||||||
|
|
||||||
|
"github.com/golang-jwt/jwt/v5"
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
ErrOIDCDiscoveryFailed = errors.New("OIDC discovery failed")
|
||||||
|
ErrOIDCKeyNotFound = errors.New("no matching key in JWKS")
|
||||||
|
ErrOIDCInvalidToken = errors.New("invalid OIDC token")
|
||||||
|
)
|
||||||
|
|
||||||
|
// OIDCDiscovery is the subset of the provider metadata document Cloudreve
|
||||||
|
// needs to consume an external OIDC identity provider.
|
||||||
|
type OIDCDiscovery struct {
|
||||||
|
Issuer string `json:"issuer"`
|
||||||
|
AuthorizationEndpoint string `json:"authorization_endpoint"`
|
||||||
|
TokenEndpoint string `json:"token_endpoint"`
|
||||||
|
UserinfoEndpoint string `json:"userinfo_endpoint"`
|
||||||
|
JWKSURI string `json:"jwks_uri"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// OIDCTokenResponse is the token endpoint payload.
|
||||||
|
type OIDCTokenResponse struct {
|
||||||
|
AccessToken string `json:"access_token"`
|
||||||
|
TokenType string `json:"token_type"`
|
||||||
|
ExpiresIn int `json:"expires_in"`
|
||||||
|
RefreshToken string `json:"refresh_token,omitempty"`
|
||||||
|
IDToken string `json:"id_token"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// OIDCUserInfo is the userinfo endpoint payload subset.
|
||||||
|
type OIDCUserInfo struct {
|
||||||
|
Sub string `json:"sub"`
|
||||||
|
Name string `json:"name,omitempty"`
|
||||||
|
PreferredUsername string `json:"preferred_username,omitempty"`
|
||||||
|
Picture string `json:"picture,omitempty"`
|
||||||
|
Email string `json:"email,omitempty"`
|
||||||
|
EmailVerified *bool `json:"email_verified,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate checks the discovery document has the endpoints the code flow requires.
|
||||||
|
func (d *OIDCDiscovery) Validate() error {
|
||||||
|
if d.Issuer == "" || d.AuthorizationEndpoint == "" || d.TokenEndpoint == "" || d.JWKSURI == "" {
|
||||||
|
return fmt.Errorf("incomplete discovery document: %w", ErrOIDCDiscoveryFailed)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// VerifyOIDCIDToken validates signature, issuer, audience, expiry and nonce of
|
||||||
|
// an ID token against the provider's JWKS. Returns ErrOIDCKeyNotFound when the
|
||||||
|
// signing key is absent so callers can refetch JWKS once for key rotation.
|
||||||
|
func VerifyOIDCIDToken(idToken, issuer, clientID, nonce string, jwks *JWKSet) (*OIDCIDTokenClaims, error) {
|
||||||
|
parser := jwt.NewParser()
|
||||||
|
unverified, _, err := parser.ParseUnverified(idToken, &OIDCIDTokenClaims{})
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("malformed id_token: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
kid, _ := unverified.Header["kid"].(string)
|
||||||
|
key, err := findJWK(jwks, kid)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
claims := &OIDCIDTokenClaims{}
|
||||||
|
_, err = jwt.ParseWithClaims(idToken, claims, func(t *jwt.Token) (any, error) {
|
||||||
|
return key, nil
|
||||||
|
},
|
||||||
|
jwt.WithValidMethods([]string{"RS256"}),
|
||||||
|
jwt.WithIssuer(issuer),
|
||||||
|
jwt.WithAudience(clientID),
|
||||||
|
jwt.WithExpirationRequired(),
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("%w: %v", ErrOIDCInvalidToken, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if claims.Nonce != nonce {
|
||||||
|
return nil, fmt.Errorf("nonce mismatch: %w", ErrOIDCInvalidToken)
|
||||||
|
}
|
||||||
|
|
||||||
|
return claims, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// findJWK selects the signing key by kid. A token without kid is only matched
|
||||||
|
// when the provider publishes a single key.
|
||||||
|
func findJWK(jwks *JWKSet, kid string) (*rsa.PublicKey, error) {
|
||||||
|
for _, k := range jwks.Keys {
|
||||||
|
if k.Kty != "RSA" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if k.Kid == kid || (kid == "" && len(jwks.Keys) == 1) {
|
||||||
|
return jwkToRSAPublicKey(k)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil, ErrOIDCKeyNotFound
|
||||||
|
}
|
||||||
|
|
||||||
|
func jwkToRSAPublicKey(jwk JWK) (*rsa.PublicKey, error) {
|
||||||
|
nBytes, err := base64.RawURLEncoding.DecodeString(jwk.N)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid JWK modulus: %w", err)
|
||||||
|
}
|
||||||
|
eBytes, err := base64.RawURLEncoding.DecodeString(jwk.E)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid JWK exponent: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
e := new(big.Int).SetBytes(eBytes).Int64()
|
||||||
|
if e < 3 || e > int64(1<<31-1) {
|
||||||
|
return nil, fmt.Errorf("invalid JWK exponent value")
|
||||||
|
}
|
||||||
|
|
||||||
|
return &rsa.PublicKey{
|
||||||
|
N: new(big.Int).SetBytes(nBytes),
|
||||||
|
E: int(e),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
@ -0,0 +1,103 @@
|
|||||||
|
package auth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/rand"
|
||||||
|
"crypto/rsa"
|
||||||
|
"crypto/x509"
|
||||||
|
"encoding/pem"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/golang-jwt/jwt/v5"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func testRSAKeyPEM(t *testing.T) string {
|
||||||
|
key, err := rsa.GenerateKey(rand.Reader, 2048)
|
||||||
|
require.NoError(t, err)
|
||||||
|
return string(pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(key)}))
|
||||||
|
}
|
||||||
|
|
||||||
|
func signConsumerToken(t *testing.T, pemKey string, claims *OIDCIDTokenClaims) string {
|
||||||
|
token, err := SignOIDCIDToken(pemKey, claims)
|
||||||
|
require.NoError(t, err)
|
||||||
|
return token
|
||||||
|
}
|
||||||
|
|
||||||
|
func validClaims() *OIDCIDTokenClaims {
|
||||||
|
return &OIDCIDTokenClaims{
|
||||||
|
RegisteredClaims: jwt.RegisteredClaims{
|
||||||
|
Issuer: "https://idp.example.com",
|
||||||
|
Audience: jwt.ClaimStrings{"cloudreve"},
|
||||||
|
Subject: "user-1",
|
||||||
|
ExpiresAt: jwt.NewNumericDate(time.Now().Add(time.Hour)),
|
||||||
|
IssuedAt: jwt.NewNumericDate(time.Now()),
|
||||||
|
},
|
||||||
|
Nonce: "nonce-1",
|
||||||
|
Email: "user@example.com",
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestVerifyOIDCIDToken(t *testing.T) {
|
||||||
|
pemKey := testRSAKeyPEM(t)
|
||||||
|
jwks, err := OIDCJWKSet(pemKey)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
t.Run("valid token", func(t *testing.T) {
|
||||||
|
token := signConsumerToken(t, pemKey, validClaims())
|
||||||
|
claims, err := VerifyOIDCIDToken(token, "https://idp.example.com", "cloudreve", "nonce-1", jwks)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, "user@example.com", claims.Email)
|
||||||
|
require.Equal(t, "user-1", claims.Subject)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("wrong issuer", func(t *testing.T) {
|
||||||
|
token := signConsumerToken(t, pemKey, validClaims())
|
||||||
|
_, err := VerifyOIDCIDToken(token, "https://evil.example.com", "cloudreve", "nonce-1", jwks)
|
||||||
|
require.ErrorIs(t, err, ErrOIDCInvalidToken)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("wrong audience", func(t *testing.T) {
|
||||||
|
token := signConsumerToken(t, pemKey, validClaims())
|
||||||
|
_, err := VerifyOIDCIDToken(token, "https://idp.example.com", "other-client", "nonce-1", jwks)
|
||||||
|
require.ErrorIs(t, err, ErrOIDCInvalidToken)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("nonce mismatch", func(t *testing.T) {
|
||||||
|
token := signConsumerToken(t, pemKey, validClaims())
|
||||||
|
_, err := VerifyOIDCIDToken(token, "https://idp.example.com", "cloudreve", "other-nonce", jwks)
|
||||||
|
require.ErrorIs(t, err, ErrOIDCInvalidToken)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("expired token", func(t *testing.T) {
|
||||||
|
claims := validClaims()
|
||||||
|
claims.ExpiresAt = jwt.NewNumericDate(time.Now().Add(-time.Hour))
|
||||||
|
token := signConsumerToken(t, pemKey, claims)
|
||||||
|
_, err := VerifyOIDCIDToken(token, "https://idp.example.com", "cloudreve", "nonce-1", jwks)
|
||||||
|
require.ErrorIs(t, err, ErrOIDCInvalidToken)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("forged signature", func(t *testing.T) {
|
||||||
|
otherPEM := testRSAKeyPEM(t)
|
||||||
|
token := signConsumerToken(t, otherPEM, validClaims())
|
||||||
|
_, err := VerifyOIDCIDToken(token, "https://idp.example.com", "cloudreve", "nonce-1", jwks)
|
||||||
|
// Foreign key produces a different kid -> key not found in JWKS.
|
||||||
|
require.ErrorIs(t, err, ErrOIDCKeyNotFound)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("malformed token", func(t *testing.T) {
|
||||||
|
_, err := VerifyOIDCIDToken("not-a-jwt", "https://idp.example.com", "cloudreve", "nonce-1", jwks)
|
||||||
|
require.Error(t, err)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOIDCDiscoveryValidate(t *testing.T) {
|
||||||
|
require.NoError(t, (&OIDCDiscovery{
|
||||||
|
Issuer: "https://idp.example.com",
|
||||||
|
AuthorizationEndpoint: "https://idp.example.com/authorize",
|
||||||
|
TokenEndpoint: "https://idp.example.com/token",
|
||||||
|
JWKSURI: "https://idp.example.com/jwks",
|
||||||
|
}).Validate())
|
||||||
|
|
||||||
|
require.ErrorIs(t, (&OIDCDiscovery{Issuer: "https://idp.example.com"}).Validate(), ErrOIDCDiscoveryFailed)
|
||||||
|
}
|
||||||
@ -0,0 +1,584 @@
|
|||||||
|
package user
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/sha256"
|
||||||
|
"encoding/base64"
|
||||||
|
"encoding/hex"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
"net/url"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/application/dependency"
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/ent"
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/ent/user"
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/inventory"
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/pkg/auth"
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/pkg/request"
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/pkg/serializer"
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/pkg/setting"
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/pkg/util"
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Parameter contexts for SSO services.
|
||||||
|
type (
|
||||||
|
SSOLoginParameterCtx struct{}
|
||||||
|
SSOCallbackParameterCtx struct{}
|
||||||
|
SSOExchangeParameterCtx struct{}
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
ssoStateTTL = 600 // seconds the authorization state stays valid
|
||||||
|
ssoTicketTTL = 60 // seconds the one-time login ticket stays valid
|
||||||
|
ssoDiscoveryTTL = 3600
|
||||||
|
ssoJWKSTTL = 900
|
||||||
|
)
|
||||||
|
|
||||||
|
// ssoState is the one-time server-side state stored for an in-flight SSO flow.
|
||||||
|
// Storing it in KV (single-use) prevents replay and binds the callback to the
|
||||||
|
// exact login attempt that generated it.
|
||||||
|
type ssoState struct {
|
||||||
|
Nonce string
|
||||||
|
Redirect string
|
||||||
|
}
|
||||||
|
|
||||||
|
// SSOLoginService starts an inbound OIDC flow by redirecting the browser to
|
||||||
|
// the configured identity provider.
|
||||||
|
type SSOLoginService struct {
|
||||||
|
Redirect string `form:"redirect" json:"redirect"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// SSOLogin builds the authorization URL and redirects to the IdP.
|
||||||
|
func (service *SSOLoginService) SSOLogin(c *gin.Context) {
|
||||||
|
dep := dependency.FromContext(c)
|
||||||
|
settings := dep.SettingProvider()
|
||||||
|
sso := settings.SSO(c)
|
||||||
|
|
||||||
|
redirect := sanitizeSSORedirect(service.Redirect)
|
||||||
|
if err := validateSSOConfig(sso); err != nil {
|
||||||
|
redirectToSigninWithError(c, settings, "sso_not_configured")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
discovery, err := ssoDiscovery(c, dep, sso)
|
||||||
|
if err != nil {
|
||||||
|
dep.Logger().Warning("SSO discovery failed: %s", err)
|
||||||
|
redirectToSigninWithError(c, settings, "sso_discovery_failed")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
state := ssoState{
|
||||||
|
Nonce: util.RandStringRunesCrypto(32),
|
||||||
|
Redirect: redirect,
|
||||||
|
}
|
||||||
|
stateKey := ssoStateKey()
|
||||||
|
if err := dep.KV().Set(stateKey, state, ssoStateTTL); err != nil {
|
||||||
|
dep.Logger().Warning("Failed to persist SSO state: %s", err)
|
||||||
|
redirectToSigninWithError(c, settings, "sso_state_failed")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
authorize, err := url.Parse(discovery.AuthorizationEndpoint)
|
||||||
|
if err != nil {
|
||||||
|
redirectToSigninWithError(c, settings, "sso_discovery_failed")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
q := authorize.Query()
|
||||||
|
q.Set("response_type", "code")
|
||||||
|
q.Set("client_id", sso.ClientID)
|
||||||
|
q.Set("redirect_uri", ssoCallbackURL(settings, c))
|
||||||
|
q.Set("scope", sso.Scopes)
|
||||||
|
q.Set("state", stateKey)
|
||||||
|
q.Set("nonce", state.Nonce)
|
||||||
|
authorize.RawQuery = q.Encode()
|
||||||
|
|
||||||
|
c.Redirect(http.StatusFound, authorize.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
// SSOCallbackService completes the inbound OIDC flow.
|
||||||
|
type SSOCallbackService struct {
|
||||||
|
Code string `form:"code" json:"code"`
|
||||||
|
State string `form:"state" json:"state"`
|
||||||
|
Error string `form:"error" json:"error"`
|
||||||
|
ErrorDescription string `form:"error_description" json:"error_description"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// SSOCallback validates the response, resolves the user and redirects to the
|
||||||
|
// SPA with a one-time ticket. Tokens never appear in the URL.
|
||||||
|
func (service *SSOCallbackService) SSOCallback(c *gin.Context) {
|
||||||
|
dep := dependency.FromContext(c)
|
||||||
|
settings := dep.SettingProvider()
|
||||||
|
|
||||||
|
fail := func(code string) {
|
||||||
|
redirectToSigninWithError(c, settings, code)
|
||||||
|
}
|
||||||
|
|
||||||
|
if service.Error != "" {
|
||||||
|
dep.Logger().Info("SSO callback error from IdP: %s (%s)", service.Error, service.ErrorDescription)
|
||||||
|
fail("sso_denied")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
sso := settings.SSO(c)
|
||||||
|
if err := validateSSOConfig(sso); err != nil {
|
||||||
|
fail("sso_not_configured")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if service.State == "" || service.Code == "" {
|
||||||
|
fail("sso_invalid_response")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Single-use state: Get then Delete so a captured state cannot be replayed.
|
||||||
|
rawState, ok := dep.KV().Get(service.State)
|
||||||
|
if !ok {
|
||||||
|
fail("sso_state_expired")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
_ = dep.KV().Delete("", service.State)
|
||||||
|
state, ok := rawState.(ssoState)
|
||||||
|
if !ok {
|
||||||
|
fail("sso_state_expired")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
discovery, err := ssoDiscovery(c, dep, sso)
|
||||||
|
if err != nil {
|
||||||
|
dep.Logger().Warning("SSO discovery failed: %s", err)
|
||||||
|
fail("sso_discovery_failed")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
validate := ssoURLValidator(c, sso)
|
||||||
|
httpClient := dep.RequestClient()
|
||||||
|
|
||||||
|
tokens, err := exchangeOIDCCode(c, httpClient, discovery.TokenEndpoint, service.Code, ssoCallbackURL(settings, c), sso.ClientID, sso.ClientSecret, validate)
|
||||||
|
if err != nil {
|
||||||
|
dep.Logger().Warning("SSO token exchange failed: %s", err)
|
||||||
|
fail("sso_exchange_failed")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
claims, err := ssoVerifyIDToken(c, dep, sso, discovery, tokens.IDToken, state.Nonce)
|
||||||
|
if err != nil {
|
||||||
|
dep.Logger().Warning("SSO id_token validation failed: %s", err)
|
||||||
|
fail("sso_token_invalid")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
email := strings.ToLower(strings.TrimSpace(claims.Email))
|
||||||
|
name := claims.Name
|
||||||
|
preferred := claims.PreferredUsername
|
||||||
|
|
||||||
|
// Some providers omit email from the ID token; fall back to userinfo.
|
||||||
|
if email == "" && tokens.AccessToken != "" && discovery.UserinfoEndpoint != "" {
|
||||||
|
info, err := fetchOIDCUserInfo(c, httpClient, discovery.UserinfoEndpoint, tokens.AccessToken, validate)
|
||||||
|
if err != nil {
|
||||||
|
dep.Logger().Warning("SSO userinfo request failed: %s", err)
|
||||||
|
} else {
|
||||||
|
email = strings.ToLower(strings.TrimSpace(info.Email))
|
||||||
|
if name == "" {
|
||||||
|
name = info.Name
|
||||||
|
}
|
||||||
|
if preferred == "" {
|
||||||
|
preferred = info.PreferredUsername
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if email == "" {
|
||||||
|
fail("sso_no_email")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
targetUser, err := ssoResolveUser(c, dep, sso, email, name, preferred)
|
||||||
|
if err != nil {
|
||||||
|
dep.Logger().Info("SSO login rejected for %q: %s", email, err)
|
||||||
|
fail("sso_account_unavailable")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// One-time ticket -> frontend exchanges it for a token pair via JSON API.
|
||||||
|
ticket := ssoTicketKey()
|
||||||
|
if err := dep.KV().Set(ticket, targetUser.ID, ssoTicketTTL); err != nil {
|
||||||
|
dep.Logger().Warning("Failed to persist SSO ticket: %s", err)
|
||||||
|
fail("sso_state_failed")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
callback := settings.SiteURL(c).ResolveReference(&url.URL{Path: "callback/sso"})
|
||||||
|
q := callback.Query()
|
||||||
|
q.Set("ticket", ticket)
|
||||||
|
if state.Redirect != "" {
|
||||||
|
q.Set("redirect", state.Redirect)
|
||||||
|
}
|
||||||
|
callback.RawQuery = q.Encode()
|
||||||
|
c.Redirect(http.StatusFound, callback.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
// SSOExchangeService trades a one-time callback ticket for a session token.
|
||||||
|
type SSOExchangeService struct {
|
||||||
|
Ticket string `form:"ticket" json:"ticket" binding:"required"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// SSOExchange validates the ticket and issues a builtin token pair.
|
||||||
|
func (service *SSOExchangeService) SSOExchange(c *gin.Context) (any, error) {
|
||||||
|
dep := dependency.FromContext(c)
|
||||||
|
|
||||||
|
raw, ok := dep.KV().Get(service.Ticket)
|
||||||
|
if !ok {
|
||||||
|
return nil, serializer.NewError(serializer.CodeCredentialInvalid, "Invalid or expired SSO ticket", nil)
|
||||||
|
}
|
||||||
|
_ = dep.KV().Delete("", service.Ticket)
|
||||||
|
|
||||||
|
uid, ok := raw.(int)
|
||||||
|
if !ok {
|
||||||
|
return nil, serializer.NewError(serializer.CodeCredentialInvalid, "Invalid SSO ticket", nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
u, err := dep.UserClient().GetByID(c, uid)
|
||||||
|
if err != nil {
|
||||||
|
return nil, serializer.NewError(serializer.CodeUserNotFound, "User not found", err)
|
||||||
|
}
|
||||||
|
if err := checkUserStatus(u); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
util.WithValue(c, inventory.UserCtx{}, u)
|
||||||
|
return IssueToken(c)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ssoResolveUser finds or provisions the local account for an SSO identity.
|
||||||
|
func ssoResolveUser(c *gin.Context, dep dependency.Dep, sso *setting.SSO, email, name, preferred string) (*ent.User, error) {
|
||||||
|
userClient := dep.UserClient()
|
||||||
|
|
||||||
|
u, err := userClient.GetByEmail(c, email)
|
||||||
|
if err == nil {
|
||||||
|
if err := checkUserStatus(u); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return u, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if !sso.RegisterEnabled {
|
||||||
|
return nil, serializer.NewError(serializer.CodeNoPermissionErr, "SSO account provisioning is disabled", nil)
|
||||||
|
}
|
||||||
|
if err := CheckEmailAllowed(dep.SettingProvider().EmailFilter(c), email); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
nick := preferred
|
||||||
|
if nick == "" {
|
||||||
|
nick = name
|
||||||
|
}
|
||||||
|
if nick == "" {
|
||||||
|
nick = strings.Split(email, "@")[0]
|
||||||
|
}
|
||||||
|
|
||||||
|
newUser, err := userClient.Create(c, &inventory.NewUserArgs{
|
||||||
|
Email: email,
|
||||||
|
Nick: nick,
|
||||||
|
Status: user.StatusActive,
|
||||||
|
GroupID: dep.SettingProvider().DefaultGroup(c),
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to provision SSO user: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return newUser, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func checkUserStatus(u *ent.User) error {
|
||||||
|
switch u.Status {
|
||||||
|
case user.StatusSysBanned, user.StatusManualBanned:
|
||||||
|
return serializer.NewError(serializer.CodeUserBaned, "User is banned", nil)
|
||||||
|
case user.StatusInactive:
|
||||||
|
return serializer.NewError(serializer.CodeUserNotActivated, "User is not activated", nil)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// CheckEmailAllowed enforces the sign-up email filter. Shared by the classic
|
||||||
|
// register endpoint and SSO account provisioning.
|
||||||
|
func CheckEmailAllowed(filter *setting.EmailFilter, email string) error {
|
||||||
|
local, domain, found := strings.Cut(email, "@")
|
||||||
|
if !found {
|
||||||
|
return serializer.NewError(serializer.CodeParamErr, "Invalid email", nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
if filter.DisableSubAddress && strings.Contains(local, "+") {
|
||||||
|
return serializer.NewError(serializer.CodeParamErr, "Sub-address emails are not allowed", nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
domain = strings.ToLower(domain)
|
||||||
|
listed := false
|
||||||
|
for _, entry := range filter.List {
|
||||||
|
if entry == domain || strings.HasSuffix(domain, "."+entry) {
|
||||||
|
listed = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
switch filter.Mode {
|
||||||
|
case setting.EmailFilterWhitelist:
|
||||||
|
if !listed {
|
||||||
|
return serializer.NewError(serializer.CodeParamErr, "Email domain is not allowed", nil)
|
||||||
|
}
|
||||||
|
case setting.EmailFilterBlacklist:
|
||||||
|
if listed {
|
||||||
|
return serializer.NewError(serializer.CodeEmailProviderBaned, "Email domain is not allowed", nil)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func validateSSOConfig(sso *setting.SSO) error {
|
||||||
|
if !sso.Enabled || sso.Issuer == "" || sso.ClientID == "" {
|
||||||
|
return errors.New("sso not enabled or not configured")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ssoURLValidator wraps SSRF validation: the operator-trusted issuer host is
|
||||||
|
// allowlisted, endpoints on other internal hosts are rejected.
|
||||||
|
func ssoURLValidator(c *gin.Context, sso *setting.SSO) func(raw string) error {
|
||||||
|
allowed := []string{}
|
||||||
|
if u, err := url.Parse(sso.Issuer); err == nil && u.Hostname() != "" {
|
||||||
|
allowed = append(allowed, u.Hostname())
|
||||||
|
}
|
||||||
|
return func(raw string) error {
|
||||||
|
return request.ValidateExternalURL(c, raw, request.SSRFOptions{AllowedHosts: allowed})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func ssoCallbackURL(settings setting.Provider, c *gin.Context) string {
|
||||||
|
return settings.SiteURL(c).ResolveReference(&url.URL{Path: "api/v4/session/sso/callback"}).String()
|
||||||
|
}
|
||||||
|
|
||||||
|
func redirectToSigninWithError(c *gin.Context, settings setting.Provider, code string) {
|
||||||
|
dest := settings.SiteURL(c).ResolveReference(&url.URL{Path: "session"})
|
||||||
|
q := dest.Query()
|
||||||
|
q.Set("sso_error", code)
|
||||||
|
dest.RawQuery = q.Encode()
|
||||||
|
c.Redirect(http.StatusFound, dest.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
func ssoStateKey() string { return "sso_state_" + util.RandStringRunesCrypto(32) }
|
||||||
|
func ssoTicketKey() string { return "sso_ticket_" + util.RandStringRunesCrypto(32) }
|
||||||
|
|
||||||
|
// sanitizeSSORedirect keeps only same-origin relative paths so a crafted
|
||||||
|
// login URL cannot bounce the browser to an external site after sign-in.
|
||||||
|
func sanitizeSSORedirect(raw string) string {
|
||||||
|
if raw == "" || !strings.HasPrefix(raw, "/") || strings.HasPrefix(raw, "//") {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return raw
|
||||||
|
}
|
||||||
|
|
||||||
|
// ssoDiscovery returns the provider metadata, KV-cached by issuer.
|
||||||
|
func ssoDiscovery(c *gin.Context, dep dependency.Dep, sso *setting.SSO) (*auth.OIDCDiscovery, error) {
|
||||||
|
key := "sso_discovery_" + ssoIssuerHash(sso.Issuer)
|
||||||
|
if cached, ok := dep.KV().Get(key); ok {
|
||||||
|
if doc, ok := cached.(auth.OIDCDiscovery); ok {
|
||||||
|
return &doc, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
validate := ssoURLValidator(c, sso)
|
||||||
|
doc, err := fetchOIDCDiscovery(c, dep.RequestClient(), sso.Issuer, validate)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate each endpoint as well; a compromised discovery doc must not
|
||||||
|
// become an SSRF proxy.
|
||||||
|
for _, endpoint := range []string{doc.TokenEndpoint, doc.JWKSURI, doc.UserinfoEndpoint} {
|
||||||
|
if endpoint == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if err := validate(endpoint); err != nil {
|
||||||
|
return nil, fmt.Errorf("unsafe OIDC endpoint %q: %w", endpoint, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
_ = dep.KV().Set(key, *doc, ssoDiscoveryTTL)
|
||||||
|
return doc, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ssoVerifyIDToken verifies the ID token, refetching JWKS once when the
|
||||||
|
// signing key is unknown (key rotation).
|
||||||
|
func ssoVerifyIDToken(c *gin.Context, dep dependency.Dep, sso *setting.SSO, doc *auth.OIDCDiscovery, idToken, nonce string) (*auth.OIDCIDTokenClaims, error) {
|
||||||
|
jwksKey := "sso_jwks_" + ssoIssuerHash(doc.JWKSURI)
|
||||||
|
validate := ssoURLValidator(c, sso)
|
||||||
|
|
||||||
|
fetch := func() (*auth.JWKSet, error) {
|
||||||
|
jwks, err := fetchOIDCJWKS(c, dep.RequestClient(), doc.JWKSURI, validate)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
_ = dep.KV().Set(jwksKey, *jwks, ssoJWKSTTL)
|
||||||
|
return jwks, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var jwks *auth.JWKSet
|
||||||
|
if cached, ok := dep.KV().Get(jwksKey); ok {
|
||||||
|
if j, ok := cached.(auth.JWKSet); ok {
|
||||||
|
jwks = &j
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if jwks == nil {
|
||||||
|
var err error
|
||||||
|
jwks, err = fetch()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
claims, err := auth.VerifyOIDCIDToken(idToken, doc.Issuer, sso.ClientID, nonce, jwks)
|
||||||
|
if errors.Is(err, auth.ErrOIDCKeyNotFound) {
|
||||||
|
if jwks, err = fetch(); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
claims, err = auth.VerifyOIDCIDToken(idToken, doc.Issuer, sso.ClientID, nonce, jwks)
|
||||||
|
}
|
||||||
|
|
||||||
|
return claims, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func ssoIssuerHash(issuer string) string {
|
||||||
|
sum := sha256.Sum256([]byte(issuer))
|
||||||
|
return hex.EncodeToString(sum[:8])
|
||||||
|
}
|
||||||
|
|
||||||
|
// fetchOIDCDiscovery retrieves the provider metadata document.
|
||||||
|
func fetchOIDCDiscovery(c *gin.Context, client request.Client, issuer string, validate func(raw string) error) (*auth.OIDCDiscovery, error) {
|
||||||
|
wellKnown := strings.TrimRight(issuer, "/") + "/.well-known/openid-configuration"
|
||||||
|
if err := validate(wellKnown); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
resp, err := client.
|
||||||
|
Request(http.MethodGet, wellKnown, nil,
|
||||||
|
request.WithContext(c),
|
||||||
|
request.WithTimeout(10*time.Second),
|
||||||
|
).
|
||||||
|
CheckHTTPResponse(http.StatusOK).
|
||||||
|
GetResponse()
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("%w: %v", auth.ErrOIDCDiscoveryFailed, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var doc auth.OIDCDiscovery
|
||||||
|
if err := json.Unmarshal([]byte(resp), &doc); err != nil {
|
||||||
|
return nil, fmt.Errorf("malformed discovery document: %w", auth.ErrOIDCDiscoveryFailed)
|
||||||
|
}
|
||||||
|
if err := doc.Validate(); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return &doc, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// fetchOIDCJWKS retrieves the provider's signing keys.
|
||||||
|
func fetchOIDCJWKS(c *gin.Context, client request.Client, jwksURI string, validate func(raw string) error) (*auth.JWKSet, error) {
|
||||||
|
if err := validate(jwksURI); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
resp, err := client.
|
||||||
|
Request(http.MethodGet, jwksURI, nil,
|
||||||
|
request.WithContext(c),
|
||||||
|
request.WithTimeout(10*time.Second),
|
||||||
|
).
|
||||||
|
CheckHTTPResponse(http.StatusOK).
|
||||||
|
GetResponse()
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to fetch JWKS: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var jwks auth.JWKSet
|
||||||
|
if err := json.Unmarshal([]byte(resp), &jwks); err != nil {
|
||||||
|
return nil, fmt.Errorf("malformed JWKS: %w", err)
|
||||||
|
}
|
||||||
|
if len(jwks.Keys) == 0 {
|
||||||
|
return nil, auth.ErrOIDCKeyNotFound
|
||||||
|
}
|
||||||
|
|
||||||
|
return &jwks, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// exchangeOIDCCode trades the authorization code for tokens at the token
|
||||||
|
// endpoint. Client secret is sent via HTTP basic auth (client_secret_basic).
|
||||||
|
func exchangeOIDCCode(c *gin.Context, client request.Client, tokenEndpoint, code, redirectURI, clientID, clientSecret string, validate func(raw string) error) (*auth.OIDCTokenResponse, error) {
|
||||||
|
if err := validate(tokenEndpoint); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
form := url.Values{
|
||||||
|
"grant_type": {"authorization_code"},
|
||||||
|
"code": {code},
|
||||||
|
"redirect_uri": {redirectURI},
|
||||||
|
}
|
||||||
|
header := http.Header{
|
||||||
|
"Content-Type": {"application/x-www-form-urlencoded"},
|
||||||
|
"Authorization": {"Basic " + base64.StdEncoding.EncodeToString([]byte(url.QueryEscape(clientID)+":"+url.QueryEscape(clientSecret)))},
|
||||||
|
"Accept": {"application/json"},
|
||||||
|
}
|
||||||
|
|
||||||
|
resp, err := client.
|
||||||
|
Request(http.MethodPost, tokenEndpoint, strings.NewReader(form.Encode()),
|
||||||
|
request.WithContext(c),
|
||||||
|
request.WithTimeout(15*time.Second),
|
||||||
|
request.WithHeader(header),
|
||||||
|
).
|
||||||
|
CheckHTTPResponse(http.StatusOK).
|
||||||
|
GetResponse()
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("token exchange failed: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var tokens auth.OIDCTokenResponse
|
||||||
|
if err := json.Unmarshal([]byte(resp), &tokens); err != nil {
|
||||||
|
return nil, fmt.Errorf("malformed token response: %w", err)
|
||||||
|
}
|
||||||
|
if tokens.IDToken == "" {
|
||||||
|
return nil, fmt.Errorf("missing id_token: %w", auth.ErrOIDCInvalidToken)
|
||||||
|
}
|
||||||
|
|
||||||
|
return &tokens, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// fetchOIDCUserInfo retrieves claims from the userinfo endpoint, used when the
|
||||||
|
// ID token omits the email claim.
|
||||||
|
func fetchOIDCUserInfo(c *gin.Context, client request.Client, endpoint, accessToken string, validate func(raw string) error) (*auth.OIDCUserInfo, error) {
|
||||||
|
if err := validate(endpoint); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
header := http.Header{
|
||||||
|
"Authorization": {"Bearer " + accessToken},
|
||||||
|
"Accept": {"application/json"},
|
||||||
|
}
|
||||||
|
resp, err := client.
|
||||||
|
Request(http.MethodGet, endpoint, nil,
|
||||||
|
request.WithContext(c),
|
||||||
|
request.WithTimeout(10*time.Second),
|
||||||
|
request.WithHeader(header),
|
||||||
|
).
|
||||||
|
CheckHTTPResponse(http.StatusOK).
|
||||||
|
GetResponse()
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("userinfo request failed: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var info auth.OIDCUserInfo
|
||||||
|
if err := json.Unmarshal([]byte(resp), &info); err != nil {
|
||||||
|
return nil, fmt.Errorf("malformed userinfo response: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return &info, nil
|
||||||
|
}
|
||||||
@ -0,0 +1,137 @@
|
|||||||
|
package user
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/pkg/serializer"
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/pkg/setting"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestCheckEmailAllowed(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
filter *setting.EmailFilter
|
||||||
|
email string
|
||||||
|
wantErr bool
|
||||||
|
code int
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "disabled filter allows any domain",
|
||||||
|
filter: &setting.EmailFilter{Mode: setting.EmailFilterDisabled},
|
||||||
|
email: "a@anything.com",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "whitelist allows listed domain",
|
||||||
|
filter: &setting.EmailFilter{
|
||||||
|
Mode: setting.EmailFilterWhitelist,
|
||||||
|
List: []string{"example.com"},
|
||||||
|
},
|
||||||
|
email: "a@example.com",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "whitelist allows subdomain of listed domain",
|
||||||
|
filter: &setting.EmailFilter{
|
||||||
|
Mode: setting.EmailFilterWhitelist,
|
||||||
|
List: []string{"example.com"},
|
||||||
|
},
|
||||||
|
email: "a@mail.example.com",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "whitelist rejects unlisted domain",
|
||||||
|
filter: &setting.EmailFilter{
|
||||||
|
Mode: setting.EmailFilterWhitelist,
|
||||||
|
List: []string{"example.com"},
|
||||||
|
},
|
||||||
|
email: "a@other.com",
|
||||||
|
wantErr: true,
|
||||||
|
code: serializer.CodeParamErr,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "whitelist does not suffix-match partial domain",
|
||||||
|
filter: &setting.EmailFilter{
|
||||||
|
Mode: setting.EmailFilterWhitelist,
|
||||||
|
List: []string{"ample.com"},
|
||||||
|
},
|
||||||
|
email: "a@example.com",
|
||||||
|
wantErr: true,
|
||||||
|
code: serializer.CodeParamErr,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "blacklist rejects listed domain",
|
||||||
|
filter: &setting.EmailFilter{
|
||||||
|
Mode: setting.EmailFilterBlacklist,
|
||||||
|
List: []string{"spam.com"},
|
||||||
|
},
|
||||||
|
email: "a@spam.com",
|
||||||
|
wantErr: true,
|
||||||
|
code: serializer.CodeEmailProviderBaned,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "blacklist rejects subdomain of listed domain",
|
||||||
|
filter: &setting.EmailFilter{
|
||||||
|
Mode: setting.EmailFilterBlacklist,
|
||||||
|
List: []string{"spam.com"},
|
||||||
|
},
|
||||||
|
email: "a@mx.spam.com",
|
||||||
|
wantErr: true,
|
||||||
|
code: serializer.CodeEmailProviderBaned,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "blacklist allows unlisted domain",
|
||||||
|
filter: &setting.EmailFilter{
|
||||||
|
Mode: setting.EmailFilterBlacklist,
|
||||||
|
List: []string{"spam.com"},
|
||||||
|
},
|
||||||
|
email: "a@ok.com",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "sub-address rejected when disabled",
|
||||||
|
filter: &setting.EmailFilter{
|
||||||
|
Mode: setting.EmailFilterDisabled,
|
||||||
|
DisableSubAddress: true,
|
||||||
|
},
|
||||||
|
email: "a+tag@example.com",
|
||||||
|
wantErr: true,
|
||||||
|
code: serializer.CodeParamErr,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "sub-address allowed when enabled",
|
||||||
|
filter: &setting.EmailFilter{
|
||||||
|
Mode: setting.EmailFilterDisabled,
|
||||||
|
DisableSubAddress: false,
|
||||||
|
},
|
||||||
|
email: "a+tag@example.com",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "invalid email rejected",
|
||||||
|
filter: &setting.EmailFilter{Mode: setting.EmailFilterDisabled},
|
||||||
|
email: "not-an-email",
|
||||||
|
wantErr: true,
|
||||||
|
code: serializer.CodeParamErr,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
err := CheckEmailAllowed(tt.filter, tt.email)
|
||||||
|
if tt.wantErr {
|
||||||
|
require.Error(t, err)
|
||||||
|
if appErr, ok := err.(serializer.AppError); ok {
|
||||||
|
require.Equal(t, tt.code, appErr.Code)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
require.NoError(t, err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSanitizeSSORedirect(t *testing.T) {
|
||||||
|
require.Equal(t, "", sanitizeSSORedirect(""))
|
||||||
|
require.Equal(t, "", sanitizeSSORedirect("https://evil.com"))
|
||||||
|
require.Equal(t, "", sanitizeSSORedirect("//evil.com"))
|
||||||
|
require.Equal(t, "", sanitizeSSORedirect("evil.com/path"))
|
||||||
|
require.Equal(t, "/home", sanitizeSSORedirect("/home"))
|
||||||
|
require.Equal(t, "/share/s/abc", sanitizeSSORedirect("/share/s/abc"))
|
||||||
|
}
|
||||||
Loading…
Reference in new issue