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
43 changes: 42 additions & 1 deletion pkg/auth/oidc.go
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@ package auth

import (
"context"
"crypto"
"crypto/tls"
"crypto/x509"
"fmt"
Expand Down Expand Up @@ -76,6 +77,17 @@ func createOIDCHTTPClient(trustedCAFile string, insecureSkipVerify bool, proxyUR
return &http.Client{Transport: transport}, nil
}

func VerifierFromPublicKeys(cfg v1.AuthOIDCServerConfig, keys []crypto.PublicKey) (*oidc.IDTokenVerifier, error) {
verifierConf := &oidc.Config{
ClientID: cfg.Audience,
SkipClientIDCheck: cfg.Audience == "",
SkipExpiryCheck: cfg.SkipExpiryCheck,
SkipIssuerCheck: cfg.SkipIssuerCheck,
}
keySet := &oidc.StaticKeySet{PublicKeys: keys}
return oidc.NewVerifier(cfg.Issuer, keySet, verifierConf), nil
}

// nonCachingTokenSource wraps a clientcredentials.Config to fetch a fresh
// token on every call. This is used as a fallback when the OIDC provider
// does not return expires_in, which would cause a caching TokenSource to
Expand Down Expand Up @@ -273,10 +285,39 @@ type OidcAuthConsumer struct {
subjectsFromLogin map[string]struct{}
}

func NewTokenVerifierFromStatic(cfg v1.AuthOIDCServerConfig) (TokenVerifier, error) {
if cfg.IssuerSpec.PemFile != "" {
pemBytes, err := os.ReadFile(cfg.IssuerSpec.PemFile)
if err == nil {
key, err := DecodePemCert(pemBytes)
if err != nil {
return nil, err
}
return VerifierFromPublicKeys(cfg, key)
}
}
if cfg.IssuerSpec.JWKSFile != "" {
jwksBytes, err := os.ReadFile(cfg.IssuerSpec.JWKSFile)
if err != nil {
return nil, err
}
jwks, err := DecodeJWKSFile(jwksBytes)
if err != nil {
return nil, err
}
return VerifierFromPublicKeys(cfg, DecodeJWKS(jwks))
}
return VerifierFromPublicKeys(cfg, DecodeJWKS(cfg.IssuerSpec.JWKS))
}

func NewTokenVerifier(cfg v1.AuthOIDCServerConfig) TokenVerifier {
provider, err := oidc.NewProvider(context.Background(), cfg.Issuer)
if err != nil {
panic(err)
verifier, err := NewTokenVerifierFromStatic(cfg)
if err != nil {
panic(err)
}
return verifier
}
verifierConf := oidc.Config{
ClientID: cfg.Audience,
Expand Down
260 changes: 260 additions & 0 deletions pkg/auth/oidc_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,8 @@ package auth_test

import (
"context"
"crypto/rsa"
_ "embed"
"encoding/json"
"net/http"
"net/http/httptest"
Expand All @@ -10,6 +12,8 @@ import (
"time"

"github.com/coreos/go-oidc/v3/oidc"
"github.com/go-jose/go-jose/v4"
"github.com/go-jose/go-jose/v4/jwt"
"github.com/stretchr/testify/require"

"github.com/fatedier/frp/pkg/auth"
Expand All @@ -19,6 +23,15 @@ import (

type mockTokenVerifier struct{}

//go:embed testSample/pki.json
var pkiJwkContent []byte

//go:embed testSample/pki/server.full.pem
var pkiPemContent []byte

//go:embed testSample/pem_single.pem
var singlePemContent []byte

func (m *mockTokenVerifier) Verify(ctx context.Context, subject string) (*oidc.IDToken, error) {
return &oidc.IDToken{
Subject: subject,
Expand Down Expand Up @@ -251,3 +264,250 @@ func TestNewOidcAuthSetterRejectsInvalidStaticConfig(t *testing.T) {
r.Error(err)
r.Contains(err.Error(), "cannot specify both auth.oidc.audience and auth.oidc.additionalEndpointParams.audience")
}

func setupStaticOidc(t *testing.T) (*jose.JSONWebKeySet, jwt.Builder) {
// Test Setup include Load JWKS + Generate Token
r := require.New(t)
jwks, err := auth.DecodeJWKSFile(pkiJwkContent)
r.NoError((err))
signer, err := jose.NewSigner(jose.SigningKey{
Algorithm: jose.RS256,
Key: jwks.Key("00000000-0000-ffff-0000-000000000000")[0].Key.(*rsa.PrivateKey),
}, nil)
r.NoError((err))
for i, k := range jwks.Keys {
pk := k.Key.(*rsa.PrivateKey)
jwks.Keys[i].Key = pk.Public() // We provides Private so to avoid issue with Verifier cast them as Public
}
return jwks, jwt.Signed(signer)
}

func TestPingAfterStaticLoginSucceeds(t *testing.T) {
// Test Setup include Load JWKS + Generate Token
jwks, builder := setupStaticOidc(t)
builder = builder.Claims(jwt.Claims{
Issuer: "https://kubernetes.default.svc.cluster.local",
Subject: "system:serviceaccount:default:default",
Audience: jwt.Audience{"https://kubernetes.default.svc.cluster.local", "k3s"},
NotBefore: jwt.NewNumericDate(time.Now()),
IssuedAt: jwt.NewNumericDate(time.Now()),
Expiry: jwt.NewNumericDate(time.Now().Add(1 * time.Hour)),
})

r := require.New(t)
verifier, err := auth.NewTokenVerifierFromStatic(
v1.AuthOIDCServerConfig{
Issuer: "https://kubernetes.default.svc.cluster.local",
Audience: "k3s",
IssuerSpec: v1.AuthOIDCIssuer{
JWKS: jwks,
},
})
r.NoError((err))
consumer := auth.NewOidcAuthVerifier([]v1.AuthScope{v1.AuthScopeHeartBeats}, verifier)
token, err := builder.Serialize()
r.NoError(err)

err = consumer.VerifyLogin(&msg.Login{
PrivilegeKey: token,
})
r.NoError(err)

err = consumer.VerifyPing(&msg.Ping{
PrivilegeKey: token,
Timestamp: time.Now().UnixMilli(),
})
r.NoError(err)
}

func TestExpiredTokenStaticLoginFailed(t *testing.T) {
// Test Setup include Load JWKS + Generate Token
jwks, builder := setupStaticOidc(t)
builder = builder.Claims(jwt.Claims{
Issuer: "https://kubernetes.default.svc.cluster.local",
Subject: "system:serviceaccount:default:default",
Audience: jwt.Audience{"https://kubernetes.default.svc.cluster.local", "k3s"},
NotBefore: jwt.NewNumericDate(time.Now()),
IssuedAt: jwt.NewNumericDate(time.Now()),
Expiry: jwt.NewNumericDate(time.Now().Add(-1 * time.Hour)),
})

r := require.New(t)
verifier, err := auth.NewTokenVerifierFromStatic(v1.AuthOIDCServerConfig{
Issuer: "https://kubernetes.default.svc.cluster.local",
Audience: "k3s",
IssuerSpec: v1.AuthOIDCIssuer{
JWKS: jwks,
},
})
r.NoError((err))
consumer := auth.NewOidcAuthVerifier([]v1.AuthScope{v1.AuthScopeHeartBeats}, verifier)
token, err := builder.Serialize()
r.NoError(err)

err = consumer.VerifyLogin(&msg.Login{
PrivilegeKey: token,
})
r.Error(err)
r.Contains(err.Error(), "oidc: token is expired")
}

func TestBadAudienceStaticLoginFailed(t *testing.T) {
// Test Setup include Load JWKS + Generate Token
jwks, builder := setupStaticOidc(t)
builder = builder.Claims(jwt.Claims{
Issuer: "https://kubernetes.default.svc.cluster.local",
Subject: "system:serviceaccount:default:default",
Audience: jwt.Audience{"https://kubernetes.default.svc.cluster.local", "k3s"},
NotBefore: jwt.NewNumericDate(time.Now()),
IssuedAt: jwt.NewNumericDate(time.Now()),
Expiry: jwt.NewNumericDate(time.Now().Add(-1 * time.Hour)),
})

r := require.New(t)
verifier, err := auth.NewTokenVerifierFromStatic(v1.AuthOIDCServerConfig{
Issuer: "https://kubernetes.default.svc.cluster.local",
Audience: "k8s",
IssuerSpec: v1.AuthOIDCIssuer{
JWKS: jwks,
},
})
r.NoError((err))
consumer := auth.NewOidcAuthVerifier([]v1.AuthScope{v1.AuthScopeHeartBeats}, verifier)
token, err := builder.Serialize()
r.NoError(err)

err = consumer.VerifyLogin(&msg.Login{
PrivilegeKey: token,
})
r.Error(err)
r.Contains(err.Error(), `oidc: expected audience "k8s" got ["https://kubernetes.default.svc.cluster.local" "k3s"]`)
}

func TestBadIssuerStaticLoginFailed(t *testing.T) {
// Test Setup include Load JWKS + Generate Token
jwks, builder := setupStaticOidc(t)
builder = builder.Claims(jwt.Claims{
Issuer: "https://kubernetes.default.svc.cluster.local",
Subject: "system:serviceaccount:default:default",
Audience: jwt.Audience{"https://kubernetes.default.svc.cluster.local", "k3s"},
NotBefore: jwt.NewNumericDate(time.Now()),
IssuedAt: jwt.NewNumericDate(time.Now()),
Expiry: jwt.NewNumericDate(time.Now().Add(-1 * time.Hour)),
})

r := require.New(t)
verifier, err := auth.NewTokenVerifierFromStatic(v1.AuthOIDCServerConfig{
Issuer: "https://kubernetes.default.svc.cluster",
Audience: "k3s",
IssuerSpec: v1.AuthOIDCIssuer{
JWKS: jwks,
},
})
r.NoError((err))
consumer := auth.NewOidcAuthVerifier([]v1.AuthScope{v1.AuthScopeHeartBeats}, verifier)
token, err := builder.Serialize()
r.NoError(err)

err = consumer.VerifyLogin(&msg.Login{
PrivilegeKey: token,
})
r.Error(err)
r.Contains(err.Error(), "oidc: id token issued by a different provider")
}

func TestPingAfterStaticLoginCrossJKWSPemSucceeds(t *testing.T) {
// Test Setup include Load JWKS + Generate Token
_, builder := setupStaticOidc(t)
builder = builder.Claims(jwt.Claims{
Issuer: "https://kubernetes.default.svc.cluster.local",
Subject: "system:serviceaccount:default:default",
Audience: jwt.Audience{"https://kubernetes.default.svc.cluster.local", "k3s"},
NotBefore: jwt.NewNumericDate(time.Now()),
IssuedAt: jwt.NewNumericDate(time.Now()),
Expiry: jwt.NewNumericDate(time.Now().Add(1 * time.Hour)),
})

r := require.New(t)
keys, err := auth.DecodePemCert(pkiPemContent)
r.NoError(err)
verifier, err := auth.VerifierFromPublicKeys(
v1.AuthOIDCServerConfig{Issuer: "https://kubernetes.default.svc.cluster.local", Audience: "k3s"},
keys,
)
r.NoError((err))
consumer := auth.NewOidcAuthVerifier([]v1.AuthScope{v1.AuthScopeHeartBeats}, verifier)
token, err := builder.Serialize()
r.NoError(err)

err = consumer.VerifyLogin(&msg.Login{
PrivilegeKey: token,
})
r.NoError((err))

err = consumer.VerifyPing(&msg.Ping{
PrivilegeKey: token,
Timestamp: time.Now().UnixMilli(),
})
r.NoError(err)
}

func TestBadPublicKeyStaticLoginFailed(t *testing.T) {
// Test Setup include Load JWKS + Generate Token
_, builder := setupStaticOidc(t)
builder = builder.Claims(jwt.Claims{
Issuer: "https://kubernetes.default.svc.cluster.local",
Subject: "system:serviceaccount:default:default",
Audience: jwt.Audience{"https://kubernetes.default.svc.cluster.local", "k3s"},
NotBefore: jwt.NewNumericDate(time.Now()),
IssuedAt: jwt.NewNumericDate(time.Now()),
Expiry: jwt.NewNumericDate(time.Now().Add(1 * time.Hour)),
})

r := require.New(t)
keys, err := auth.DecodePemCert(singlePemContent)
r.NoError(err)
verifier, err := auth.VerifierFromPublicKeys(
v1.AuthOIDCServerConfig{Issuer: "https://kubernetes.default.svc.cluster.local", Audience: "k3s"},
keys,
)
r.NoError((err))
consumer := auth.NewOidcAuthVerifier([]v1.AuthScope{v1.AuthScopeHeartBeats}, verifier)
token, err := builder.Serialize()
r.NoError(err)

err = consumer.VerifyLogin(&msg.Login{
PrivilegeKey: token,
})
r.Error((err))
r.Contains(err.Error(), "failed to verify signature: no public keys able to verify jwt")
}

func TestEmptyKeysSetLoginFailed(t *testing.T) {
// Test Setup include Load JWKS + Generate Token
_, builder := setupStaticOidc(t)
builder = builder.Claims(jwt.Claims{
Issuer: "https://kubernetes.default.svc.cluster.local",
Subject: "system:serviceaccount:default:default",
Audience: jwt.Audience{"https://kubernetes.default.svc.cluster.local", "k3s"},
NotBefore: jwt.NewNumericDate(time.Now()),
IssuedAt: jwt.NewNumericDate(time.Now()),
Expiry: jwt.NewNumericDate(time.Now().Add(1 * time.Hour)),
})

r := require.New(t)
verifier, err := auth.VerifierFromPublicKeys(
v1.AuthOIDCServerConfig{Issuer: "https://kubernetes.default.svc.cluster.local", Audience: "k3s"},
auth.DecodeJWKS(nil),
)
r.NoError((err))
consumer := auth.NewOidcAuthVerifier([]v1.AuthScope{v1.AuthScopeHeartBeats}, verifier)
token, err := builder.Serialize()
r.NoError(err)

err = consumer.VerifyLogin(&msg.Login{
PrivilegeKey: token,
})
r.Error(err)
r.Contains(err.Error(), "invalid OIDC token in login: failed to verify signature: no public keys able to verify jwt")
}
20 changes: 20 additions & 0 deletions pkg/auth/testSample/jwks_multiple.json
Original file line number Diff line number Diff line change
@@ -0,0 +1,20 @@
{
"keys": [
{
"kty": "RSA",
"use": "sig",
"alg": "RS256",
"kid": "00000000-0000-0000-0000-000000000000",
"n": "7e2xRhkEeedFuNG-VVhqNXv1QQ3RAJHyPS77IUQQi5QURfOelAoucNcERS0pJ6LvkMtLcWcnGRfZxAIzYfCnCunpyEcYc2xN93TOShobWMzi_TS45EV_4vk7j2elabPoBh-LhrPCvJdOZqfqkSqv-NIDz1i_4W87LtzhMpD2I5s",
"e": "AQAB"
},
{
"kty": "RSA",
"use": "sig",
"alg": "RS256",
"kid": "00000000-0000-0000-0000-000000000001",
"n": "m_Z5AIjf9H1lJkBc5d9JKQQRkTuGagfkQIoVpyNq0krIWe6SYX0Fd-dhbcUSxiPTmKQhI3QF5qIZQV9OTsyKVC3moqFpvZcrFlHcfeFdds22nBthN6SrdjHfuc7VJP4Dpevnfi7xLN8Rjv0YY6D9EvzvKMS6Esn1yM6uY3K4QmE",
"e": "AQAB"
}
]
}
12 changes: 12 additions & 0 deletions pkg/auth/testSample/jwks_single.json
Original file line number Diff line number Diff line change
@@ -0,0 +1,12 @@
{
"keys": [
{
"kty": "RSA",
"use": "sig",
"alg": "RS256",
"kid": "00000000-0000-0000-0000-000000000000",
"n": "7e2xRhkEeedFuNG-VVhqNXv1QQ3RAJHyPS77IUQQi5QURfOelAoucNcERS0pJ6LvkMtLcWcnGRfZxAIzYfCnCunpyEcYc2xN93TOShobWMzi_TS45EV_4vk7j2elabPoBh-LhrPCvJdOZqfqkSqv-NIDz1i_4W87LtzhMpD2I5s",
"e": "AQAB"
}
]
}
Loading