a73x

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"])
	}
}