mirror of
https://github.com/netbirdio/dex.git
synced 2026-05-22 18:43:53 -07:00
feat: implement id_token_hint (#4670)
Signed-off-by: maksim.nabokikh <max.nabokih@gmail.com> Signed-off-by: Maksim Nabokikh <maksim.nabokikh@flant.com>
This commit is contained in:
+20
-3
@@ -295,7 +295,7 @@ func (s *Server) getClientWithAuthError(ctx context.Context, clientID string) (s
|
||||
|
||||
func (s *Server) handleConnectorLogin(w http.ResponseWriter, r *http.Request) {
|
||||
ctx := r.Context()
|
||||
authReq, err := s.parseAuthorizationRequest(r)
|
||||
authReq, hintSubject, err := s.parseAuthorizationRequest(r)
|
||||
if err != nil {
|
||||
s.logger.ErrorContext(r.Context(), "failed to parse authorization request", "err", err)
|
||||
|
||||
@@ -374,9 +374,26 @@ func (s *Server) handleConnectorLogin(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
// handle prompt only if sessions are enabled
|
||||
if s.sessionConfig != nil {
|
||||
// Retrieve the session once for use in both hint and prompt logic.
|
||||
session := s.getValidAuthSession(ctx, w, r, authReq)
|
||||
|
||||
// id_token_hint logic (OIDC Core 1.0 3.1.2.1):
|
||||
// When a hint is provided, verify that the session user matches.
|
||||
if hintSubject != "" {
|
||||
if !sessionMatchesHint(session, hintSubject) {
|
||||
// Clear the session if the user is different from the hint.
|
||||
session = nil
|
||||
}
|
||||
if session == nil && prompt.None() {
|
||||
// Cannot authenticate silently with prompt=none.
|
||||
s.redirectWithError(w, r, authReq, errLoginRequired, "id_token_hint does not match authenticated user")
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// prompt=none: no UI allowed.
|
||||
if prompt.None() {
|
||||
redirectURL, ok := s.trySessionLogin(ctx, r, w, authReq)
|
||||
redirectURL, ok := s.trySessionLoginWithSession(ctx, r, w, authReq, session)
|
||||
if !ok {
|
||||
s.redirectWithError(w, r, authReq, errLoginRequired, "User not authenticated")
|
||||
return
|
||||
@@ -391,7 +408,7 @@ func (s *Server) handleConnectorLogin(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
if !prompt.Login() {
|
||||
// Normal flow: try session-based login (skip if prompt=login forces re-auth).
|
||||
if redirectURL, ok := s.trySessionLogin(ctx, r, w, authReq); ok {
|
||||
if redirectURL, ok := s.trySessionLoginWithSession(ctx, r, w, authReq, session); ok {
|
||||
if redirectURL != "" {
|
||||
http.Redirect(w, r, redirectURL, http.StatusSeeOther)
|
||||
}
|
||||
|
||||
+72
-26
@@ -19,6 +19,7 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/coreos/go-oidc/v3/oidc"
|
||||
"github.com/go-jose/go-jose/v4"
|
||||
|
||||
"github.com/dexidp/dex/connector"
|
||||
@@ -438,16 +439,51 @@ func (s *Server) newIDToken(ctx context.Context, clientID string, claims storage
|
||||
return idToken, expiry, nil
|
||||
}
|
||||
|
||||
// parse the initial request from the OAuth2 client.
|
||||
func (s *Server) parseAuthorizationRequest(r *http.Request) (*storage.AuthRequest, error) {
|
||||
// validateIDTokenHint verifies the signature and issuer of an id_token_hint.
|
||||
// Expired tokens are accepted per OIDC Core 1.0 §3.1.2.1.
|
||||
// Returns the raw subject claim from the token.
|
||||
func (s *Server) validateIDTokenHint(ctx context.Context, hint string) (string, error) {
|
||||
verifier := oidc.NewVerifier(s.issuerURL.String(), &signerKeySet{s.signer}, &oidc.Config{
|
||||
SkipExpiryCheck: true,
|
||||
// SkipClientIDCheck is set because the hint may originate from any client that
|
||||
// Dex issued a token to — the caller does not know the expected audience in advance.
|
||||
// The signature verification via signerKeySet already guarantees the token was
|
||||
// issued by this server, which is sufficient for a hint.
|
||||
// Dex does the client id check later in the scope of the session validation.
|
||||
SkipClientIDCheck: true,
|
||||
})
|
||||
idToken, err := verifier.Verify(ctx, hint)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return idToken.Subject, nil
|
||||
}
|
||||
|
||||
// sessionMatchesHint checks whether the session's user identity matches the
|
||||
// subject from an id_token_hint by encoding the session's (userID, connectorID)
|
||||
// via genSubject and doing a string comparison.
|
||||
func sessionMatchesHint(session *storage.AuthSession, hintSubject string) bool {
|
||||
if session == nil {
|
||||
return false
|
||||
}
|
||||
encoded, err := genSubject(session.UserID, session.ConnectorID)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return encoded == hintSubject
|
||||
}
|
||||
|
||||
// parseAuthorizationRequest parses the initial request from the OAuth2 client.
|
||||
// Returns the auth request, the raw subject from id_token_hint (empty if not provided), and any error.
|
||||
func (s *Server) parseAuthorizationRequest(r *http.Request) (*storage.AuthRequest, string, error) {
|
||||
ctx := r.Context()
|
||||
if err := r.ParseForm(); err != nil {
|
||||
return nil, newDisplayedErr(http.StatusBadRequest, "Failed to parse request.")
|
||||
return nil, "", newDisplayedErr(http.StatusBadRequest, "Failed to parse request.")
|
||||
}
|
||||
q := r.Form
|
||||
redirectURI, err := url.QueryUnescape(q.Get("redirect_uri"))
|
||||
if err != nil {
|
||||
return nil, newDisplayedErr(http.StatusBadRequest, "No redirect_uri provided.")
|
||||
return nil, "", newDisplayedErr(http.StatusBadRequest, "No redirect_uri provided.")
|
||||
}
|
||||
|
||||
clientID := q.Get("client_id")
|
||||
@@ -469,15 +505,15 @@ func (s *Server) parseAuthorizationRequest(r *http.Request) (*storage.AuthReques
|
||||
if err != nil {
|
||||
if err == storage.ErrNotFound {
|
||||
s.logger.ErrorContext(r.Context(), "invalid client_id provided", "client_id", clientID)
|
||||
return nil, newDisplayedErr(http.StatusNotFound, "Invalid client_id.")
|
||||
return nil, "", newDisplayedErr(http.StatusNotFound, "Invalid client_id.")
|
||||
}
|
||||
s.logger.ErrorContext(r.Context(), "failed to get client", "err", err)
|
||||
return nil, newDisplayedErr(http.StatusInternalServerError, "Database error.")
|
||||
return nil, "", newDisplayedErr(http.StatusInternalServerError, "Database error.")
|
||||
}
|
||||
|
||||
if !validateRedirectURI(client, redirectURI) {
|
||||
s.logger.ErrorContext(r.Context(), "unregistered redirect_uri", "redirect_uri", redirectURI, "client_id", clientID)
|
||||
return nil, newDisplayedErr(http.StatusBadRequest, "Unregistered redirect_uri.")
|
||||
return nil, "", newDisplayedErr(http.StatusBadRequest, "Unregistered redirect_uri.")
|
||||
}
|
||||
if redirectURI == deviceCallbackURI && client.Public {
|
||||
redirectURI = s.absPath(deviceCallbackURI)
|
||||
@@ -492,30 +528,30 @@ func (s *Server) parseAuthorizationRequest(r *http.Request) (*storage.AuthReques
|
||||
connectors, err := s.storage.ListConnectors(ctx)
|
||||
if err != nil {
|
||||
s.logger.ErrorContext(r.Context(), "failed to list connectors", "err", err)
|
||||
return nil, newRedirectedErr(errServerError, "Unable to retrieve connectors")
|
||||
return nil, "", newRedirectedErr(errServerError, "Unable to retrieve connectors")
|
||||
}
|
||||
if !validateConnectorID(connectors, connectorID) {
|
||||
return nil, newRedirectedErr(errInvalidRequest, "Invalid ConnectorID")
|
||||
return nil, "", newRedirectedErr(errInvalidRequest, "Invalid ConnectorID")
|
||||
}
|
||||
if !isConnectorAllowed(client.AllowedConnectors, connectorID) {
|
||||
return nil, newRedirectedErr(errInvalidRequest, "Connector not allowed for this client")
|
||||
return nil, "", newRedirectedErr(errInvalidRequest, "Connector not allowed for this client")
|
||||
}
|
||||
}
|
||||
|
||||
// dex doesn't support request parameter and must return request_not_supported error
|
||||
// https://openid.net/specs/openid-connect-core-1_0.html#6.1
|
||||
if q.Get("request") != "" {
|
||||
return nil, newRedirectedErr(errRequestNotSupported, "Server does not support request parameter.")
|
||||
return nil, "", newRedirectedErr(errRequestNotSupported, "Server does not support request parameter.")
|
||||
}
|
||||
|
||||
if codeChallenge != "" && !slices.Contains(s.pkce.CodeChallengeMethodsSupported, codeChallengeMethod) {
|
||||
return nil, newRedirectedErr(errInvalidRequest, "Unsupported PKCE challenge method (%q).", codeChallengeMethod)
|
||||
return nil, "", newRedirectedErr(errInvalidRequest, "Unsupported PKCE challenge method (%q).", codeChallengeMethod)
|
||||
}
|
||||
|
||||
// Enforce PKCE if configured.
|
||||
// https://datatracker.ietf.org/doc/html/draft-ietf-oauth-v2-1-12#section-4.1.1
|
||||
if s.pkce.Enforce && codeChallenge == "" {
|
||||
return nil, newRedirectedErr(errInvalidRequest, "PKCE is required. The code_challenge parameter must be provided.")
|
||||
return nil, "", newRedirectedErr(errInvalidRequest, "PKCE is required. The code_challenge parameter must be provided.")
|
||||
}
|
||||
|
||||
var (
|
||||
@@ -537,7 +573,7 @@ func (s *Server) parseAuthorizationRequest(r *http.Request) (*storage.AuthReques
|
||||
|
||||
isTrusted, err := s.validateCrossClientTrust(r.Context(), clientID, peerID)
|
||||
if err != nil {
|
||||
return nil, newRedirectedErr(errServerError, "Internal server error.")
|
||||
return nil, "", newRedirectedErr(errServerError, "Internal server error.")
|
||||
}
|
||||
if !isTrusted {
|
||||
invalidScopes = append(invalidScopes, scope)
|
||||
@@ -545,13 +581,13 @@ func (s *Server) parseAuthorizationRequest(r *http.Request) (*storage.AuthReques
|
||||
}
|
||||
}
|
||||
if !hasOpenIDScope {
|
||||
return nil, newRedirectedErr(errInvalidScope, `Missing required scope(s) ["openid"].`)
|
||||
return nil, "", newRedirectedErr(errInvalidScope, `Missing required scope(s) ["openid"].`)
|
||||
}
|
||||
if len(unrecognized) > 0 {
|
||||
return nil, newRedirectedErr(errInvalidScope, "Unrecognized scope(s) %q", unrecognized)
|
||||
return nil, "", newRedirectedErr(errInvalidScope, "Unrecognized scope(s) %q", unrecognized)
|
||||
}
|
||||
if len(invalidScopes) > 0 {
|
||||
return nil, newRedirectedErr(errInvalidScope, "Client can't request scope(s) %q", invalidScopes)
|
||||
return nil, "", newRedirectedErr(errInvalidScope, "Client can't request scope(s) %q", invalidScopes)
|
||||
}
|
||||
|
||||
var rt struct {
|
||||
@@ -569,23 +605,23 @@ func (s *Server) parseAuthorizationRequest(r *http.Request) (*storage.AuthReques
|
||||
case responseTypeToken:
|
||||
rt.token = true
|
||||
default:
|
||||
return nil, newRedirectedErr(errInvalidRequest, "Invalid response type %q", responseType)
|
||||
return nil, "", newRedirectedErr(errInvalidRequest, "Invalid response type %q", responseType)
|
||||
}
|
||||
|
||||
if !s.supportedResponseTypes[responseType] {
|
||||
return nil, newRedirectedErr(errUnsupportedResponseType, "Unsupported response type %q", responseType)
|
||||
return nil, "", newRedirectedErr(errUnsupportedResponseType, "Unsupported response type %q", responseType)
|
||||
}
|
||||
}
|
||||
|
||||
if len(responseTypes) == 0 {
|
||||
return nil, newRedirectedErr(errInvalidRequest, "No response_type provided")
|
||||
return nil, "", newRedirectedErr(errInvalidRequest, "No response_type provided")
|
||||
}
|
||||
|
||||
if rt.token && !rt.code && !rt.idToken {
|
||||
// "token" can't be provided by its own.
|
||||
//
|
||||
// https://openid.net/specs/openid-connect-core-1_0.html#Authentication
|
||||
return nil, newRedirectedErr(errInvalidRequest, "Response type 'token' must be provided with type 'id_token' and/or 'code'")
|
||||
return nil, "", newRedirectedErr(errInvalidRequest, "Response type 'token' must be provided with type 'id_token' and/or 'code'")
|
||||
}
|
||||
if !rt.code {
|
||||
// Either "id_token token" or "id_token" has been provided which implies the
|
||||
@@ -593,18 +629,18 @@ func (s *Server) parseAuthorizationRequest(r *http.Request) (*storage.AuthReques
|
||||
//
|
||||
// https://openid.net/specs/openid-connect-core-1_0.html#ImplicitAuthRequest
|
||||
if nonce == "" {
|
||||
return nil, newRedirectedErr(errInvalidRequest, "Response type 'token' requires a 'nonce' value.")
|
||||
return nil, "", newRedirectedErr(errInvalidRequest, "Response type 'token' requires a 'nonce' value.")
|
||||
}
|
||||
}
|
||||
if rt.token {
|
||||
if redirectURI == redirectURIOOB {
|
||||
return nil, newRedirectedErr(errInvalidRequest, "Cannot use response type 'token' with redirect_uri '%s'.", redirectURIOOB)
|
||||
return nil, "", newRedirectedErr(errInvalidRequest, "Cannot use response type 'token' with redirect_uri '%s'.", redirectURIOOB)
|
||||
}
|
||||
}
|
||||
|
||||
prompt, err := ParsePrompt(q.Get("prompt"))
|
||||
if err != nil {
|
||||
return nil, newRedirectedErr(errInvalidRequest, "Invalid prompt parameter: %v", err)
|
||||
return nil, "", newRedirectedErr(errInvalidRequest, "Invalid prompt parameter: %v", err)
|
||||
}
|
||||
|
||||
// Parse max_age: -1 means not specified.
|
||||
@@ -612,7 +648,7 @@ func (s *Server) parseAuthorizationRequest(r *http.Request) (*storage.AuthReques
|
||||
if maxAgeStr := q.Get("max_age"); maxAgeStr != "" {
|
||||
v, err := strconv.Atoi(maxAgeStr)
|
||||
if err != nil || v < 0 {
|
||||
return nil, newRedirectedErr(errInvalidRequest, "Invalid max_age value %q", maxAgeStr)
|
||||
return nil, "", newRedirectedErr(errInvalidRequest, "Invalid max_age value %q", maxAgeStr)
|
||||
}
|
||||
maxAge = v
|
||||
}
|
||||
@@ -620,6 +656,16 @@ func (s *Server) parseAuthorizationRequest(r *http.Request) (*storage.AuthReques
|
||||
// OIDC prompt=consent implies force approval.
|
||||
forceApproval := q.Get("approval_prompt") == "force" || prompt.Consent()
|
||||
|
||||
// Validate id_token_hint if provided (OIDC Core 1.0 §3.1.2.1).
|
||||
var idTokenHintSubject string
|
||||
if hint := q.Get("id_token_hint"); hint != "" {
|
||||
sub, err := s.validateIDTokenHint(ctx, hint)
|
||||
if err != nil {
|
||||
return nil, "", newRedirectedErr(errInvalidRequest, "Invalid id_token_hint.")
|
||||
}
|
||||
idTokenHintSubject = sub
|
||||
}
|
||||
|
||||
return &storage.AuthRequest{
|
||||
ID: storage.NewID(),
|
||||
ClientID: client.ID,
|
||||
@@ -637,7 +683,7 @@ func (s *Server) parseAuthorizationRequest(r *http.Request) (*storage.AuthReques
|
||||
CodeChallengeMethod: codeChallengeMethod,
|
||||
},
|
||||
HMACKey: storage.NewHMACKey(crypto.SHA256),
|
||||
}, nil
|
||||
}, idTokenHintSubject, nil
|
||||
}
|
||||
|
||||
func parseCrossClientScope(scope string) (peerID string, ok bool) {
|
||||
|
||||
+190
-1
@@ -3,13 +3,17 @@ package server
|
||||
import (
|
||||
"crypto/rand"
|
||||
"crypto/rsa"
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/go-jose/go-jose/v4"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/dexidp/dex/server/signer"
|
||||
@@ -432,7 +436,7 @@ func TestParseAuthorizationRequest(t *testing.T) {
|
||||
req = httptest.NewRequest("GET", httpServer.URL+"/auth?"+params.Encode(), nil)
|
||||
}
|
||||
|
||||
_, err := server.parseAuthorizationRequest(req)
|
||||
_, _, err := server.parseAuthorizationRequest(req)
|
||||
if tc.expectedError == nil {
|
||||
if err != nil {
|
||||
t.Errorf("%s: expected no error", tc.name)
|
||||
@@ -883,3 +887,188 @@ func TestRedirectedAuthErrHandler(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// signTestIDToken creates a signed JWT with the given claims using the test key.
|
||||
func signTestIDToken(t *testing.T, claims interface{}) string {
|
||||
t.Helper()
|
||||
payload, err := json.Marshal(claims)
|
||||
require.NoError(t, err)
|
||||
|
||||
joseSigner, err := jose.NewSigner(jose.SigningKey{Algorithm: jose.RS256, Key: testKey}, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
jws, err := joseSigner.Sign(payload)
|
||||
require.NoError(t, err)
|
||||
|
||||
token, err := jws.CompactSerialize()
|
||||
require.NoError(t, err)
|
||||
return token
|
||||
}
|
||||
|
||||
func TestValidateIDTokenHint(t *testing.T) {
|
||||
sig, err := signer.NewMockSigner(testKey)
|
||||
require.NoError(t, err)
|
||||
|
||||
issuerURL, err := url.Parse("https://issuer.example.com")
|
||||
require.NoError(t, err)
|
||||
|
||||
s := &Server{
|
||||
signer: sig,
|
||||
issuerURL: *issuerURL,
|
||||
logger: slog.Default(),
|
||||
}
|
||||
|
||||
now := time.Now()
|
||||
|
||||
t.Run("valid hint (not expired)", func(t *testing.T) {
|
||||
token := signTestIDToken(t, idTokenClaims{
|
||||
Issuer: "https://issuer.example.com",
|
||||
Subject: "CgNmb28SA2Jhcg",
|
||||
Expiry: now.Add(1 * time.Hour).Unix(),
|
||||
})
|
||||
sub, err := s.validateIDTokenHint(t.Context(), token)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "CgNmb28SA2Jhcg", sub)
|
||||
})
|
||||
|
||||
t.Run("valid hint (expired)", func(t *testing.T) {
|
||||
token := signTestIDToken(t, idTokenClaims{
|
||||
Issuer: "https://issuer.example.com",
|
||||
Subject: "CgNmb28SA2Jhcg",
|
||||
Expiry: now.Add(-1 * time.Hour).Unix(),
|
||||
})
|
||||
sub, err := s.validateIDTokenHint(t.Context(), token)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "CgNmb28SA2Jhcg", sub)
|
||||
})
|
||||
|
||||
t.Run("invalid signature", func(t *testing.T) {
|
||||
otherKey, err := rsa.GenerateKey(rand.Reader, 2048)
|
||||
require.NoError(t, err)
|
||||
|
||||
payload, err := json.Marshal(idTokenClaims{
|
||||
Issuer: "https://issuer.example.com",
|
||||
Subject: "CgNmb28SA2Jhcg",
|
||||
Expiry: now.Add(1 * time.Hour).Unix(),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
joseSigner, err := jose.NewSigner(jose.SigningKey{Algorithm: jose.RS256, Key: otherKey}, nil)
|
||||
require.NoError(t, err)
|
||||
jws, err := joseSigner.Sign(payload)
|
||||
require.NoError(t, err)
|
||||
token, err := jws.CompactSerialize()
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = s.validateIDTokenHint(t.Context(), token)
|
||||
assert.Error(t, err)
|
||||
})
|
||||
|
||||
t.Run("wrong issuer", func(t *testing.T) {
|
||||
token := signTestIDToken(t, idTokenClaims{
|
||||
Issuer: "https://wrong-issuer.example.com",
|
||||
Subject: "CgNmb28SA2Jhcg",
|
||||
Expiry: now.Add(1 * time.Hour).Unix(),
|
||||
})
|
||||
_, err := s.validateIDTokenHint(t.Context(), token)
|
||||
assert.Error(t, err)
|
||||
})
|
||||
|
||||
t.Run("malformed token", func(t *testing.T) {
|
||||
_, err := s.validateIDTokenHint(t.Context(), "not-a-valid-jwt")
|
||||
assert.Error(t, err)
|
||||
})
|
||||
}
|
||||
|
||||
func TestSessionMatchesHint(t *testing.T) {
|
||||
// genSubject("foo", "bar") == "CgNmb28SA2Jhcg" (from TestGetSubject)
|
||||
assert.True(t, sessionMatchesHint(&storage.AuthSession{UserID: "foo", ConnectorID: "bar"}, "CgNmb28SA2Jhcg"))
|
||||
assert.False(t, sessionMatchesHint(&storage.AuthSession{UserID: "other", ConnectorID: "bar"}, "CgNmb28SA2Jhcg"))
|
||||
assert.False(t, sessionMatchesHint(&storage.AuthSession{UserID: "foo", ConnectorID: "other"}, "CgNmb28SA2Jhcg"))
|
||||
assert.False(t, sessionMatchesHint(nil, "CgNmb28SA2Jhcg"))
|
||||
}
|
||||
|
||||
func TestParseAuthorizationRequest_IDTokenHint(t *testing.T) {
|
||||
sig, err := signer.NewMockSigner(testKey)
|
||||
require.NoError(t, err)
|
||||
|
||||
now := time.Now()
|
||||
|
||||
t.Run("valid id_token_hint populates subject", func(t *testing.T) {
|
||||
httpServer, server := newTestServerMultipleConnectors(t, func(c *Config) {
|
||||
c.SupportedResponseTypes = []string{"code"}
|
||||
c.Storage = storage.WithStaticClients(c.Storage, []storage.Client{
|
||||
{ID: "foo", RedirectURIs: []string{"https://example.com/foo"}},
|
||||
})
|
||||
c.Signer = sig
|
||||
})
|
||||
defer httpServer.Close()
|
||||
|
||||
token := signTestIDToken(t, idTokenClaims{
|
||||
Issuer: httpServer.URL,
|
||||
Subject: "CgNmb28SA2Jhcg",
|
||||
Expiry: now.Add(1 * time.Hour).Unix(),
|
||||
})
|
||||
|
||||
params := url.Values{
|
||||
"client_id": {"foo"},
|
||||
"redirect_uri": {"https://example.com/foo"},
|
||||
"response_type": {"code"},
|
||||
"scope": {"openid"},
|
||||
"id_token_hint": {token},
|
||||
}
|
||||
req := httptest.NewRequest("GET", httpServer.URL+"/auth?"+params.Encode(), nil)
|
||||
|
||||
_, hintSubject, err := server.parseAuthorizationRequest(req)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "CgNmb28SA2Jhcg", hintSubject)
|
||||
})
|
||||
|
||||
t.Run("invalid id_token_hint returns error", func(t *testing.T) {
|
||||
httpServer, server := newTestServerMultipleConnectors(t, func(c *Config) {
|
||||
c.SupportedResponseTypes = []string{"code"}
|
||||
c.Storage = storage.WithStaticClients(c.Storage, []storage.Client{
|
||||
{ID: "foo", RedirectURIs: []string{"https://example.com/foo"}},
|
||||
})
|
||||
c.Signer = sig
|
||||
})
|
||||
defer httpServer.Close()
|
||||
|
||||
params := url.Values{
|
||||
"client_id": {"foo"},
|
||||
"redirect_uri": {"https://example.com/foo"},
|
||||
"response_type": {"code"},
|
||||
"scope": {"openid"},
|
||||
"id_token_hint": {"invalid-token"},
|
||||
}
|
||||
req := httptest.NewRequest("GET", httpServer.URL+"/auth?"+params.Encode(), nil)
|
||||
|
||||
_, _, err := server.parseAuthorizationRequest(req)
|
||||
require.Error(t, err)
|
||||
redirectErr, ok := err.(*redirectedAuthErr)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, errInvalidRequest, redirectErr.Type)
|
||||
})
|
||||
|
||||
t.Run("no id_token_hint leaves subject empty", func(t *testing.T) {
|
||||
httpServer, server := newTestServerMultipleConnectors(t, func(c *Config) {
|
||||
c.SupportedResponseTypes = []string{"code"}
|
||||
c.Storage = storage.WithStaticClients(c.Storage, []storage.Client{
|
||||
{ID: "foo", RedirectURIs: []string{"https://example.com/foo"}},
|
||||
})
|
||||
})
|
||||
defer httpServer.Close()
|
||||
|
||||
params := url.Values{
|
||||
"client_id": {"foo"},
|
||||
"redirect_uri": {"https://example.com/foo"},
|
||||
"response_type": {"code"},
|
||||
"scope": {"openid"},
|
||||
}
|
||||
req := httptest.NewRequest("GET", httpServer.URL+"/auth?"+params.Encode(), nil)
|
||||
|
||||
_, hintSubject, err := server.parseAuthorizationRequest(req)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "", hintSubject)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -238,6 +238,13 @@ func (s *Server) createOrUpdateAuthSession(ctx context.Context, r *http.Request,
|
||||
// Returns ("", false) if session-based login is not possible.
|
||||
func (s *Server) trySessionLogin(ctx context.Context, r *http.Request, w http.ResponseWriter, authReq *storage.AuthRequest) (string, bool) {
|
||||
session := s.getValidAuthSession(ctx, w, r, authReq)
|
||||
return s.trySessionLoginWithSession(ctx, r, w, authReq, session)
|
||||
}
|
||||
|
||||
// trySessionLoginWithSession is like trySessionLogin but accepts a pre-retrieved session.
|
||||
// This allows callers to inspect the session (e.g., for id_token_hint comparison) before
|
||||
// attempting session-based login.
|
||||
func (s *Server) trySessionLoginWithSession(ctx context.Context, r *http.Request, w http.ResponseWriter, authReq *storage.AuthRequest, session *storage.AuthSession) (string, bool) {
|
||||
if session == nil {
|
||||
return "", false
|
||||
}
|
||||
|
||||
@@ -732,6 +732,83 @@ func TestTrySessionLogin_MaxAge(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestTrySessionLoginWithSession_IDTokenHint(t *testing.T) {
|
||||
ctx := t.Context()
|
||||
|
||||
// genSubject("user-1", "mock") produces a deterministic subject string.
|
||||
hintSubjectForUser1Mock, err := genSubject("user-1", "mock")
|
||||
require.NoError(t, err)
|
||||
|
||||
hintSubjectOther, err := genSubject("other-user", "mock")
|
||||
require.NoError(t, err)
|
||||
|
||||
t.Run("hint matches session user - session login succeeds", func(t *testing.T) {
|
||||
s := newTestSessionServer(t)
|
||||
s.skipApproval = true
|
||||
authReq := setupSessionLoginFixture(t, s)
|
||||
|
||||
session := s.getValidAuthSession(ctx, httptest.NewRecorder(), sessionCookieRequest("user-1", "mock", "test-nonce"), &authReq)
|
||||
require.NotNil(t, session)
|
||||
|
||||
// Verify hint matches.
|
||||
assert.True(t, sessionMatchesHint(session, hintSubjectForUser1Mock))
|
||||
|
||||
r := sessionCookieRequest("user-1", "mock", "test-nonce")
|
||||
w := httptest.NewRecorder()
|
||||
|
||||
_, ok := s.trySessionLoginWithSession(ctx, r, w, &authReq, session)
|
||||
assert.True(t, ok)
|
||||
})
|
||||
|
||||
t.Run("hint does not match session user - session invalidated", func(t *testing.T) {
|
||||
s := newTestSessionServer(t)
|
||||
s.skipApproval = true
|
||||
authReq := setupSessionLoginFixture(t, s)
|
||||
|
||||
session := s.getValidAuthSession(ctx, httptest.NewRecorder(), sessionCookieRequest("user-1", "mock", "test-nonce"), &authReq)
|
||||
require.NotNil(t, session)
|
||||
|
||||
// Verify hint does NOT match.
|
||||
assert.False(t, sessionMatchesHint(session, hintSubjectOther))
|
||||
|
||||
// Simulating the hint mismatch logic from handleConnectorLogin:
|
||||
// when hint doesn't match and prompt is not none, session is set to nil.
|
||||
var nilSession *storage.AuthSession
|
||||
r := sessionCookieRequest("user-1", "mock", "test-nonce")
|
||||
w := httptest.NewRecorder()
|
||||
|
||||
_, ok := s.trySessionLoginWithSession(ctx, r, w, &authReq, nilSession)
|
||||
assert.False(t, ok, "session login should fail when session is invalidated due to hint mismatch")
|
||||
})
|
||||
|
||||
t.Run("hint with no session - trySessionLoginWithSession returns false", func(t *testing.T) {
|
||||
s := newTestSessionServer(t)
|
||||
s.skipApproval = true
|
||||
authReq := setupSessionLoginFixture(t, s)
|
||||
|
||||
r := httptest.NewRequest(http.MethodGet, "/", nil)
|
||||
w := httptest.NewRecorder()
|
||||
|
||||
_, ok := s.trySessionLoginWithSession(ctx, r, w, &authReq, nil)
|
||||
assert.False(t, ok)
|
||||
})
|
||||
|
||||
t.Run("no hint - unchanged behavior", func(t *testing.T) {
|
||||
s := newTestSessionServer(t)
|
||||
s.skipApproval = true
|
||||
authReq := setupSessionLoginFixture(t, s)
|
||||
|
||||
session := s.getValidAuthSession(ctx, httptest.NewRecorder(), sessionCookieRequest("user-1", "mock", "test-nonce"), &authReq)
|
||||
require.NotNil(t, session)
|
||||
|
||||
r := sessionCookieRequest("user-1", "mock", "test-nonce")
|
||||
w := httptest.NewRecorder()
|
||||
|
||||
_, ok := s.trySessionLoginWithSession(ctx, r, w, &authReq, session)
|
||||
assert.True(t, ok)
|
||||
})
|
||||
}
|
||||
|
||||
func TestParseAuthRequest_PromptAndMaxAge(t *testing.T) {
|
||||
t.Run("prompt=consent sets ForceApprovalPrompt", func(t *testing.T) {
|
||||
authReq := storage.AuthRequest{
|
||||
|
||||
Reference in New Issue
Block a user