Merge pull request #141 from Dvorinka/feature/sso-oidc
feat(sso): inbound OIDC single sign-on + signup email filteringpull/3582/head
commit
6e03fdf8c1
@ -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