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 <bryan@getprobo.com>
This commit is contained in:
@@ -110,6 +110,16 @@ func (b *Builder) Build() (*probod.FullConfig, error) {
|
||||
DomainVerificationIntervalSeconds: b.getEnvIntOrDefault("SAML_DOMAIN_VERIFICATION_INTERVAL_SECONDS", 60),
|
||||
DomainVerificationResolverAddr: b.getEnvOrDefault("SAML_DOMAIN_VERIFICATION_RESOLVER_ADDR", "8.8.8.8:53"),
|
||||
},
|
||||
Google: probod.OIDCProviderConfig{
|
||||
ClientID: b.getEnv("AUTH_GOOGLE_CLIENT_ID"),
|
||||
ClientSecret: b.getEnv("AUTH_GOOGLE_CLIENT_SECRET"),
|
||||
Enabled: b.getEnv("AUTH_GOOGLE_CLIENT_ID") != "",
|
||||
},
|
||||
Microsoft: probod.OIDCProviderConfig{
|
||||
ClientID: b.getEnv("AUTH_MICROSOFT_CLIENT_ID"),
|
||||
ClientSecret: b.getEnv("AUTH_MICROSOFT_CLIENT_SECRET"),
|
||||
Enabled: b.getEnv("AUTH_MICROSOFT_CLIENT_ID") != "",
|
||||
},
|
||||
},
|
||||
TrustCenter: probod.TrustCenterConfig{
|
||||
HTTPAddr: b.getEnvOrDefault("TRUST_CENTER_HTTP_ADDR", ":80"),
|
||||
|
||||
19
pkg/coredata/migrations/20260319T150000Z.sql
Normal file
19
pkg/coredata/migrations/20260319T150000Z.sql
Normal file
@@ -0,0 +1,19 @@
|
||||
ALTER TYPE session_auth_method ADD VALUE 'GOOGLE';
|
||||
ALTER TYPE session_auth_method ADD VALUE 'MICROSOFT';
|
||||
|
||||
CREATE TYPE iam_oidc_provider AS ENUM (
|
||||
'GOOGLE',
|
||||
'MICROSOFT'
|
||||
);
|
||||
|
||||
CREATE TABLE iam_oidc_states (
|
||||
id TEXT PRIMARY KEY,
|
||||
provider iam_oidc_provider NOT NULL,
|
||||
nonce TEXT NOT NULL,
|
||||
code_verifier TEXT NOT NULL,
|
||||
continue_url TEXT NOT NULL DEFAULT '',
|
||||
created_at TIMESTAMP NOT NULL,
|
||||
expires_at TIMESTAMP NOT NULL
|
||||
);
|
||||
|
||||
CREATE INDEX idx_iam_oidc_states_expires_at ON iam_oidc_states(expires_at);
|
||||
137
pkg/coredata/oidc_state.go
Normal file
137
pkg/coredata/oidc_state.go
Normal file
@@ -0,0 +1,137 @@
|
||||
// Copyright (c) 2025 Probo Inc <hello@getprobo.com>.
|
||||
//
|
||||
// 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
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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,
|
||||
|
||||
99
pkg/iam/oidc/errors.go
Normal file
99
pkg/iam/oidc/errors.go
Normal file
@@ -0,0 +1,99 @@
|
||||
// Copyright (c) 2025 Probo Inc <hello@getprobo.com>.
|
||||
//
|
||||
// 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"
|
||||
}
|
||||
89
pkg/iam/oidc/gc.go
Normal file
89
pkg/iam/oidc/gc.go
Normal file
@@ -0,0 +1,89 @@
|
||||
// Copyright (c) 2025 Probo Inc <hello@getprobo.com>.
|
||||
//
|
||||
// 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
|
||||
},
|
||||
)
|
||||
}
|
||||
686
pkg/iam/oidc/service.go
Normal file
686
pkg/iam/oidc/service.go
Normal file
@@ -0,0 +1,686 @@
|
||||
// Copyright (c) 2025 Probo Inc <hello@getprobo.com>.
|
||||
//
|
||||
// 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
|
||||
}
|
||||
@@ -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) })
|
||||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
21
pkg/probod/oidc_config.go
Normal file
21
pkg/probod/oidc_config.go
Normal file
@@ -0,0 +1,21 @@
|
||||
// Copyright (c) 2025 Probo Inc <hello@getprobo.com>.
|
||||
//
|
||||
// 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"`
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
161
pkg/server/api/connect/v1/oidc_handler.go
Normal file
161
pkg/server/api/connect/v1/oidc_handler.go
Normal file
@@ -0,0 +1,161 @@
|
||||
// Copyright (c) 2025 Probo Inc <hello@getprobo.com>.
|
||||
//
|
||||
// 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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user