From 23084a72a25c9b121cf55a2cbfeb5f0b9cff5de7 Mon Sep 17 00:00:00 2001 From: Bryan Frimin Date: Fri, 20 Mar 2026 16:08:15 +0100 Subject: [PATCH] Add OIDC login support for Google and Microsoft providers Implements OpenID Connect authentication flow with PKCE, JWT verification, and enterprise-only account restrictions. Adds OIDC service with JWKS caching and state management, HTTP handlers for login/callback flows, GraphQL query for available providers, and sign-in UI integration. Signed-off-by: Bryan Frimin --- .../src/pages/iam/auth/sign-in/SignInPage.tsx | 50 ++ pkg/bootstrap/builder.go | 10 + pkg/coredata/migrations/20260319T150000Z.sql | 19 + pkg/coredata/oidc_state.go | 137 ++++ pkg/coredata/session.go | 2 + pkg/iam/auth_service.go | 23 + pkg/iam/oidc/errors.go | 99 +++ pkg/iam/oidc/gc.go | 89 +++ pkg/iam/oidc/service.go | 686 ++++++++++++++++++ pkg/iam/service.go | 13 + pkg/probod/auth_config.go | 16 +- pkg/probod/oidc_config.go | 21 + pkg/probod/probod.go | 11 + pkg/server/api/connect/v1/oidc_handler.go | 161 ++++ pkg/server/api/connect/v1/resolver.go | 4 + pkg/server/api/connect/v1/schema.graphql | 8 + pkg/server/api/connect/v1/v1_resolver.go | 16 + 17 files changed, 1358 insertions(+), 7 deletions(-) create mode 100644 pkg/coredata/migrations/20260319T150000Z.sql create mode 100644 pkg/coredata/oidc_state.go create mode 100644 pkg/iam/oidc/errors.go create mode 100644 pkg/iam/oidc/gc.go create mode 100644 pkg/iam/oidc/service.go create mode 100644 pkg/probod/oidc_config.go create mode 100644 pkg/server/api/connect/v1/oidc_handler.go diff --git a/apps/console/src/pages/iam/auth/sign-in/SignInPage.tsx b/apps/console/src/pages/iam/auth/sign-in/SignInPage.tsx index f3bbdea82..909f398e2 100644 --- a/apps/console/src/pages/iam/auth/sign-in/SignInPage.tsx +++ b/apps/console/src/pages/iam/auth/sign-in/SignInPage.tsx @@ -1,6 +1,52 @@ import { useTranslate } from "@probo/i18n"; import { Button } from "@probo/ui"; +import { Suspense } from "react"; +import { useLazyLoadQuery } from "react-relay"; import { Link, useLocation } from "react-router"; +import { graphql } from "relay-runtime"; + +import type { SignInPageQuery } from "#/__generated__/iam/SignInPageQuery.graphql"; +import { useSafeContinueUrl } from "#/hooks/useSafeContinueUrl"; + +const oidcProvidersQuery = graphql` + query SignInPageQuery { + oidcProviders { + name + loginURL + } + } +`; + +function OIDCButtons() { + const { __ } = useTranslate(); + const safeContinueUrl = useSafeContinueUrl(); + + const data = useLazyLoadQuery(oidcProvidersQuery, {}); + + if (data.oidcProviders.length === 0) { + return null; + } + + return ( + <> + {data.oidcProviders.map((provider) => ( + + ))} + + ); +} export default function SignInPage() { const { __ } = useTranslate(); @@ -23,6 +69,10 @@ export default function SignInPage() { {__("Login with Email")} + + + +
. +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package coredata + +import ( + "context" + "errors" + "fmt" + "time" + + "github.com/jackc/pgx/v5" + "go.gearno.de/kit/pg" +) + +type ( + OIDCProvider string + + OIDCState struct { + ID string `db:"id"` + Provider OIDCProvider `db:"provider"` + Nonce string `db:"nonce"` + CodeVerifier string `db:"code_verifier"` + ContinueURL string `db:"continue_url"` + CreatedAt time.Time `db:"created_at"` + ExpiresAt time.Time `db:"expires_at"` + } +) + +const ( + OIDCProviderGoogle OIDCProvider = "GOOGLE" + OIDCProviderMicrosoft OIDCProvider = "MICROSOFT" +) + +func (p OIDCProvider) IsValid() bool { + switch p { + case OIDCProviderGoogle, OIDCProviderMicrosoft: + return true + } + return false +} + +func (p OIDCProvider) String() string { return string(p) } + +func (p *OIDCProvider) UnmarshalText(text []byte) error { + *p = OIDCProvider(text) + if !p.IsValid() { + return fmt.Errorf("%s is not a valid OIDCProvider", string(text)) + } + return nil +} + +func (p OIDCProvider) MarshalText() ([]byte, error) { + return []byte(p.String()), nil +} + +func (s *OIDCState) Insert(ctx context.Context, conn pg.Conn) error { + query := ` +INSERT INTO iam_oidc_states (id, provider, nonce, code_verifier, continue_url, created_at, expires_at) +VALUES (@id, @provider, @nonce, @code_verifier, @continue_url, @created_at, @expires_at) +` + + args := pgx.StrictNamedArgs{ + "id": s.ID, + "provider": s.Provider, + "nonce": s.Nonce, + "code_verifier": s.CodeVerifier, + "continue_url": s.ContinueURL, + "created_at": s.CreatedAt, + "expires_at": s.ExpiresAt, + } + + _, err := conn.Exec(ctx, query, args) + if err != nil { + return fmt.Errorf("cannot insert oidc_state: %w", err) + } + + return nil +} + +func (s *OIDCState) LoadByIDForUpdate(ctx context.Context, conn pg.Conn, id string) error { + query := ` +SELECT id, provider, nonce, code_verifier, continue_url, created_at, expires_at +FROM iam_oidc_states +WHERE id = @id +FOR UPDATE +` + + rows, err := conn.Query(ctx, query, pgx.StrictNamedArgs{"id": id}) + if err != nil { + return fmt.Errorf("cannot query oidc_state: %w", err) + } + + state, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[OIDCState]) + if err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return ErrResourceNotFound + } + return fmt.Errorf("cannot collect oidc_state: %w", err) + } + + *s = state + return nil +} + +func (s *OIDCState) Delete(ctx context.Context, conn pg.Conn) error { + query := `DELETE FROM iam_oidc_states WHERE id = @id` + + _, err := conn.Exec(ctx, query, pgx.StrictNamedArgs{"id": s.ID}) + if err != nil { + return fmt.Errorf("cannot delete oidc_state: %w", err) + } + + return nil +} + +func DeleteExpiredOIDCStates(ctx context.Context, conn pg.Conn, now time.Time) (int64, error) { + query := `DELETE FROM iam_oidc_states WHERE expires_at < @now` + + result, err := conn.Exec(ctx, query, pgx.StrictNamedArgs{"now": now}) + if err != nil { + return 0, fmt.Errorf("cannot delete expired oidc_states: %w", err) + } + + return result.RowsAffected(), nil +} diff --git a/pkg/coredata/session.go b/pkg/coredata/session.go index 09f97746d..f0c392f6c 100644 --- a/pkg/coredata/session.go +++ b/pkg/coredata/session.go @@ -57,6 +57,8 @@ const ( AuthMethodMagicLink AuthMethod = "MAGIC_LINK" AuthMethodPassword AuthMethod = "PASSWORD" AuthMethodSAML AuthMethod = "SAML" + AuthMethodGoogle AuthMethod = "GOOGLE" + AuthMethodMicrosoft AuthMethod = "MICROSOFT" ) func NewRootSession(identityID gid.GID, method AuthMethod, duration time.Duration) *Session { diff --git a/pkg/iam/auth_service.go b/pkg/iam/auth_service.go index 7152d2bc7..8c68c2575 100644 --- a/pkg/iam/auth_service.go +++ b/pkg/iam/auth_service.go @@ -444,6 +444,29 @@ func (s AuthService) OpenSessionWithSAML(ctx context.Context, identityID gid.GID return session, nil } +func (s AuthService) OpenSessionWithOIDC(ctx context.Context, identityID gid.GID, authMethod coredata.AuthMethod) (*coredata.Session, error) { + session := &coredata.Session{} + + err := s.pg.WithTx( + ctx, + func(conn pg.Conn) (err error) { + session = coredata.NewRootSession(identityID, authMethod, s.sessionDuration) + err = session.Insert(ctx, conn) + if err != nil { + return fmt.Errorf("cannot insert session: %w", err) + } + + return nil + }, + ) + + if err != nil { + return nil, err + } + + return session, nil +} + func (s AuthService) CheckCredentials( ctx context.Context, email mail.Addr, diff --git a/pkg/iam/oidc/errors.go b/pkg/iam/oidc/errors.go new file mode 100644 index 000000000..7a8c75aa4 --- /dev/null +++ b/pkg/iam/oidc/errors.go @@ -0,0 +1,99 @@ +// Copyright (c) 2025 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package oidc + +import ( + "fmt" + + "go.probo.inc/probo/pkg/coredata" +) + +type ErrProviderNotEnabled struct { + Provider coredata.OIDCProvider +} + +func NewProviderNotEnabledError(provider coredata.OIDCProvider) error { + return &ErrProviderNotEnabled{Provider: provider} +} + +func (e ErrProviderNotEnabled) Error() string { + return fmt.Sprintf("cannot authenticate: OIDC provider %q is not enabled", e.Provider) +} + +type ErrInvalidState struct{} + +func NewInvalidStateError() error { + return &ErrInvalidState{} +} + +func (e ErrInvalidState) Error() string { + return "cannot validate OIDC state: invalid or expired" +} + +type ErrCodeExchange struct { + Err error +} + +func NewCodeExchangeError(err error) error { + return &ErrCodeExchange{Err: err} +} + +func (e ErrCodeExchange) Error() string { + return fmt.Sprintf("cannot exchange authorization code: %v", e.Err) +} + +func (e ErrCodeExchange) Unwrap() error { + return e.Err +} + +type ErrIDTokenMissing struct{} + +func NewIDTokenMissingError() error { + return &ErrIDTokenMissing{} +} + +func (e ErrIDTokenMissing) Error() string { + return "cannot extract id_token: not present in token response" +} + +type ErrMissingEmailClaim struct{} + +func NewMissingEmailClaimError() error { + return &ErrMissingEmailClaim{} +} + +func (e ErrMissingEmailClaim) Error() string { + return "cannot extract email: claim missing from id token" +} + +type ErrEmailNotVerified struct{} + +func NewEmailNotVerifiedError() error { + return &ErrEmailNotVerified{} +} + +func (e ErrEmailNotVerified) Error() string { + return "cannot authenticate: email address is not verified by the OIDC provider" +} + +type ErrPersonalAccountNotAllowed struct{} + +func NewPersonalAccountNotAllowedError() error { + return &ErrPersonalAccountNotAllowed{} +} + +func (e ErrPersonalAccountNotAllowed) Error() string { + return "cannot authenticate: personal accounts are not allowed, use an enterprise account" +} diff --git a/pkg/iam/oidc/gc.go b/pkg/iam/oidc/gc.go new file mode 100644 index 000000000..38b65d7de --- /dev/null +++ b/pkg/iam/oidc/gc.go @@ -0,0 +1,89 @@ +// Copyright (c) 2025 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package oidc + +import ( + "context" + "fmt" + "time" + + "go.gearno.de/kit/log" + "go.gearno.de/kit/pg" + "go.probo.inc/probo/pkg/coredata" +) + +const ( + DefaultGarbageCollectionInterval = 1 * time.Hour +) + +type GarbageCollector struct { + pg *pg.Client + interval time.Duration + logger *log.Logger +} + +func NewGarbageCollector( + pg *pg.Client, + interval time.Duration, + logger *log.Logger, +) *GarbageCollector { + return &GarbageCollector{ + pg: pg, + interval: interval, + logger: logger.Named("oidc.garbage_collector").With(log.Duration("interval", interval)), + } +} + +func (gc *GarbageCollector) Run(ctx context.Context) error { + gc.logger.InfoCtx(ctx, "oidc garbage collector starting") + + if err := gc.cleanup(ctx); err != nil { + gc.logger.ErrorCtx(ctx, "cannot run initial cleanup", log.Error(err)) + } + + for { + select { + case <-ctx.Done(): + gc.logger.InfoCtx(ctx, "oidc garbage collector shutting down") + return ctx.Err() + case <-time.After(gc.interval): + if err := gc.cleanup(ctx); err != nil { + gc.logger.ErrorCtx(ctx, "cannot run periodic cleanup", log.Error(err)) + } + } + } +} + +func (gc *GarbageCollector) cleanup(ctx context.Context) error { + now := time.Now() + + return gc.pg.WithTx( + ctx, + func(tx pg.Conn) error { + deleted, err := coredata.DeleteExpiredOIDCStates(ctx, tx, now) + if err != nil { + return fmt.Errorf("cannot delete expired oidc states: %w", err) + } + + gc.logger.InfoCtx( + ctx, + "oidc garbage collector cleaned up expired states", + log.Int64("deleted", deleted), + ) + + return nil + }, + ) +} diff --git a/pkg/iam/oidc/service.go b/pkg/iam/oidc/service.go new file mode 100644 index 000000000..38ca481db --- /dev/null +++ b/pkg/iam/oidc/service.go @@ -0,0 +1,686 @@ +// Copyright (c) 2025 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package oidc + +import ( + "context" + "crypto" + "crypto/ecdsa" + "crypto/elliptic" + cryptorand "crypto/rand" + "crypto/rsa" + "crypto/sha256" + "encoding/asn1" + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "io" + "math/big" + "net/http" + "strings" + "sync" + "time" + + "go.gearno.de/kit/log" + "go.gearno.de/kit/pg" + "go.probo.inc/probo/pkg/coredata" + "go.probo.inc/probo/pkg/gid" + "go.probo.inc/probo/pkg/mail" + "golang.org/x/oauth2" +) + +var cryptoRandReader io.Reader = cryptorand.Reader + +type ( + ProviderConfig struct { + ClientID string + ClientSecret string + Enabled bool + } + + providerInfo struct { + oauth2Config oauth2.Config + jwksURL string + issuerValidator func(string) bool + enterpriseChecker func(*idTokenClaims) bool + } + + UserInfo struct { + Email mail.Addr + FullName string + } + + Service struct { + pg *pg.Client + baseURL string + logger *log.Logger + providers map[coredata.OIDCProvider]*providerInfo + + jwksMu sync.RWMutex + jwksCache map[string]*jwksEntry + } + + jwksEntry struct { + keys []jwk + fetchedAt time.Time + } + + jwk struct { + Kty string `json:"kty"` + Kid string `json:"kid"` + Use string `json:"use"` + N string `json:"n"` + E string `json:"e"` + Crv string `json:"crv"` + X string `json:"x"` + Y string `json:"y"` + } + + jwksResponse struct { + Keys []jwk `json:"keys"` + } + + idTokenClaims struct { + Issuer string `json:"iss"` + Subject string `json:"sub"` + Audience any `json:"aud"` + ExpiresAt float64 `json:"exp"` + Nonce string `json:"nonce"` + Email string `json:"email"` + EmailVerified any `json:"email_verified"` + Name string `json:"name"` + HostedDomain string `json:"hd"` + } +) + +func (c *idTokenClaims) hasAudience(clientID string) bool { + switch aud := c.Audience.(type) { + case string: + return aud == clientID + case []any: + for _, v := range aud { + if s, ok := v.(string); ok && s == clientID { + return true + } + } + } + return false +} + +func (c *idTokenClaims) isEmailVerified() bool { + switch v := c.EmailVerified.(type) { + case bool: + return v + case string: + return strings.EqualFold(v, "true") + } + return false +} + +var ( + googleEndpoint = oauth2.Endpoint{ + AuthURL: "https://accounts.google.com/o/oauth2/v2/auth", + TokenURL: "https://oauth2.googleapis.com/token", + } + + microsoftEndpoint = oauth2.Endpoint{ + AuthURL: "https://login.microsoftonline.com/common/oauth2/v2.0/authorize", + TokenURL: "https://login.microsoftonline.com/common/oauth2/v2.0/token", + } +) + +const ( + googleJWKSURL = "https://www.googleapis.com/oauth2/v3/certs" + microsoftJWKSURL = "https://login.microsoftonline.com/common/discovery/v2.0/keys" + microsoftConsumerTenantID = "9188040d-6c67-4c5b-b112-36a304b66dad" + jwksCacheTTL = 1 * time.Hour +) + +func NewService( + pgClient *pg.Client, + baseURL string, + google ProviderConfig, + microsoft ProviderConfig, + logger *log.Logger, +) *Service { + s := &Service{ + pg: pgClient, + baseURL: baseURL, + logger: logger.Named("oidc"), + providers: make(map[coredata.OIDCProvider]*providerInfo), + jwksCache: make(map[string]*jwksEntry), + } + + if google.Enabled { + s.providers[coredata.OIDCProviderGoogle] = &providerInfo{ + oauth2Config: oauth2.Config{ + ClientID: google.ClientID, + ClientSecret: google.ClientSecret, + Endpoint: googleEndpoint, + RedirectURL: baseURL + "/api/connect/v1/oidc/google/callback", + Scopes: []string{"openid", "email", "profile"}, + }, + jwksURL: googleJWKSURL, + issuerValidator: func(iss string) bool { + return iss == "https://accounts.google.com" + }, + enterpriseChecker: func(claims *idTokenClaims) bool { + // The "hd" (hosted domain) claim is only present for + // Google Workspace accounts. Personal gmail.com accounts + // do not have this claim. + return claims.HostedDomain != "" + }, + } + } + + if microsoft.Enabled { + s.providers[coredata.OIDCProviderMicrosoft] = &providerInfo{ + oauth2Config: oauth2.Config{ + ClientID: microsoft.ClientID, + ClientSecret: microsoft.ClientSecret, + Endpoint: microsoftEndpoint, + RedirectURL: baseURL + "/api/connect/v1/oidc/microsoft/callback", + Scopes: []string{"openid", "email", "profile"}, + }, + jwksURL: microsoftJWKSURL, + issuerValidator: func(iss string) bool { + return strings.HasPrefix(iss, "https://login.microsoftonline.com/") && + strings.HasSuffix(iss, "/v2.0") + }, + enterpriseChecker: func(claims *idTokenClaims) bool { + // Personal Microsoft accounts (live.com, outlook.com, + // hotmail.com) use the consumer tenant ID. Reject them. + return !strings.Contains( + claims.Issuer, + microsoftConsumerTenantID, + ) + }, + } + } + + return s +} + +func (s *Service) Run(ctx context.Context) error { + gc := NewGarbageCollector(s.pg, DefaultGarbageCollectionInterval, s.logger) + + gcCtx, stopGC := context.WithCancel(ctx) + defer stopGC() + + errCh := make(chan error, 1) + go func() { + errCh <- gc.Run(gcCtx) + }() + + select { + case <-ctx.Done(): + stopGC() + <-errCh + return ctx.Err() + case err := <-errCh: + if err != nil { + s.logger.ErrorCtx(ctx, "oidc garbage collector failed", log.Error(err)) + return err + } + + return nil + } +} + +func (s *Service) IsProviderEnabled(provider coredata.OIDCProvider) bool { + _, ok := s.providers[provider] + return ok +} + +func (s *Service) EnabledProviders() []coredata.OIDCProvider { + providers := make([]coredata.OIDCProvider, 0, len(s.providers)) + for p := range s.providers { + providers = append(providers, p) + } + return providers +} + +func (s *Service) InitiateLogin( + ctx context.Context, + provider coredata.OIDCProvider, + continueURL string, +) (string, error) { + info, ok := s.providers[provider] + if !ok { + return "", NewProviderNotEnabledError(provider) + } + + state, err := generateRandomString(32) + if err != nil { + return "", fmt.Errorf("cannot generate state: %w", err) + } + + nonce, err := generateRandomString(32) + if err != nil { + return "", fmt.Errorf("cannot generate nonce: %w", err) + } + + codeVerifier, err := generateRandomString(64) + if err != nil { + return "", fmt.Errorf("cannot generate code verifier: %w", err) + } + + now := time.Now() + oidcState := &coredata.OIDCState{ + ID: state, + Provider: provider, + Nonce: nonce, + CodeVerifier: codeVerifier, + ContinueURL: continueURL, + CreatedAt: now, + ExpiresAt: now.Add(10 * time.Minute), + } + + err = s.pg.WithTx( + ctx, + func(tx pg.Conn) error { + if err := oidcState.Insert(ctx, tx); err != nil { + return fmt.Errorf("cannot store oidc state: %w", err) + } + return nil + }, + ) + if err != nil { + return "", err + } + + codeChallenge := computeCodeChallenge(codeVerifier) + + authURL := info.oauth2Config.AuthCodeURL( + state, + oauth2.SetAuthURLParam("nonce", nonce), + oauth2.SetAuthURLParam("code_challenge", codeChallenge), + oauth2.SetAuthURLParam("code_challenge_method", "S256"), + ) + + return authURL, nil +} + +func (s *Service) HandleCallback( + ctx context.Context, + provider coredata.OIDCProvider, + stateParam string, + code string, +) (*coredata.Identity, string, error) { + info, ok := s.providers[provider] + if !ok { + return nil, "", NewProviderNotEnabledError(provider) + } + + var oidcState coredata.OIDCState + err := s.pg.WithTx( + ctx, + func(tx pg.Conn) error { + if err := oidcState.LoadByIDForUpdate(ctx, tx, stateParam); err != nil { + if errors.Is(err, coredata.ErrResourceNotFound) { + return NewInvalidStateError() + } + return fmt.Errorf("cannot load oidc state: %w", err) + } + + if err := oidcState.Delete(ctx, tx); err != nil { + return fmt.Errorf("cannot delete oidc state: %w", err) + } + + return nil + }, + ) + if err != nil { + return nil, "", err + } + + if time.Now().After(oidcState.ExpiresAt) { + return nil, "", NewInvalidStateError() + } + + if oidcState.Provider != provider { + return nil, "", NewInvalidStateError() + } + + token, err := info.oauth2Config.Exchange( + ctx, + code, + oauth2.SetAuthURLParam("code_verifier", oidcState.CodeVerifier), + ) + if err != nil { + return nil, "", NewCodeExchangeError(err) + } + + rawIDToken, ok := token.Extra("id_token").(string) + if !ok { + return nil, "", NewIDTokenMissingError() + } + + claims, err := s.verifyAndParseIDToken(ctx, info, rawIDToken, oidcState.Nonce) + if err != nil { + return nil, "", fmt.Errorf("cannot verify id token: %w", err) + } + + if claims.Email == "" { + return nil, "", NewMissingEmailClaimError() + } + + if !claims.isEmailVerified() { + return nil, "", NewEmailNotVerifiedError() + } + + if !info.enterpriseChecker(claims) { + return nil, "", NewPersonalAccountNotAllowedError() + } + + email, err := mail.ParseAddr(claims.Email) + if err != nil { + return nil, "", fmt.Errorf("cannot parse email from id token: %w", err) + } + + var identity *coredata.Identity + now := time.Now() + + err = s.pg.WithTx( + ctx, + func(tx pg.Conn) error { + identity = &coredata.Identity{} + err := identity.LoadByEmail(ctx, tx, email) + if err != nil { + if !errors.Is(err, coredata.ErrResourceNotFound) { + return fmt.Errorf("cannot load identity by email: %w", err) + } + + identity = &coredata.Identity{ + ID: gid.New(gid.NilTenant, coredata.IdentityEntityType), + EmailAddress: email, + FullName: claims.Name, + EmailAddressVerified: true, + CreatedAt: now, + UpdatedAt: now, + } + + if err := identity.Insert(ctx, tx); err != nil { + return fmt.Errorf("cannot insert identity: %w", err) + } + + return nil + } + + if !identity.EmailAddressVerified { + identity.EmailAddressVerified = true + identity.UpdatedAt = now + + if err := identity.Update(ctx, tx); err != nil { + return fmt.Errorf("cannot update identity: %w", err) + } + } + + return nil + }, + ) + if err != nil { + return nil, "", err + } + + return identity, oidcState.ContinueURL, nil +} + +func (s *Service) verifyAndParseIDToken(ctx context.Context, info *providerInfo, rawIDToken string, expectedNonce string) (*idTokenClaims, error) { + parts := strings.Split(rawIDToken, ".") + if len(parts) != 3 { + return nil, fmt.Errorf("cannot parse id token: invalid format") + } + + headerJSON, err := base64.RawURLEncoding.DecodeString(parts[0]) + if err != nil { + return nil, fmt.Errorf("cannot decode id token header: %w", err) + } + + var header struct { + Alg string `json:"alg"` + Kid string `json:"kid"` + } + if err := json.Unmarshal(headerJSON, &header); err != nil { + return nil, fmt.Errorf("cannot parse id token header: %w", err) + } + + key, err := s.getSigningKey(ctx, info.jwksURL, header.Kid) + if err != nil { + return nil, fmt.Errorf("cannot get signing key: %w", err) + } + + signedContent := parts[0] + "." + parts[1] + signature, err := base64.RawURLEncoding.DecodeString(parts[2]) + if err != nil { + return nil, fmt.Errorf("cannot decode signature: %w", err) + } + + if err := verifySignature(header.Alg, key, []byte(signedContent), signature); err != nil { + return nil, fmt.Errorf("cannot verify id token signature: %w", err) + } + + payloadJSON, err := base64.RawURLEncoding.DecodeString(parts[1]) + if err != nil { + return nil, fmt.Errorf("cannot decode id token payload: %w", err) + } + + var claims idTokenClaims + if err := json.Unmarshal(payloadJSON, &claims); err != nil { + return nil, fmt.Errorf("cannot parse id token claims: %w", err) + } + + if !info.issuerValidator(claims.Issuer) { + return nil, fmt.Errorf("cannot validate issuer: unexpected issuer %q", claims.Issuer) + } + + if !claims.hasAudience(info.oauth2Config.ClientID) { + return nil, fmt.Errorf("cannot validate audience: expected %q", info.oauth2Config.ClientID) + } + + if claims.Nonce != expectedNonce { + return nil, fmt.Errorf("cannot validate nonce: expected %q, got %q", expectedNonce, claims.Nonce) + } + + if time.Now().After(time.Unix(int64(claims.ExpiresAt), 0)) { + return nil, fmt.Errorf("cannot validate id token: token has expired") + } + + return &claims, nil +} + +func (s *Service) getSigningKey(ctx context.Context, jwksURL string, kid string) (crypto.PublicKey, error) { + s.jwksMu.RLock() + entry, ok := s.jwksCache[jwksURL] + s.jwksMu.RUnlock() + + if !ok || time.Since(entry.fetchedAt) > jwksCacheTTL { + keys, err := fetchJWKS(ctx, jwksURL) + if err != nil { + return nil, err + } + + entry = &jwksEntry{keys: keys, fetchedAt: time.Now()} + s.jwksMu.Lock() + s.jwksCache[jwksURL] = entry + s.jwksMu.Unlock() + } + + for _, k := range entry.keys { + if k.Kid == kid { + return parseJWK(k) + } + } + + // Key not found in cache, try refreshing + keys, err := fetchJWKS(ctx, jwksURL) + if err != nil { + return nil, err + } + + entry = &jwksEntry{keys: keys, fetchedAt: time.Now()} + s.jwksMu.Lock() + s.jwksCache[jwksURL] = entry + s.jwksMu.Unlock() + + for _, k := range entry.keys { + if k.Kid == kid { + return parseJWK(k) + } + } + + return nil, fmt.Errorf("cannot find signing key %q in JWKS", kid) +} + +func fetchJWKS(ctx context.Context, jwksURL string) ([]jwk, error) { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, jwksURL, nil) + if err != nil { + return nil, fmt.Errorf("cannot create jwks request: %w", err) + } + + resp, err := http.DefaultClient.Do(req) + if err != nil { + return nil, fmt.Errorf("cannot fetch jwks: %w", err) + } + defer func() { _ = resp.Body.Close() }() + + body, err := io.ReadAll(io.LimitReader(resp.Body, 1<<20)) + if err != nil { + return nil, fmt.Errorf("cannot read jwks response: %w", err) + } + + if resp.StatusCode != http.StatusOK { + return nil, fmt.Errorf("cannot fetch jwks: unexpected status %d", resp.StatusCode) + } + + var jwksResp jwksResponse + if err := json.Unmarshal(body, &jwksResp); err != nil { + return nil, fmt.Errorf("cannot parse jwks response: %w", err) + } + + return jwksResp.Keys, nil +} + +func parseJWK(k jwk) (crypto.PublicKey, error) { + switch k.Kty { + case "RSA": + nBytes, err := base64.RawURLEncoding.DecodeString(k.N) + if err != nil { + return nil, fmt.Errorf("cannot decode RSA modulus: %w", err) + } + eBytes, err := base64.RawURLEncoding.DecodeString(k.E) + if err != nil { + return nil, fmt.Errorf("cannot decode RSA exponent: %w", err) + } + + n := new(big.Int).SetBytes(nBytes) + e := 0 + for _, b := range eBytes { + e = e<<8 + int(b) + } + + return &rsa.PublicKey{N: n, E: e}, nil + + case "EC": + var curve elliptic.Curve + switch k.Crv { + case "P-256": + curve = elliptic.P256() + case "P-384": + curve = elliptic.P384() + case "P-521": + curve = elliptic.P521() + default: + return nil, fmt.Errorf("unsupported EC curve: %s", k.Crv) + } + + xBytes, err := base64.RawURLEncoding.DecodeString(k.X) + if err != nil { + return nil, fmt.Errorf("cannot decode EC X: %w", err) + } + yBytes, err := base64.RawURLEncoding.DecodeString(k.Y) + if err != nil { + return nil, fmt.Errorf("cannot decode EC Y: %w", err) + } + + return &ecdsa.PublicKey{ + Curve: curve, + X: new(big.Int).SetBytes(xBytes), + Y: new(big.Int).SetBytes(yBytes), + }, nil + + default: + return nil, fmt.Errorf("unsupported key type: %s", k.Kty) + } +} + +func verifySignature(alg string, key crypto.PublicKey, signedContent []byte, signature []byte) error { + hash := sha256.Sum256(signedContent) + + switch alg { + case "RS256": + rsaKey, ok := key.(*rsa.PublicKey) + if !ok { + return fmt.Errorf("cannot verify RS256 signature: expected RSA public key") + } + return rsa.VerifyPKCS1v15(rsaKey, crypto.SHA256, hash[:], signature) + + case "ES256": + ecKey, ok := key.(*ecdsa.PublicKey) + if !ok { + return fmt.Errorf("cannot verify ES256 signature: expected ECDSA public key") + } + + // JWS (RFC 7515) encodes ECDSA signatures as raw R||S + // concatenation (2x32 bytes for P-256), not ASN.1 DER. + // Convert to ASN.1 for ecdsa.VerifyASN1. + keySize := (ecKey.Curve.Params().BitSize + 7) / 8 + if len(signature) != 2*keySize { + return fmt.Errorf("cannot verify ES256 signature: invalid length %d, expected %d", len(signature), 2*keySize) + } + + r := new(big.Int).SetBytes(signature[:keySize]) + sigS := new(big.Int).SetBytes(signature[keySize:]) + + derSig, err := asn1.Marshal(struct { + R, S *big.Int + }{r, sigS}) + if err != nil { + return fmt.Errorf("cannot encode ECDSA signature to ASN.1: %w", err) + } + + if !ecdsa.VerifyASN1(ecKey, hash[:], derSig) { + return fmt.Errorf("cannot verify ECDSA signature") + } + return nil + + default: + return fmt.Errorf("cannot verify signature: unsupported algorithm %s", alg) + } +} + +func computeCodeChallenge(codeVerifier string) string { + hash := sha256.Sum256([]byte(codeVerifier)) + return base64.RawURLEncoding.EncodeToString(hash[:]) +} + +func generateRandomString(length int) (string, error) { + b := make([]byte, length) + if _, err := io.ReadFull(cryptoRandReader, b); err != nil { + return "", fmt.Errorf("cannot generate random bytes: %w", err) + } + return base64.RawURLEncoding.EncodeToString(b), nil +} diff --git a/pkg/iam/service.go b/pkg/iam/service.go index 3a91d9adc..5f67e9014 100644 --- a/pkg/iam/service.go +++ b/pkg/iam/service.go @@ -18,6 +18,7 @@ import ( "go.probo.inc/probo/pkg/crypto/passwdhash" "go.probo.inc/probo/pkg/filemanager" "go.probo.inc/probo/pkg/gid" + "go.probo.inc/probo/pkg/iam/oidc" "go.probo.inc/probo/pkg/iam/saml" "go.probo.inc/probo/pkg/iam/scim" "golang.org/x/sync/errgroup" @@ -46,6 +47,7 @@ type ( SessionService *SessionService AuthService *AuthService SAMLService *saml.Service + OIDCService *oidc.Service SCIMService *scim.Service APIKeyService *APIKeyService Authorizer *Authorizer @@ -73,6 +75,8 @@ type ( DomainVerificationResolverAddr string SCIMBridgeSyncInterval time.Duration SCIMBridgePollInterval time.Duration + GoogleOIDC oidc.ProviderConfig + MicrosoftOIDC oidc.ProviderConfig } ) @@ -135,6 +139,14 @@ func NewService( } svc.SAMLService = samlService + svc.OIDCService = oidc.NewService( + svc.pg, + svc.baseURL, + cfg.GoogleOIDC, + cfg.MicrosoftOIDC, + cfg.Logger, + ) + svc.SCIMService = scim.NewService( svc.pg, cfg.Logger.Named("scim"), @@ -166,6 +178,7 @@ func (s *Service) Run(ctx context.Context) error { g, ctx := errgroup.WithContext(ctx) g.Go(func() error { return s.SAMLService.Run(ctx) }) + g.Go(func() error { return s.OIDCService.Run(ctx) }) g.Go(func() error { return s.samlDomainVerifier.Run(ctx) }) g.Go(func() error { return s.SCIMService.Run(ctx) }) diff --git a/pkg/probod/auth_config.go b/pkg/probod/auth_config.go index 32ae0c9ea..0ed402a46 100644 --- a/pkg/probod/auth_config.go +++ b/pkg/probod/auth_config.go @@ -20,13 +20,15 @@ import ( ) type AuthConfig struct { - Cookie CookieConfig `json:"cookie"` - Password PasswordConfig `json:"password"` - DisableSignup bool `json:"disable-signup"` - InvitationConfirmationTokenValidity int `json:"invitation-confirmation-token-validity"` - PasswordResetTokenValidity int `json:"password-reset-token-validity"` - MagicLinkTokenValidity int `json:"magic-link-token-validity"` - SAML SAMLConfig `json:"saml"` + Cookie CookieConfig `json:"cookie"` + Password PasswordConfig `json:"password"` + DisableSignup bool `json:"disable-signup"` + InvitationConfirmationTokenValidity int `json:"invitation-confirmation-token-validity"` + PasswordResetTokenValidity int `json:"password-reset-token-validity"` + MagicLinkTokenValidity int `json:"magic-link-token-validity"` + SAML SAMLConfig `json:"saml"` + Google OIDCProviderConfig `json:"google"` + Microsoft OIDCProviderConfig `json:"microsoft"` } type CookieConfig struct { diff --git a/pkg/probod/oidc_config.go b/pkg/probod/oidc_config.go new file mode 100644 index 000000000..4638a963f --- /dev/null +++ b/pkg/probod/oidc_config.go @@ -0,0 +1,21 @@ +// Copyright (c) 2025 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package probod + +type OIDCProviderConfig struct { + ClientID string `json:"client-id"` + ClientSecret string `json:"client-secret"` + Enabled bool `json:"enabled"` +} diff --git a/pkg/probod/probod.go b/pkg/probod/probod.go index 9bc066a36..aa79122b9 100644 --- a/pkg/probod/probod.go +++ b/pkg/probod/probod.go @@ -55,6 +55,7 @@ import ( "go.probo.inc/probo/pkg/filemanager" "go.probo.inc/probo/pkg/html2pdf" "go.probo.inc/probo/pkg/iam" + "go.probo.inc/probo/pkg/iam/oidc" "go.probo.inc/probo/pkg/mailer" "go.probo.inc/probo/pkg/mailman" "go.probo.inc/probo/pkg/probo" @@ -381,6 +382,16 @@ func (impl *Implm) Run( DomainVerificationResolverAddr: impl.cfg.Auth.SAML.DomainVerificationResolverAddr, SCIMBridgeSyncInterval: time.Duration(impl.cfg.SCIMBridge.SyncInterval) * time.Second, SCIMBridgePollInterval: time.Duration(impl.cfg.SCIMBridge.PollInterval) * time.Second, + GoogleOIDC: oidc.ProviderConfig{ + ClientID: impl.cfg.Auth.Google.ClientID, + ClientSecret: impl.cfg.Auth.Google.ClientSecret, + Enabled: impl.cfg.Auth.Google.Enabled, + }, + MicrosoftOIDC: oidc.ProviderConfig{ + ClientID: impl.cfg.Auth.Microsoft.ClientID, + ClientSecret: impl.cfg.Auth.Microsoft.ClientSecret, + Enabled: impl.cfg.Auth.Microsoft.Enabled, + }, }, ) if err != nil { diff --git a/pkg/server/api/connect/v1/oidc_handler.go b/pkg/server/api/connect/v1/oidc_handler.go new file mode 100644 index 000000000..cc9c0c98a --- /dev/null +++ b/pkg/server/api/connect/v1/oidc_handler.go @@ -0,0 +1,161 @@ +// Copyright (c) 2025 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package connect_v1 + +import ( + "errors" + "net/http" + "strings" + + "github.com/go-chi/chi/v5" + "go.gearno.de/kit/httpserver" + "go.gearno.de/kit/log" + "go.probo.inc/probo/pkg/baseurl" + "go.probo.inc/probo/pkg/coredata" + "go.probo.inc/probo/pkg/iam" + "go.probo.inc/probo/pkg/saferedirect" + "go.probo.inc/probo/pkg/securecookie" + "go.probo.inc/probo/pkg/server/api/authn" +) + +type OIDCHandler struct { + iam *iam.Service + sessionCookie *authn.Cookie + baseURL *baseurl.BaseURL + logger *log.Logger + safeRedirect *saferedirect.SafeRedirect +} + +func NewOIDCHandler(iam *iam.Service, cookieConfig securecookie.Config, baseURL *baseurl.BaseURL, logger *log.Logger) *OIDCHandler { + return &OIDCHandler{ + iam: iam, + sessionCookie: authn.NewCookie(&cookieConfig), + baseURL: baseURL, + logger: logger, + safeRedirect: &saferedirect.SafeRedirect{AllowedHost: baseURL.Host()}, + } +} + +func (h *OIDCHandler) LoginHandler(w http.ResponseWriter, r *http.Request) { + ctx := r.Context() + + provider, err := parseOIDCProvider(chi.URLParam(r, "provider")) + if err != nil { + httpserver.RenderError(w, http.StatusBadRequest, errors.New("invalid provider")) + return + } + + if !h.iam.OIDCService.IsProviderEnabled(provider) { + httpserver.RenderError(w, http.StatusNotFound, errors.New("provider not enabled")) + return + } + + continueURL := r.URL.Query().Get("continue") + + authURL, err := h.iam.OIDCService.InitiateLogin(ctx, provider, continueURL) + if err != nil { + h.logger.ErrorCtx(ctx, "cannot initiate OIDC login", log.Error(err)) + httpserver.RenderError(w, http.StatusInternalServerError, errors.New("internal server error")) + return + } + + http.Redirect(w, r, authURL, http.StatusFound) +} + +func (h *OIDCHandler) CallbackHandler(w http.ResponseWriter, r *http.Request) { + ctx := r.Context() + + provider, err := parseOIDCProvider(chi.URLParam(r, "provider")) + if err != nil { + httpserver.RenderError(w, http.StatusBadRequest, errors.New("invalid provider")) + return + } + + errParam := r.URL.Query().Get("error") + if errParam != "" { + h.logger.WarnCtx( + ctx, + "OIDC provider returned error", + log.String("error", errParam), + log.String("error_description", r.URL.Query().Get("error_description")), + ) + httpserver.RenderError(w, http.StatusUnauthorized, errors.New("authentication failed")) + return + } + + stateParam := r.URL.Query().Get("state") + code := r.URL.Query().Get("code") + + if stateParam == "" || code == "" { + httpserver.RenderError(w, http.StatusBadRequest, errors.New("missing state or code")) + return + } + + identity, continueURL, err := h.iam.OIDCService.HandleCallback(ctx, provider, stateParam, code) + if err != nil { + h.logger.ErrorCtx(ctx, "cannot handle OIDC callback", log.Error(err)) + httpserver.RenderError(w, http.StatusUnauthorized, errors.New("authentication failed")) + return + } + + var authMethod coredata.AuthMethod + switch provider { + case coredata.OIDCProviderGoogle: + authMethod = coredata.AuthMethodGoogle + case coredata.OIDCProviderMicrosoft: + authMethod = coredata.AuthMethodMicrosoft + } + + rootSession := authn.SessionFromContext(ctx) + + switch { + case rootSession == nil: + rootSession, err = h.iam.AuthService.OpenSessionWithOIDC(ctx, identity.ID, authMethod) + if err != nil { + h.logger.ErrorCtx(ctx, "cannot open root session", log.Error(err)) + httpserver.RenderError(w, http.StatusInternalServerError, errors.New("internal server error")) + return + } + case rootSession.IdentityID != identity.ID: + err = h.iam.SessionService.CloseSession(ctx, rootSession.ID) + if err != nil { + h.logger.ErrorCtx(ctx, "cannot close session", log.Error(err)) + httpserver.RenderError(w, http.StatusInternalServerError, errors.New("internal server error")) + return + } + + rootSession, err = h.iam.AuthService.OpenSessionWithOIDC(ctx, identity.ID, authMethod) + if err != nil { + h.logger.ErrorCtx(ctx, "cannot open root session", log.Error(err)) + httpserver.RenderError(w, http.StatusInternalServerError, errors.New("internal server error")) + return + } + } + + h.sessionCookie.Set(w, rootSession) + + h.safeRedirect.Redirect(w, r, continueURL, "/", http.StatusFound) +} + +func parseOIDCProvider(s string) (coredata.OIDCProvider, error) { + switch strings.ToLower(s) { + case "google": + return coredata.OIDCProviderGoogle, nil + case "microsoft": + return coredata.OIDCProviderMicrosoft, nil + default: + return "", errors.New("unknown provider") + } +} diff --git a/pkg/server/api/connect/v1/resolver.go b/pkg/server/api/connect/v1/resolver.go index ddc367a26..24c050a6a 100644 --- a/pkg/server/api/connect/v1/resolver.go +++ b/pkg/server/api/connect/v1/resolver.go @@ -52,10 +52,14 @@ func NewMux(logger *log.Logger, svc *iam.Service, cookieConfig securecookie.Conf router := r.With(sessionMiddleware, apiKeyMiddleware) + oidcHandler := NewOIDCHandler(svc, cookieConfig, baseURL, logger) + router.Handle("/graphql", graphqlHandler) router.Get("/saml/2.0/metadata", samlHandler.MetadataHandler) router.Post("/saml/2.0/consume", samlHandler.ConsumeHandler) router.Get("/saml/2.0/{samlConfigID}", samlHandler.LoginHandler) + router.Get("/oidc/{provider}/login", oidcHandler.LoginHandler) + router.Get("/oidc/{provider}/callback", oidcHandler.CallbackHandler) // SCIM 2.0 endpoints - these use their own bearer token authentication scimServer := NewSCIMServer(scimHandler) diff --git a/pkg/server/api/connect/v1/schema.graphql b/pkg/server/api/connect/v1/schema.graphql index fb4f68371..f8f672b83 100644 --- a/pkg/server/api/connect/v1/schema.graphql +++ b/pkg/server/api/connect/v1/schema.graphql @@ -41,12 +41,20 @@ interface Node { id: ID! } +type OIDCProviderInfo { + name: String! + loginURL: String! +} + type Query { node(id: ID!): Node @session(required: PRESENT) viewer: Identity @session(required: PRESENT) ssoLoginURL(email: EmailAddr!): String @goField(forceResolver: true) @session(required: OPTIONAL) + oidcProviders: [OIDCProviderInfo!]! + @goField(forceResolver: true) + @session(required: OPTIONAL) } type Mutation { diff --git a/pkg/server/api/connect/v1/v1_resolver.go b/pkg/server/api/connect/v1/v1_resolver.go index 1b6564f97..64fef54d8 100644 --- a/pkg/server/api/connect/v1/v1_resolver.go +++ b/pkg/server/api/connect/v1/v1_resolver.go @@ -9,6 +9,7 @@ import ( "context" "errors" "fmt" + "strings" "time" "github.com/99designs/gqlgen/graphql" @@ -1746,6 +1747,21 @@ func (r *queryResolver) SsoLoginURL(ctx context.Context, email mail.Addr) (*stri return &loginURL, nil } +// OidcProviders is the resolver for the oidcProviders field. +func (r *queryResolver) OidcProviders(ctx context.Context) ([]*types.OIDCProviderInfo, error) { + providers := r.iam.OIDCService.EnabledProviders() + result := make([]*types.OIDCProviderInfo, 0, len(providers)) + + for _, p := range providers { + result = append(result, &types.OIDCProviderInfo{ + Name: strings.ToLower(p.String()), + LoginURL: r.baseURL.WithPath("/api/connect/v1/oidc/" + strings.ToLower(p.String()) + "/login").MustString(), + }) + } + + return result, nil +} + // TestLoginURL is the resolver for the testLoginUrl field. func (r *sAMLConfigurationResolver) TestLoginURL(ctx context.Context, obj *types.SAMLConfiguration) (string, error) { return r.baseURL.WithPath("/api/connect/v1/saml/2.0/" + obj.ID.String()).MustString(), nil