mirror of
https://github.com/oauth2-proxy/oauth2-proxy.git
synced 2026-10-03 13:13:15 +02:00
Add DiscoveryProvider to perform OIDC discovery
This commit is contained in:
@@ -0,0 +1,81 @@
|
||||
package oidc
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/oauth2-proxy/oauth2-proxy/v7/pkg/logger"
|
||||
"github.com/oauth2-proxy/oauth2-proxy/v7/pkg/requests"
|
||||
)
|
||||
|
||||
// providerJSON resresents the information we need from an OIDC discovery
|
||||
type providerJSON struct {
|
||||
Issuer string `json:"issuer"`
|
||||
AuthURL string `json:"authorization_endpoint"`
|
||||
TokenURL string `json:"token_endpoint"`
|
||||
JWKsURL string `json:"jwks_uri"`
|
||||
UserInfoURL string `json:"userinfo_endpoint"`
|
||||
}
|
||||
|
||||
// Endpoints represents the endpoints discovered as part of the OIDC discovery process
|
||||
// that will be used by the authentication providers.
|
||||
type Endpoints struct {
|
||||
AuthURL string
|
||||
TokenURL string
|
||||
JWKsURL string
|
||||
UserInfoURL string
|
||||
}
|
||||
|
||||
// DiscoveryProvider holds information about an identity provider having
|
||||
// used OIDC discovery to retrieve the information.
|
||||
type DiscoveryProvider interface {
|
||||
Endpoints() Endpoints
|
||||
}
|
||||
|
||||
// NewProvider allows a user to perform an OIDC discovery and returns the DiscoveryProvider.
|
||||
// We implement this here as opposed to using oidc.Provider so that we can override the Issuer verification check.
|
||||
// As we have our own verifier and fetch the userinfo separately, the rest of the oidc.Provider implementation is not
|
||||
// useful to us.
|
||||
func NewProvider(ctx context.Context, issuerURL string, skipIssuerVerification bool) (DiscoveryProvider, error) {
|
||||
// go-oidc doesn't let us pass bypass the issuer check this in the oidc.NewProvider call
|
||||
// (which uses discovery to get the URLs), so we'll do a quick check ourselves and if
|
||||
// we get the URLs, we'll just use the non-discovery path.
|
||||
|
||||
logger.Printf("Performing OIDC Discovery...")
|
||||
|
||||
var p providerJSON
|
||||
requestURL := strings.TrimSuffix(issuerURL, "/") + "/.well-known/openid-configuration"
|
||||
if err := requests.New(requestURL).WithContext(ctx).Do().UnmarshalInto(&p); err != nil {
|
||||
return nil, fmt.Errorf("failed to discover OIDC configuration: %v", err)
|
||||
}
|
||||
|
||||
if !skipIssuerVerification && p.Issuer != issuerURL {
|
||||
return nil, fmt.Errorf("oidc: issuer did not match the issuer returned by provider, expected %q got %q", issuerURL, p.Issuer)
|
||||
}
|
||||
|
||||
return &discoveryProvider{
|
||||
authURL: p.AuthURL,
|
||||
tokenURL: p.TokenURL,
|
||||
jwksURL: p.JWKsURL,
|
||||
userInfoURL: p.UserInfoURL,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// discoveryProvider holds the discovered endpoints
|
||||
type discoveryProvider struct {
|
||||
authURL string
|
||||
tokenURL string
|
||||
jwksURL string
|
||||
userInfoURL string
|
||||
}
|
||||
|
||||
// Endpoints returns the discovered endpoints needed for an authentication provider.
|
||||
func (p *discoveryProvider) Endpoints() Endpoints {
|
||||
return Endpoints{
|
||||
AuthURL: p.authURL,
|
||||
TokenURL: p.tokenURL,
|
||||
JWKsURL: p.jwksURL,
|
||||
UserInfoURL: p.userInfoURL,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,114 @@
|
||||
package oidc
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net"
|
||||
"net/http"
|
||||
|
||||
"github.com/oauth2-proxy/mockoidc"
|
||||
. "github.com/onsi/ginkgo"
|
||||
. "github.com/onsi/ginkgo/extensions/table"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
var _ = Describe("Provider", func() {
|
||||
type newProviderTableInput struct {
|
||||
skipIssuerVerification bool
|
||||
expectedError string
|
||||
middlewares func(*mockoidc.MockOIDC) []func(http.Handler) http.Handler
|
||||
}
|
||||
|
||||
DescribeTable("NewProvider", func(in *newProviderTableInput) {
|
||||
m, err := mockoidc.NewServer(nil)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
|
||||
if in.middlewares != nil {
|
||||
middlewares := in.middlewares(m)
|
||||
for _, middlware := range middlewares {
|
||||
m.AddMiddleware(middlware)
|
||||
}
|
||||
}
|
||||
|
||||
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
|
||||
Expect(m.Start(ln, nil)).To(Succeed())
|
||||
defer func() {
|
||||
Expect(m.Shutdown()).To(Succeed())
|
||||
}()
|
||||
|
||||
provider, err := NewProvider(context.Background(), m.Issuer(), in.skipIssuerVerification)
|
||||
if in.expectedError != "" {
|
||||
Expect(err).To(MatchError(HavePrefix(in.expectedError)))
|
||||
return
|
||||
}
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
|
||||
endpoints := provider.Endpoints()
|
||||
Expect(endpoints.AuthURL).To(Equal(m.AuthorizationEndpoint()))
|
||||
Expect(endpoints.TokenURL).To(Equal(m.TokenEndpoint()))
|
||||
Expect(endpoints.JWKsURL).To(Equal(m.JWKSEndpoint()))
|
||||
Expect(endpoints.UserInfoURL).To(Equal(m.UserinfoEndpoint()))
|
||||
},
|
||||
Entry("with issuer verification and the issuer matches", &newProviderTableInput{
|
||||
skipIssuerVerification: false,
|
||||
}),
|
||||
Entry("with skip issuer verification and the issuer matches", &newProviderTableInput{
|
||||
skipIssuerVerification: true,
|
||||
}),
|
||||
Entry("with issuer verification and an invalid issuer", &newProviderTableInput{
|
||||
skipIssuerVerification: false,
|
||||
middlewares: func(m *mockoidc.MockOIDC) []func(http.Handler) http.Handler {
|
||||
return []func(http.Handler) http.Handler{
|
||||
newInvalidIssuerMiddleware(m),
|
||||
}
|
||||
},
|
||||
expectedError: "oidc: issuer did not match the issuer returned by provider",
|
||||
}),
|
||||
Entry("with skip issuer verification and an invalid issuer", &newProviderTableInput{
|
||||
skipIssuerVerification: true,
|
||||
middlewares: func(m *mockoidc.MockOIDC) []func(http.Handler) http.Handler {
|
||||
return []func(http.Handler) http.Handler{
|
||||
newInvalidIssuerMiddleware(m),
|
||||
}
|
||||
},
|
||||
}),
|
||||
Entry("when the issuer returns a bad response", &newProviderTableInput{
|
||||
skipIssuerVerification: false,
|
||||
middlewares: func(m *mockoidc.MockOIDC) []func(http.Handler) http.Handler {
|
||||
return []func(http.Handler) http.Handler{
|
||||
newBadRequestMiddleware(),
|
||||
}
|
||||
},
|
||||
expectedError: "failed to discover OIDC configuration: unexpected status \"400\"",
|
||||
}),
|
||||
)
|
||||
})
|
||||
|
||||
func newInvalidIssuerMiddleware(m *mockoidc.MockOIDC) func(http.Handler) http.Handler {
|
||||
return func(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(rw http.ResponseWriter, req *http.Request) {
|
||||
p := providerJSON{
|
||||
Issuer: "invalid",
|
||||
AuthURL: m.AuthorizationEndpoint(),
|
||||
TokenURL: m.TokenEndpoint(),
|
||||
JWKsURL: m.JWKSEndpoint(),
|
||||
UserInfoURL: m.UserinfoEndpoint(),
|
||||
}
|
||||
data, err := json.Marshal(p)
|
||||
if err != nil {
|
||||
rw.WriteHeader(500)
|
||||
}
|
||||
rw.Write(data)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func newBadRequestMiddleware() func(http.Handler) http.Handler {
|
||||
return func(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(rw http.ResponseWriter, req *http.Request) {
|
||||
rw.WriteHeader(400)
|
||||
})
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user