mirror of
https://github.com/pocket-id/pocket-id.git
synced 2026-09-30 23:09:05 +02:00
Co-authored-by: Alessandro (Ale) Segala <43508+ItalyPaleAle@users.noreply.github.com> Co-authored-by: Elias Schneider <login@eliasschneider.com>
264 lines
10 KiB
Go
264 lines
10 KiB
Go
package backchannellogout
|
|
|
|
import (
|
|
"context"
|
|
"fmt"
|
|
"log/slog"
|
|
"net/http"
|
|
"net/url"
|
|
"strings"
|
|
"time"
|
|
|
|
"github.com/italypaleale/francis/actor"
|
|
francishost "github.com/italypaleale/francis/host"
|
|
"github.com/italypaleale/francis/host/local"
|
|
"gorm.io/gorm"
|
|
|
|
"github.com/pocket-id/pocket-id/backend/internal/model"
|
|
)
|
|
|
|
// requestTimeout bounds each notification POST, so one unreachable client cannot stall the others
|
|
const requestTimeout = 10 * time.Second
|
|
|
|
// TokenSigner mints the logout tokens delivered to clients
|
|
type TokenSigner interface {
|
|
GenerateLogoutToken(userID string, clientID string) (string, error)
|
|
}
|
|
|
|
// Service sends OIDC Back-Channel Logout 1.0 tokens to clients when a user's access is revoked
|
|
// Deliveries are scheduled as durable jobs, so they survive a restart and failed attempts are retried a capped number of times before giving up
|
|
type Service struct {
|
|
db *gorm.DB
|
|
tokenSigner TokenSigner
|
|
httpClient *http.Client
|
|
actors *actor.Service
|
|
}
|
|
|
|
func NewService(db *gorm.DB, tokenSigner TokenSigner, httpClient *http.Client, actorsHost francishost.Host) (*Service, error) {
|
|
s := &Service{
|
|
db: db,
|
|
tokenSigner: tokenSigner,
|
|
httpClient: newHTTPClient(httpClient),
|
|
actors: actorsHost.Service(),
|
|
}
|
|
|
|
err := actorsHost.RegisterActor(
|
|
ActorType,
|
|
s.newNotifierActor,
|
|
local.WithCapacityGroup(ActorType, deliveryConcurrency),
|
|
local.WithMaxAttempts(deliveryMaxAttempts),
|
|
)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("error registering the %s actor: %w", ActorType, err)
|
|
}
|
|
|
|
return s, nil
|
|
}
|
|
|
|
// newHTTPClient refuses to follow redirects, as Go would turn the POST into a body-less GET and the logout token would be silently dropped
|
|
// Returning the redirect response instead makes the delivery fail loudly on the status check
|
|
func newHTTPClient(source *http.Client) *http.Client {
|
|
if source == nil {
|
|
source = http.DefaultClient
|
|
}
|
|
|
|
return &http.Client{
|
|
Transport: source.Transport,
|
|
CheckRedirect: func(_ *http.Request, _ []*http.Request) error {
|
|
return http.ErrUseLastResponse
|
|
},
|
|
}
|
|
}
|
|
|
|
// target is a single client to notify that a user's session should end
|
|
type target struct {
|
|
UserID string
|
|
ClientID string
|
|
LogoutURL string
|
|
}
|
|
|
|
// targetsQuery selects the authorizations of clients that are registered for back-channel logout
|
|
// Callers narrow it down to the users or the client whose access was revoked
|
|
func (s *Service) targetsQuery(ctx context.Context, tx *gorm.DB) *gorm.DB {
|
|
return tx.
|
|
WithContext(ctx).
|
|
Model(&model.UserAuthorizedOidcClient{}).
|
|
Select("user_authorized_oidc_clients.user_id", "user_authorized_oidc_clients.client_id", "oidc_clients.backchannel_logout_url AS logout_url").
|
|
Joins("JOIN oidc_clients ON oidc_clients.id = user_authorized_oidc_clients.client_id").
|
|
Where("oidc_clients.backchannel_logout_url <> ''")
|
|
}
|
|
|
|
// targetsForUsers returns every client the given users have authorized that is registered for back-channel logout
|
|
func (s *Service) targetsForUsers(ctx context.Context, tx *gorm.DB, userIDs []string) ([]target, error) {
|
|
if len(userIDs) == 0 {
|
|
return nil, nil
|
|
}
|
|
|
|
var targets []target
|
|
err := s.targetsQuery(ctx, tx).
|
|
Where("user_authorized_oidc_clients.user_id IN (?)", userIDs).
|
|
Scan(&targets).
|
|
Error
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return targets, nil
|
|
}
|
|
|
|
// targetsForAuthorization returns the client to notify when a single authorization is revoked, and nothing when that client is not registered for back-channel logout
|
|
func (s *Service) targetsForAuthorization(ctx context.Context, tx *gorm.DB, userID string, clientID string) ([]target, error) {
|
|
var targets []target
|
|
err := s.targetsQuery(ctx, tx).
|
|
Where("user_authorized_oidc_clients.user_id = ?", userID).
|
|
Where("user_authorized_oidc_clients.client_id = ?", clientID).
|
|
Scan(&targets).
|
|
Error
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return targets, nil
|
|
}
|
|
|
|
// targetsForClient returns every user who has authorized the given client, when that client is registered for back-channel logout
|
|
func (s *Service) targetsForClient(ctx context.Context, tx *gorm.DB, clientID string) ([]target, error) {
|
|
var targets []target
|
|
err := s.targetsQuery(ctx, tx).
|
|
Where("user_authorized_oidc_clients.client_id = ?", clientID).
|
|
Scan(&targets).
|
|
Error
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return targets, nil
|
|
}
|
|
|
|
// targetsForLostGroupAccess returns clients registered for back-channel logout that the matched users have authorized but can no longer access because of the client's group restriction
|
|
// It must run after the group membership or allowed-group change has been committed
|
|
func (s *Service) targetsForLostGroupAccess(ctx context.Context, tx *gorm.DB, userIDs []string, clientID string) ([]target, error) {
|
|
// Require at least one filter, so a caller that computes an empty user list can never match every user of every client
|
|
if len(userIDs) == 0 && clientID == "" {
|
|
return nil, nil
|
|
}
|
|
|
|
query := s.targetsQuery(ctx, tx).
|
|
Where("oidc_clients.is_group_restricted = ?", true).
|
|
Where("NOT EXISTS (SELECT 1 FROM oidc_clients_allowed_user_groups ag JOIN user_groups_users ugu ON ugu.user_group_id = ag.user_group_id WHERE ag.oidc_client_id = oidc_clients.id AND ugu.user_id = user_authorized_oidc_clients.user_id)")
|
|
if len(userIDs) > 0 {
|
|
query = query.Where("user_authorized_oidc_clients.user_id IN (?)", userIDs)
|
|
}
|
|
if clientID != "" {
|
|
query = query.Where("user_authorized_oidc_clients.client_id = ?", clientID)
|
|
}
|
|
|
|
var targets []target
|
|
err := query.Scan(&targets).Error
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return targets, nil
|
|
}
|
|
|
|
// PrepareUserNotifications resolves, within the given transaction, the logout notifications for users whose access is being revoked
|
|
// It exists for callers that delete the users or their authorizations, which are gone once the transaction commits
|
|
// The returned function delivers the notifications in the background and must only be called after the transaction has committed
|
|
// It is never nil, so callers that treat a failed lookup as non-fatal can call it unconditionally
|
|
func (s *Service) PrepareUserNotifications(ctx context.Context, tx *gorm.DB, userIDs []string) (func(), error) {
|
|
targets, err := s.targetsForUsers(ctx, tx, userIDs)
|
|
if err != nil {
|
|
return func() {}, err
|
|
}
|
|
return func() { s.notifyClients(ctx, targets) }, nil
|
|
}
|
|
|
|
// PrepareAuthorizationNotification resolves, within the given transaction, the logout notification for a single authorization that is being revoked
|
|
// The returned function behaves like the one from PrepareUserNotifications
|
|
func (s *Service) PrepareAuthorizationNotification(ctx context.Context, tx *gorm.DB, userID string, clientID string) (func(), error) {
|
|
targets, err := s.targetsForAuthorization(ctx, tx, userID, clientID)
|
|
if err != nil {
|
|
return func() {}, err
|
|
}
|
|
return func() { s.notifyClients(ctx, targets) }, nil
|
|
}
|
|
|
|
// PrepareClientNotifications resolves, within the given transaction, the logout notifications for every user of a client that is being deleted
|
|
// The returned function behaves like the one from PrepareUserNotifications
|
|
func (s *Service) PrepareClientNotifications(ctx context.Context, tx *gorm.DB, clientID string) (func(), error) {
|
|
targets, err := s.targetsForClient(ctx, tx, clientID)
|
|
if err != nil {
|
|
return func() {}, err
|
|
}
|
|
return func() { s.notifyClients(ctx, targets) }, nil
|
|
}
|
|
|
|
// NotifyUser delivers logout tokens to every client the user has authorized that is registered for back-channel logout
|
|
// It must be called after the change that revoked the user's access has been committed, and logs instead of failing because delivery is best effort
|
|
func (s *Service) NotifyUser(ctx context.Context, userID string) {
|
|
targets, err := s.targetsForUsers(ctx, s.db, []string{userID})
|
|
if err != nil {
|
|
slog.ErrorContext(ctx, "Failed to find clients to notify for back-channel logout", slog.String("userId", userID), slog.Any("error", err))
|
|
return
|
|
}
|
|
s.notifyClients(ctx, targets)
|
|
}
|
|
|
|
// NotifyLostGroupAccess delivers logout tokens for group-restricted clients that the matched users can no longer access
|
|
// Callers pass the users whose membership changed, the client whose allowed groups changed, or both to narrow the match
|
|
// It must be called after the group change has been committed, and logs instead of failing because delivery is best effort
|
|
func (s *Service) NotifyLostGroupAccess(ctx context.Context, userIDs []string, clientID string) {
|
|
targets, err := s.targetsForLostGroupAccess(ctx, s.db, userIDs, clientID)
|
|
if err != nil {
|
|
slog.ErrorContext(ctx, "Failed to find clients to notify for back-channel logout", slog.Any("error", err))
|
|
return
|
|
}
|
|
s.notifyClients(ctx, targets)
|
|
}
|
|
|
|
// notifyClients schedules a durable delivery job for each of the given clients, so callers are never blocked on slow or unreachable clients
|
|
// It must be called after the change that revoked the user's access has been committed
|
|
func (s *Service) notifyClients(ctx context.Context, targets []target) {
|
|
for _, t := range targets {
|
|
// One actor per authorization serializes its deliveries, while the capacity group limits running jobs without counting idle actors
|
|
actorID := t.ClientID + ":" + t.UserID
|
|
_, _, err := s.actors.Dispatch(ctx, ActorType, actorID, methodDeliver, t)
|
|
if err != nil {
|
|
slog.ErrorContext(ctx, "Failed to schedule back-channel logout notification",
|
|
slog.String("clientId", t.ClientID),
|
|
slog.String("userId", t.UserID),
|
|
slog.Any("error", err),
|
|
)
|
|
}
|
|
}
|
|
}
|
|
|
|
func (s *Service) sendLogoutToken(parentCtx context.Context, t target) error {
|
|
logoutToken, err := s.tokenSigner.GenerateLogoutToken(t.UserID, t.ClientID)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
ctx, cancel := context.WithTimeout(parentCtx, requestTimeout)
|
|
defer cancel()
|
|
|
|
body := url.Values{"logout_token": []string{logoutToken}}
|
|
req, err := http.NewRequestWithContext(ctx, http.MethodPost, t.LogoutURL, strings.NewReader(body.Encode()))
|
|
if err != nil {
|
|
return fmt.Errorf("%w: invalid logout request: %w", actor.ErrJobPermanentFailure, err)
|
|
}
|
|
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
|
|
|
res, err := s.httpClient.Do(req)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer res.Body.Close()
|
|
|
|
if res.StatusCode < 200 || res.StatusCode > 299 {
|
|
// Only retry responses that indicate a potentially temporary failure at the client
|
|
if res.StatusCode != http.StatusRequestTimeout && res.StatusCode != http.StatusTooManyRequests && res.StatusCode < http.StatusInternalServerError {
|
|
return fmt.Errorf("%w: client responded with status %d", actor.ErrJobPermanentFailure, res.StatusCode)
|
|
}
|
|
return fmt.Errorf("client responded with status %d", res.StatusCode)
|
|
}
|
|
return nil
|
|
}
|