Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
56 changes: 43 additions & 13 deletions admin/server/auth/auth.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,8 @@ package auth

import (
"context"
"fmt"
"strings"

"github.com/coreos/go-oidc/v3/oidc"
"github.com/rilldata/rill/admin"
Expand All @@ -28,21 +30,33 @@ type AuthenticatorOptions struct {
// It provides endpoints for login/logout, creates users, issues cookie-based auth tokens, and provides middleware for authenticating requests.
// The implementation was derived from: https://auth0.com/docs/quickstart/webapp/golang/01-login.
type Authenticator struct {
logger *zap.Logger
admin *admin.Service
cookies *cookies.Store
opts *AuthenticatorOptions
oidc *oidc.Provider
oauth2 oauth2.Config
logger *zap.Logger
admin *admin.Service
cookies *cookies.Store
opts *AuthenticatorOptions
oidc *oidc.Provider
oauth2 oauth2.Config
endSessionEndpoint string
}

// NewAuthenticator creates an Authenticator.
func NewAuthenticator(logger *zap.Logger, adm *admin.Service, cookieStore *cookies.Store, opts *AuthenticatorOptions) (*Authenticator, error) {
oidcProvider, err := oidc.NewProvider(context.Background(), "https://"+opts.AuthDomain+"/")
issuer := issuerURL(opts.AuthDomain)
oidcProvider, err := oidc.NewProvider(context.Background(), issuer)
if err != nil {
return nil, err
}

var claims struct {
EndSessionEndpoint string `json:"end_session_endpoint"`
}
if err := oidcProvider.Claims(&claims); err != nil {
return nil, fmt.Errorf("failed to parse the auth provider's discovery document: %w", err)
}
if claims.EndSessionEndpoint == "" && !isBareDomain(opts.AuthDomain) {
logger.Warn("auth provider does not publish an end_session_endpoint, so logging out will only end the Rill session", zap.String("issuer", issuer))
}

oauth2Config := oauth2.Config{
ClientID: opts.AuthClientID,
ClientSecret: opts.AuthClientSecret,
Expand All @@ -52,13 +66,29 @@ func NewAuthenticator(logger *zap.Logger, adm *admin.Service, cookieStore *cooki
}

a := &Authenticator{
logger: logger,
admin: adm,
cookies: cookieStore,
opts: opts,
oidc: oidcProvider,
oauth2: oauth2Config,
logger: logger,
admin: adm,
cookies: cookieStore,
opts: opts,
oidc: oidcProvider,
oauth2: oauth2Config,
endSessionEndpoint: claims.EndSessionEndpoint,
}

return a, nil
}

// issuerURL returns the OIDC issuer for authDomain.
// AuthDomain with "://" is a full issuer URL (Keycloak, Dex, etc.) used verbatim;
// without it, assume Auth0-style domain and append trailing slash.
func issuerURL(authDomain string) string {
if isBareDomain(authDomain) {
return "https://" + authDomain + "/"
}
return authDomain
}

// isBareDomain reports whether authDomain is an Auth0-style domain (e.g. "rill.auth0.com") rather than a full issuer URL.
func isBareDomain(authDomain string) bool {
return !strings.Contains(authDomain, "://")
}
111 changes: 111 additions & 0 deletions admin/server/auth/auth_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,111 @@
package auth

import (
"crypto/rand"
"crypto/rsa"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"

"github.com/go-jose/go-jose/v3"
"github.com/go-jose/go-jose/v3/jwt"
"github.com/rilldata/rill/admin"
"github.com/stretchr/testify/require"
"go.uber.org/zap"
)

// testProvider is a minimal OIDC provider serving a discovery document and a JWKS, and signing ID tokens with its key.
type testProvider struct {
*httptest.Server
key *rsa.PrivateKey
}

// newTestProvider starts a testProvider. The discovery document can be extended (or fields overridden) with extra.
func newTestProvider(t *testing.T, extra map[string]any) *testProvider {
key, err := rsa.GenerateKey(rand.Reader, 2048)
require.NoError(t, err)

p := &testProvider{key: key}
mux := http.NewServeMux()
mux.HandleFunc("/.well-known/openid-configuration", func(w http.ResponseWriter, r *http.Request) {
doc := map[string]any{
"issuer": p.URL,
"authorization_endpoint": p.URL + "/authorize",
"token_endpoint": p.URL + "/token",
"jwks_uri": p.URL + "/jwks",
"id_token_signing_alg_values_supported": []string{"RS256"},
}
for k, v := range extra {
doc[k] = v
}
_ = json.NewEncoder(w).Encode(doc)
})
mux.HandleFunc("/jwks", func(w http.ResponseWriter, r *http.Request) {
_ = json.NewEncoder(w).Encode(jose.JSONWebKeySet{Keys: []jose.JSONWebKey{{Key: &key.PublicKey, KeyID: "test", Algorithm: "RS256", Use: "sig"}}})
})
p.Server = httptest.NewServer(mux)
t.Cleanup(p.Close)
return p
}

// signIDToken returns an ID token for the given claims, signed with key (the provider's own key if nil).
func (p *testProvider) signIDToken(t *testing.T, key *rsa.PrivateKey, claims map[string]any) string {
if key == nil {
key = p.key
}
signer, err := jose.NewSigner(jose.SigningKey{Algorithm: jose.RS256, Key: key}, (&jose.SignerOptions{}).WithHeader("kid", "test"))
require.NoError(t, err)
raw, err := jwt.Signed(signer).Claims(claims).CompactSerialize()
require.NoError(t, err)
return raw
}

func TestIssuerURL(t *testing.T) {
tests := []struct {
authDomain string
want string
}{
{"rill.auth0.com", "https://rill.auth0.com/"},
{"https://idp.example.com/realms/rill", "https://idp.example.com/realms/rill"},
{"https://idp.example.com/realms/rill/", "https://idp.example.com/realms/rill/"},
{"http://localhost:5556/dex", "http://localhost:5556/dex"},
}
for _, tt := range tests {
require.Equal(t, tt.want, issuerURL(tt.authDomain), tt.authDomain)
}
}

func TestNewAuthenticator(t *testing.T) {
urls, err := admin.NewURLs("http://localhost:8080", "http://localhost:3000")
require.NoError(t, err)
adm := &admin.Service{URLs: urls}

tests := []struct {
name string
endSession any // value of end_session_endpoint in the discovery document; nil leaves it out
wantEndSession string
wantErrorSubstring string
}{
{name: "with end_session_endpoint", endSession: "https://idp.example.com/logout", wantEndSession: "https://idp.example.com/logout"},
{name: "without end_session_endpoint"},
{name: "malformed end_session_endpoint", endSession: 42, wantErrorSubstring: "discovery document"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
var discovery map[string]any
if tt.endSession != nil {
discovery = map[string]any{"end_session_endpoint": tt.endSession}
}
p := newTestProvider(t, discovery)

a, err := NewAuthenticator(zap.NewNop(), adm, nil, &AuthenticatorOptions{AuthDomain: p.URL, AuthClientID: "rill-client"})
if tt.wantErrorSubstring != "" {
require.ErrorContains(t, err, tt.wantErrorSubstring)
return
}
require.NoError(t, err)
require.Equal(t, tt.wantEndSession, a.endSessionEndpoint)
})
}
}
Loading
Loading