Files

227 lines
6.8 KiB
Go

package oidc
import (
"context"
"database/sql"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"net/url"
"strings"
"time"
"github.com/golang-jwt/jwt/v5"
"trankilou.fr/lassistanoque/backend/internal/domain"
"trankilou.fr/lassistanoque/backend/internal/service/auth"
)
var (
ErrRegistrationUnsupported = errors.New("registration via OIDC is not supported: users are provisioned at login")
ErrNoEmail = errors.New("id token does not contain an email claim")
ErrUserDisabled = errors.New("user is disabled")
)
// tokenResponse est la réponse de l'endpoint token du fournisseur.
type tokenResponse struct {
AccessToken string `json:"access_token"`
IDToken string `json:"id_token"`
TokenType string `json:"token_type"`
ExpiresIn int `json:"expires_in"`
}
// idTokenClaims sont les claims du token ID émis par le fournisseur.
type idTokenClaims struct {
Email string `json:"email"`
GivenName string `json:"given_name"`
FamilyName string `json:"family_name"`
Name string `json:"name"`
PreferredUsername string `json:"preferred_username"`
Picture string `json:"picture"`
jwt.RegisteredClaims
}
type OIDCAuthenticator struct {
tokenManager auth.TokenManager
oidcRepository domain.OIDCProviderRepository
userRepository domain.UserRepository
discovery *discovery
jwks *jwksStore
httpClient *http.Client
}
func NewOIDCAuthenticator(
tokenManager auth.TokenManager,
oidcRepository domain.OIDCProviderRepository,
userRepository domain.UserRepository,
) *OIDCAuthenticator {
return &OIDCAuthenticator{
tokenManager: tokenManager,
oidcRepository: oidcRepository,
userRepository: userRepository,
discovery: newDiscovery(),
jwks: newJwksStore(),
httpClient: &http.Client{Timeout: 10 * time.Second},
}
}
// AuthorizeURL construit l'URL d'autorisation vers laquelle rediriger l'utilisateur.
func (s *OIDCAuthenticator) AuthorizeURL(ctx context.Context, provider *domain.OIDCProvider, redirectURI string, state string) (string, error) {
metadata, err := s.discovery.Metadata(provider)
if err != nil {
return "", err
}
params := url.Values{}
params.Set("client_id", provider.ClientID)
params.Set("redirect_uri", redirectURI)
params.Set("response_type", "code")
params.Set("scope", "openid profile email")
params.Set("state", state)
return metadata.AuthorizationEndpoint + "?" + params.Encode(), nil
}
// ExchangeCode échange le code d'autorisation contre un token ID, le valide,
// puis délivre une session applicative.
func (s *OIDCAuthenticator) ExchangeCode(ctx context.Context, provider *domain.OIDCProvider, code string, redirectURI string) (*auth.Session, error) {
metadata, err := s.discovery.Metadata(provider)
if err != nil {
return nil, err
}
tokens, err := s.requestTokens(metadata.TokenEndpoint, provider, code, redirectURI)
if err != nil {
return nil, err
}
if tokens.IDToken == "" {
return nil, fmt.Errorf("provider %s did not return an id token", provider.ID)
}
claims, err := s.validateIDToken(provider, metadata, tokens.IDToken)
if err != nil {
return nil, err
}
if claims.Email == "" {
return nil, ErrNoEmail
}
user, err := s.provisionUser(claims)
if err != nil {
return nil, err
}
accessToken, expiresAt, err := s.tokenManager.GenerateAccessToken(user)
if err != nil {
return nil, fmt.Errorf("unable to generate access token: %s", err)
}
refreshToken, _, err := s.tokenManager.GenerateRefreshToken(user.ID)
if err != nil {
return nil, fmt.Errorf("unable to generate refresh token: %s", err)
}
return &auth.Session{
User: user,
AccessToken: accessToken,
RefreshToken: refreshToken,
ExpiresAt: expiresAt,
}, nil
}
// Authenticate implémente auth.Authenticator: le flux OIDC passe par le state
// signé qui transporte l'id du fournisseur, puis l'échange du code.
func (s *OIDCAuthenticator) Authenticate(ctx context.Context, creds auth.Credentials) (*auth.Session, error) {
providerID, err := s.tokenManager.ParseStateToken(creds.OIDCState)
if err != nil {
return nil, fmt.Errorf("invalid oidc state: %s", err)
}
provider, err := s.oidcRepository.FindOIDCProvider(providerID)
if err != nil {
return nil, fmt.Errorf("unknown oidc provider: %s", providerID)
}
return s.ExchangeCode(ctx, provider, creds.OIDCCode, creds.RedirectURI)
}
func (s *OIDCAuthenticator) Register(ctx context.Context, registration auth.Registration) (*auth.Session, error) {
return nil, ErrRegistrationUnsupported
}
func (s *OIDCAuthenticator) requestTokens(tokenEndpoint string, provider *domain.OIDCProvider, code string, redirectURI string) (*tokenResponse, error) {
form := url.Values{}
form.Set("grant_type", "authorization_code")
form.Set("code", code)
form.Set("redirect_uri", redirectURI)
form.Set("client_id", provider.ClientID)
form.Set("client_secret", provider.ClientSecret)
res, err := s.httpClient.Post(tokenEndpoint, "application/x-www-form-urlencoded", strings.NewReader(form.Encode()))
if err != nil {
return nil, fmt.Errorf("error calling token endpoint: %s", err)
}
defer res.Body.Close()
body, err := io.ReadAll(res.Body)
if err != nil {
return nil, err
}
if res.StatusCode != http.StatusOK {
return nil, fmt.Errorf("token endpoint returned status %d: %s", res.StatusCode, string(body))
}
tokens := &tokenResponse{}
if err := json.Unmarshal(body, tokens); err != nil {
return nil, err
}
return tokens, nil
}
func (s *OIDCAuthenticator) validateIDToken(provider *domain.OIDCProvider, metadata *providerMetadata, idToken string) (*idTokenClaims, error) {
claims := &idTokenClaims{}
_, err := jwt.ParseWithClaims(idToken, claims, s.jwks.KeyFunc(provider.ID, metadata.JwksURI),
jwt.WithValidMethods([]string{"RS256", "RS384", "RS512", "ES256", "ES384", "ES521"}),
jwt.WithIssuer(metadata.Issuer),
jwt.WithAudience(provider.ClientID),
)
if err != nil {
return nil, fmt.Errorf("invalid id token: %s", err)
}
return claims, nil
}
func (s *OIDCAuthenticator) provisionUser(claims *idTokenClaims) (*domain.User, error) {
user, err := s.userRepository.FindUserByEmail(claims.Email)
if err == nil {
if !user.Enabled {
return nil, ErrUserDisabled
}
return user, nil
}
if !errors.Is(err, sql.ErrNoRows) {
return nil, err
}
firstname := claims.GivenName
if firstname == "" {
firstname = claims.PreferredUsername
}
user, err = s.userRepository.CreateUser(&domain.User{
Email: claims.Email,
Firstname: firstname,
Lastname: claims.FamilyName,
Enabled: true,
})
if err != nil {
return nil, fmt.Errorf("error creating oidc user %s: %s", claims.Email, err)
}
_, err = s.userRepository.CreateTeam(user.ID, &domain.Team{
Label: "Espace personnel",
DateCreated: time.Now(),
})
if err != nil {
return nil, fmt.Errorf("error creating personal team for oidc user %s: %s", claims.Email, err)
}
return user, nil
}