internal/oidcprovider/provider_test.go
Ref: Size: 10.7 KiB History
package oidcprovider
import (
"context"
"crypto/sha256"
"encoding/base64"
"encoding/json"
"net/http"
"net/http/httptest"
"net/url"
"path/filepath"
"strings"
"testing"
"time"
oidc "github.com/coreos/go-oidc/v3/oidc"
)
// newTestProvider stands up a provider with one ci user and returns it plus its
// live httptest server (issuer already set to the reachable URL).
func newTestProvider(t *testing.T) (*Provider, *httptest.Server) {
t.Helper()
dir := t.TempDir()
usersPath := filepath.Join(dir, "users.json")
if err := AddUser(usersPath, "ci@example.com", "hunter2hunter2"); err != nil {
t.Fatal(err)
}
p, err := New(Config{
UsersFile: usersPath,
SigningKey: filepath.Join(dir, "signing.key"),
Clients: []Client{{ID: "eitri-console", RedirectURL: "http://client.example/auth/callback"}},
})
if err != nil {
t.Fatal(err)
}
srv := httptest.NewServer(p.Handler())
t.Cleanup(srv.Close)
p.SetIssuer(srv.URL)
return p, srv
}
func pkcePair() (verifier, challenge string) {
verifier = strings.Repeat("v", 43) // any 43-128 char unreserved string
sum := sha256.Sum256([]byte(verifier))
return verifier, base64.RawURLEncoding.EncodeToString(sum[:])
}
func authorizeURL(base, challenge string) string {
return base + "/authorize?" + url.Values{
"response_type": {"code"},
"client_id": {"eitri-console"},
"redirect_uri": {"http://client.example/auth/callback"},
"state": {"st4te"},
"scope": {"openid email"},
"code_challenge": {challenge},
"code_challenge_method": {"S256"},
}.Encode()
}
func noRedirectClient() *http.Client {
return &http.Client{CheckRedirect: func(*http.Request, []*http.Request) error {
return http.ErrUseLastResponse
}}
}
// getCode drives GET+POST /authorize with good credentials and returns the code.
func getCode(t *testing.T, srv *httptest.Server, challenge string) string {
t.Helper()
c := noRedirectClient()
authURL := authorizeURL(srv.URL, challenge)
r, err := c.Get(authURL)
if err != nil || r.StatusCode != 200 {
t.Fatalf("authorize GET: %v status=%d", err, r.StatusCode)
}
r, err = c.PostForm(authURL, url.Values{
"email": {"ci@example.com"}, "password": {"hunter2hunter2"},
})
if err != nil || r.StatusCode != http.StatusFound {
t.Fatalf("authorize POST: %v status=%d", err, r.StatusCode)
}
loc, _ := url.Parse(r.Header.Get("Location"))
if got := loc.Query().Get("state"); got != "st4te" {
t.Fatalf("state = %q", got)
}
return loc.Query().Get("code")
}
// TestCodeFlowAgainstGoOIDC drives the full authorization-code+PKCE flow
// through the real go-oidc verifier — the exact client eitri-server uses.
func TestCodeFlowAgainstGoOIDC(t *testing.T) {
_, srv := newTestProvider(t)
ctx := context.Background()
prov, err := oidc.NewProvider(ctx, srv.URL)
if err != nil {
t.Fatalf("go-oidc discovery: %v", err)
}
verifier := prov.Verifier(&oidc.Config{ClientID: "eitri-console"})
pkceVerifier, challenge := pkcePair()
code := getCode(t, srv, challenge)
tr, err := http.PostForm(srv.URL+"/token", url.Values{
"grant_type": {"authorization_code"},
"code": {code},
"code_verifier": {pkceVerifier},
"client_id": {"eitri-console"},
"redirect_uri": {"http://client.example/auth/callback"},
})
if err != nil || tr.StatusCode != 200 {
t.Fatalf("token: %v status=%d", err, tr.StatusCode)
}
var tokResp struct {
IDToken string `json:"id_token"`
}
if err := json.NewDecoder(tr.Body).Decode(&tokResp); err != nil {
t.Fatal(err)
}
idToken, err := verifier.Verify(ctx, tokResp.IDToken)
if err != nil {
t.Fatalf("verify: %v", err)
}
var claims struct {
Email string `json:"email"`
EmailVerified bool `json:"email_verified"`
}
if err := idToken.Claims(&claims); err != nil || claims.Email != "ci@example.com" {
t.Fatalf("claims: %v %+v", err, claims)
}
if !claims.EmailVerified {
t.Fatalf("email_verified = false, want true")
}
}
func TestWrongPasswordReRendersFormAfterDelay(t *testing.T) {
_, srv := newTestProvider(t)
_, challenge := pkcePair()
c := noRedirectClient()
authURL := authorizeURL(srv.URL, challenge)
start := time.Now()
r, err := c.PostForm(authURL, url.Values{
"email": {"ci@example.com"}, "password": {"wrong"},
})
if err != nil {
t.Fatal(err)
}
if r.StatusCode != http.StatusOK {
t.Fatalf("wrong password status = %d, want 200", r.StatusCode)
}
if loc := r.Header.Get("Location"); loc != "" {
t.Fatalf("wrong password issued redirect: %q", loc)
}
if elapsed := time.Since(start); elapsed < time.Second {
t.Fatalf("no brute-force brake: elapsed %v < 1s", elapsed)
}
}
func TestPKCEMismatchRejected(t *testing.T) {
_, srv := newTestProvider(t)
_, challenge := pkcePair()
code := getCode(t, srv, challenge)
tr, _ := http.PostForm(srv.URL+"/token", url.Values{
"grant_type": {"authorization_code"},
"code": {code},
"code_verifier": {strings.Repeat("x", 43)}, // wrong verifier
"client_id": {"eitri-console"},
"redirect_uri": {"http://client.example/auth/callback"},
})
if tr.StatusCode != http.StatusBadRequest {
t.Fatalf("pkce mismatch token status = %d, want 400", tr.StatusCode)
}
}
func TestUnknownClientRejectedBeforeForm(t *testing.T) {
_, srv := newTestProvider(t)
_, challenge := pkcePair()
u := srv.URL + "/authorize?" + url.Values{
"response_type": {"code"},
"client_id": {"nope"},
"redirect_uri": {"http://client.example/auth/callback"},
"code_challenge": {challenge},
"code_challenge_method": {"S256"},
}.Encode()
r, _ := noRedirectClient().Get(u)
if r.StatusCode != http.StatusBadRequest {
t.Fatalf("unknown client status = %d, want 400", r.StatusCode)
}
}
func TestMismatchedRedirectURIRejected(t *testing.T) {
_, srv := newTestProvider(t)
_, challenge := pkcePair()
u := srv.URL + "/authorize?" + url.Values{
"response_type": {"code"},
"client_id": {"eitri-console"},
"redirect_uri": {"http://evil.example/steal"},
"code_challenge": {challenge},
"code_challenge_method": {"S256"},
}.Encode()
r, _ := noRedirectClient().Get(u)
if r.StatusCode != http.StatusBadRequest {
t.Fatalf("mismatched redirect status = %d, want 400", r.StatusCode)
}
}
func TestCodeSingleUse(t *testing.T) {
_, srv := newTestProvider(t)
verifier, challenge := pkcePair()
code := getCode(t, srv, challenge)
form := url.Values{
"grant_type": {"authorization_code"},
"code": {code},
"code_verifier": {verifier},
"client_id": {"eitri-console"},
"redirect_uri": {"http://client.example/auth/callback"},
}
if r, _ := http.PostForm(srv.URL+"/token", form); r.StatusCode != 200 {
t.Fatalf("first exchange status = %d", r.StatusCode)
}
if r, _ := http.PostForm(srv.URL+"/token", form); r.StatusCode != http.StatusBadRequest {
t.Fatalf("code reuse status = %d, want 400", r.StatusCode)
}
}
func TestTokenBasicAuthClientID(t *testing.T) {
// x/oauth2's AuthStyleAutoDetect sends client credentials as HTTP Basic
// first; the exchange must succeed on that very request.
_, srv := newTestProvider(t)
verifier, challenge := pkcePair()
code := getCode(t, srv, challenge)
form := url.Values{
"grant_type": {"authorization_code"},
"code": {code},
"code_verifier": {verifier},
"redirect_uri": {"http://client.example/auth/callback"},
}
req, err := http.NewRequest(http.MethodPost, srv.URL+"/token", strings.NewReader(form.Encode()))
if err != nil {
t.Fatal(err)
}
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
req.SetBasicAuth("eitri-console", "")
r, err := http.DefaultClient.Do(req)
if err != nil {
t.Fatal(err)
}
if r.StatusCode != 200 {
t.Fatalf("basic-auth exchange status = %d, want 200", r.StatusCode)
}
}
func TestWrongClientProbeDoesNotBurnCode(t *testing.T) {
_, srv := newTestProvider(t)
verifier, challenge := pkcePair()
code := getCode(t, srv, challenge)
form := url.Values{
"grant_type": {"authorization_code"},
"code": {code},
"code_verifier": {verifier},
"client_id": {"not-the-console"},
"redirect_uri": {"http://client.example/auth/callback"},
}
if r, _ := http.PostForm(srv.URL+"/token", form); r.StatusCode != http.StatusBadRequest {
t.Fatalf("wrong-client exchange status = %d, want 400", r.StatusCode)
}
// The mismatched attempt must not have consumed the code.
form.Set("client_id", "eitri-console")
if r, _ := http.PostForm(srv.URL+"/token", form); r.StatusCode != 200 {
t.Fatalf("retry after wrong-client probe status = %d, want 200", r.StatusCode)
}
}
func TestPKCEFailureBurnsCode(t *testing.T) {
// A failed verifier is a real exchange attempt: single-use must hold, or
// an attacker could brute-force verifiers against one stolen code.
_, srv := newTestProvider(t)
_, challenge := pkcePair()
code := getCode(t, srv, challenge)
form := url.Values{
"grant_type": {"authorization_code"},
"code": {code},
"code_verifier": {strings.Repeat("x", 43)},
"client_id": {"eitri-console"},
"redirect_uri": {"http://client.example/auth/callback"},
}
if r, _ := http.PostForm(srv.URL+"/token", form); r.StatusCode != http.StatusBadRequest {
t.Fatalf("bad-verifier exchange status = %d, want 400", r.StatusCode)
}
if r, _ := http.PostForm(srv.URL+"/token", form); r.StatusCode != http.StatusBadRequest {
t.Fatalf("reuse after PKCE failure status = %d, want 400 invalid code", r.StatusCode)
}
}
func TestExpiredCodeRejected(t *testing.T) {
p, srv := newTestProvider(t)
// Freeze the clock 3 minutes in the past when minting so the code (2m TTL)
// is already stale by exchange time.
p.now = func() time.Time { return time.Now().Add(-3 * time.Minute) }
verifier, challenge := pkcePair()
code := getCode(t, srv, challenge)
p.now = time.Now
r, _ := http.PostForm(srv.URL+"/token", url.Values{
"grant_type": {"authorization_code"},
"code": {code},
"code_verifier": {verifier},
"client_id": {"eitri-console"},
"redirect_uri": {"http://client.example/auth/callback"},
})
if r.StatusCode != http.StatusBadRequest {
t.Fatalf("expired code status = %d, want 400", r.StatusCode)
}
}
func TestDiscoveryDocument(t *testing.T) {
_, srv := newTestProvider(t)
r, err := http.Get(srv.URL + "/.well-known/openid-configuration")
if err != nil || r.StatusCode != 200 {
t.Fatalf("discovery: %v status=%d", err, r.StatusCode)
}
var doc map[string]any
if err := json.NewDecoder(r.Body).Decode(&doc); err != nil {
t.Fatal(err)
}
if doc["issuer"] != srv.URL {
t.Fatalf("issuer = %v", doc["issuer"])
}
if doc["authorization_endpoint"] != srv.URL+"/authorize" {
t.Fatalf("authorization_endpoint = %v", doc["authorization_endpoint"])
}
if doc["token_endpoint"] != srv.URL+"/token" {
t.Fatalf("token_endpoint = %v", doc["token_endpoint"])
}
if doc["jwks_uri"] != srv.URL+"/jwks.json" {
t.Fatalf("jwks_uri = %v", doc["jwks_uri"])
}
}