diff --git a/consent/strategy_default.go b/consent/strategy_default.go index 768f9dd399..097d912cc2 100644 --- a/consent/strategy_default.go +++ b/consent/strategy_default.go @@ -637,7 +637,10 @@ func (s *defaultStrategy) verifyConsent(ctx context.Context, _ http.ResponseWrit if f.ConsentError.IsError() { f.ConsentError.SetDefaults(flow.ConsentRequestDeniedErrorName) - return nil, errors.WithStack(f.ConsentError.ToRFCError()) + // Return the flow alongside the error so device-flow callers can identify + // a denied consent and mark the device code session as rejected. + // The auth-code caller discards the flow on error, so this is safe. + return f, errors.WithStack(f.ConsentError.ToRFCError()) } if err := s.r.ConsentManager().CreateConsentSession(ctx, f); errors.Is(err, sqlcon.ErrUniqueViolation()) { @@ -1343,3 +1346,4 @@ func (s *defaultStrategy) verifyDevice(ctx context.Context, _ http.ResponseWrite func (s *defaultStrategy) getDeviceVerificationPath(ctx context.Context) *url.URL { return urlx.AppendPaths(s.r.Config().PublicURL(ctx), deviceVerificationPath) } + diff --git a/oauth2/handler.go b/oauth2/handler.go index b17ba34478..6225f0ae64 100644 --- a/oauth2/handler.go +++ b/oauth2/handler.go @@ -766,6 +766,23 @@ func (h *Handler) performOAuth2DeviceVerificationFlow(w http.ResponseWriter, r * return } else if err != nil { x.LogError(r, err, h.r.Logger()) + + // If consent for a device authorization flow was denied, propagate the + // rejection to the device session so the polling token client receives + // access_denied (RFC 8628 ยง3.5) instead of authorization_pending until + // the device code expires. HandleOAuth2DeviceAuthorizationRequest returns + // the flow alongside the error when the denial came from verifyConsent. + if f != nil && f.DeviceCodeRequestID.Valid { + if rq, sig, err := h.r.OAuth2Storage().GetDeviceCodeSessionByRequestID(ctx, f.DeviceCodeRequestID.String(), &Session{}); err == nil { + rq.SetUserCodeState(fosite.UserCodeRejected) + if err := h.r.OAuth2Storage().UpdateDeviceCodeSessionBySignature(ctx, sig, rq); err != nil { + x.LogError(r, err, h.r.Logger()) + } + } else { + x.LogError(r, err, h.r.Logger()) + } + } + h.r.Writer().WriteError(w, r, err) return } @@ -1628,3 +1645,4 @@ func (h *Handler) createVerifiableCredential(w http.ResponseWriter, r *http.Requ response.Credential = rawToken h.r.Writer().Write(w, r, &response) } +