Files
dokkku-google-auth/internal/authproxy/flow_test.go
T
2026-07-15 07:49:11 -04:00

386 lines
12 KiB
Go

package authproxy
import (
"encoding/base64"
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"net/url"
"strings"
"testing"
"time"
)
const (
testClientID = "test-client-id.apps.googleusercontent.com"
appHost = "myapp.example.com"
authHost = "auth.example.com"
)
// fakeGoogle stands in for Google's token endpoint. The returned id_token
// carries the given claims; signature contents don't matter because the
// token arrives over a direct TLS channel in production.
func fakeGoogle(t *testing.T, claims map[string]any) *httptest.Server {
t.Helper()
return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if err := r.ParseForm(); err != nil {
t.Errorf("parse form: %v", err)
}
if r.FormValue("grant_type") != "authorization_code" {
t.Errorf("unexpected grant_type %q", r.FormValue("grant_type"))
}
if r.FormValue("code") != "good-code" {
w.WriteHeader(http.StatusBadRequest)
fmt.Fprint(w, `{"error":"invalid_grant"}`)
return
}
payload, _ := json.Marshal(claims)
idt := b64(`{"alg":"RS256","typ":"JWT"}`) + "." +
base64.RawURLEncoding.EncodeToString(payload) + "." + b64("sig")
w.Header().Set("Content-Type", "application/json")
fmt.Fprintf(w, `{"access_token":"at","id_token":%q}`, idt)
}))
}
func b64(s string) string {
return base64.RawURLEncoding.EncodeToString([]byte(s))
}
func goodClaims() map[string]any {
return map[string]any{
"iss": "https://accounts.google.com",
"aud": testClientID,
"sub": "1234567890",
"exp": time.Now().Add(time.Hour).Unix(),
"email": "Greyson@Signal.org",
"email_verified": true,
"hd": "signal.org",
"name": "Greyson",
}
}
func newTestServer(t *testing.T, tokenURL string) *Server {
t.Helper()
cfg := Config{
ClientID: testClientID,
ClientSecret: "test-secret",
CookieSecret: strings.Repeat("k", 32),
AuthHost: authHost,
AllowedDomains: []string{"signal.org"},
CookieName: "_google_auth",
SessionTTL: time.Hour,
ListenAddr: ":0",
AuthorizeURL: "https://accounts.google.com/o/oauth2/v2/auth",
TokenURL: tokenURL,
}
if err := cfg.validate(); err != nil {
t.Fatalf("test config invalid: %v", err)
}
s, err := New(cfg)
if err != nil {
t.Fatal(err)
}
return s
}
func do(h http.Handler, r *http.Request) *httptest.ResponseRecorder {
w := httptest.NewRecorder()
h.ServeHTTP(w, r)
return w
}
// TestFullFlow walks the whole journey: unauthenticated verify, start,
// callback on the auth host, finish on the app host, authenticated verify.
func TestFullFlow(t *testing.T) {
google := fakeGoogle(t, goodClaims())
defer google.Close()
s := newTestServer(t, google.URL)
h := s.Routes()
// 1. verify without a cookie → 401
r := httptest.NewRequest("GET", "http://"+appHost+RoutePrefix+"/verify", nil)
if w := do(h, r); w.Code != http.StatusUnauthorized {
t.Fatalf("verify without cookie: got %d, want 401", w.Code)
}
// 2. start (as nginx would proxy it) → 302 to Google
origURI := "/secret/page?a=1&b=2"
r = httptest.NewRequest("GET", "http://"+appHost+RoutePrefix+"/start", nil)
r.Header.Set("Accept", "text/html,application/xhtml+xml")
r.Header.Set("X-Auth-Request-Redirect", origURI)
r.Header.Set("X-Forwarded-Proto", "https")
w := do(h, r)
if w.Code != http.StatusFound {
t.Fatalf("start: got %d, want 302", w.Code)
}
loc, err := url.Parse(w.Header().Get("Location"))
if err != nil {
t.Fatal(err)
}
if loc.Host != "accounts.google.com" {
t.Fatalf("start redirected to %s, want accounts.google.com", loc.Host)
}
if got := loc.Query().Get("client_id"); got != testClientID {
t.Fatalf("client_id = %q", got)
}
if got := loc.Query().Get("redirect_uri"); got != "https://"+authHost+RoutePrefix+"/callback" {
t.Fatalf("redirect_uri = %q", got)
}
if got := loc.Query().Get("hd"); got != "signal.org" {
t.Fatalf("hd = %q", got)
}
state := loc.Query().Get("state")
if state == "" {
t.Fatal("no state in Google redirect")
}
// 3. callback on the auth host → 302 to finish on the app host
r = httptest.NewRequest("GET",
"http://"+authHost+RoutePrefix+"/callback?code=good-code&state="+url.QueryEscape(state), nil)
w = do(h, r)
if w.Code != http.StatusFound {
t.Fatalf("callback: got %d, want 302; body: %s", w.Code, w.Body.String())
}
finishURL, err := url.Parse(w.Header().Get("Location"))
if err != nil {
t.Fatal(err)
}
if finishURL.Scheme != "https" || finishURL.Host != appHost || finishURL.Path != RoutePrefix+"/finish" {
t.Fatalf("callback redirected to %s", finishURL)
}
// 4. finish on the app host → session cookie + redirect to original URI
r = httptest.NewRequest("GET", finishURL.String(), nil)
w = do(h, r)
if w.Code != http.StatusFound {
t.Fatalf("finish: got %d, want 302; body: %s", w.Code, w.Body.String())
}
if got := w.Header().Get("Location"); got != origURI {
t.Fatalf("finish redirected to %q, want %q", got, origURI)
}
cookies := w.Result().Cookies()
if len(cookies) != 1 || cookies[0].Name != "_google_auth" {
t.Fatalf("expected one session cookie, got %v", cookies)
}
sessionCookie := cookies[0]
if !sessionCookie.Secure || !sessionCookie.HttpOnly {
t.Fatalf("session cookie should be Secure+HttpOnly: %+v", sessionCookie)
}
// 5. verify with the cookie → 200 with identity headers
r = httptest.NewRequest("GET", "http://"+appHost+RoutePrefix+"/verify", nil)
r.AddCookie(sessionCookie)
w = do(h, r)
if w.Code != http.StatusOK {
t.Fatalf("verify with cookie: got %d, want 200", w.Code)
}
if got := w.Header().Get("X-Auth-Request-Email"); got != "greyson@signal.org" {
t.Fatalf("X-Auth-Request-Email = %q", got)
}
if got := w.Header().Get("X-Auth-Request-User"); got != "1234567890" {
t.Fatalf("X-Auth-Request-User = %q", got)
}
// 6. the same cookie must NOT work on a different host
r = httptest.NewRequest("GET", "http://other.example.com"+RoutePrefix+"/verify", nil)
r.AddCookie(sessionCookie)
if w := do(h, r); w.Code != http.StatusUnauthorized {
t.Fatalf("verify on wrong host: got %d, want 401", w.Code)
}
// 7. replaying the finish token must fail (single-use nonce)
r = httptest.NewRequest("GET", finishURL.String(), nil)
if w := do(h, r); w.Code != http.StatusForbidden {
t.Fatalf("finish replay: got %d, want 403", w.Code)
}
// 8. logout clears the cookie
r = httptest.NewRequest("GET", "http://"+appHost+RoutePrefix+"/logout", nil)
r.AddCookie(sessionCookie)
w = do(h, r)
found := false
for _, c := range w.Result().Cookies() {
if c.Name == "_google_auth" && c.MaxAge < 0 {
found = true
}
}
if !found {
t.Fatal("logout did not clear the session cookie")
}
}
func TestCallbackRejectsDisallowedDomain(t *testing.T) {
claims := goodClaims()
claims["email"] = "intruder@evil.com"
claims["hd"] = "evil.com"
google := fakeGoogle(t, claims)
defer google.Close()
s := newTestServer(t, google.URL)
h := s.Routes()
state := mustState(t, s, appHost, "/")
r := httptest.NewRequest("GET",
"http://"+authHost+RoutePrefix+"/callback?code=good-code&state="+url.QueryEscape(state), nil)
w := do(h, r)
if w.Code != http.StatusForbidden {
t.Fatalf("got %d, want 403", w.Code)
}
if !strings.Contains(w.Body.String(), "not allowed") {
t.Fatalf("body should explain denial: %s", w.Body.String())
}
}
func TestCallbackRejectsWrongAudience(t *testing.T) {
claims := goodClaims()
claims["aud"] = "someone-else"
google := fakeGoogle(t, claims)
defer google.Close()
s := newTestServer(t, google.URL)
state := mustState(t, s, appHost, "/")
r := httptest.NewRequest("GET",
"http://"+authHost+RoutePrefix+"/callback?code=good-code&state="+url.QueryEscape(state), nil)
if w := do(s.Routes(), r); w.Code != http.StatusForbidden {
t.Fatalf("got %d, want 403", w.Code)
}
}
func TestCallbackRejectsUnverifiedEmail(t *testing.T) {
claims := goodClaims()
claims["email_verified"] = false
google := fakeGoogle(t, claims)
defer google.Close()
s := newTestServer(t, google.URL)
state := mustState(t, s, appHost, "/")
r := httptest.NewRequest("GET",
"http://"+authHost+RoutePrefix+"/callback?code=good-code&state="+url.QueryEscape(state), nil)
if w := do(s.Routes(), r); w.Code != http.StatusForbidden {
t.Fatalf("got %d, want 403", w.Code)
}
}
func TestCallbackOnlyServedOnAuthHost(t *testing.T) {
s := newTestServer(t, "http://unused.invalid")
r := httptest.NewRequest("GET", "http://"+appHost+RoutePrefix+"/callback?code=x&state=y", nil)
if w := do(s.Routes(), r); w.Code != http.StatusNotFound {
t.Fatalf("got %d, want 404", w.Code)
}
}
func TestCallbackRejectsExpiredState(t *testing.T) {
s := newTestServer(t, "http://unused.invalid")
st := stateClaims{Host: appHost, RD: "/", Proto: "https", Nonce: "n",
Exp: time.Now().Add(-time.Minute).Unix()}
tok, err := s.box.seal("state", st)
if err != nil {
t.Fatal(err)
}
r := httptest.NewRequest("GET",
"http://"+authHost+RoutePrefix+"/callback?code=good-code&state="+url.QueryEscape(tok), nil)
if w := do(s.Routes(), r); w.Code != http.StatusForbidden {
t.Fatalf("got %d, want 403", w.Code)
}
}
func TestFinishRejectsWrongHost(t *testing.T) {
s := newTestServer(t, "http://unused.invalid")
hand := handoffClaims{Email: "a@signal.org", User: "1", Host: appHost, RD: "/",
Proto: "https", Nonce: randToken(), Exp: time.Now().Add(time.Minute).Unix()}
tok, _ := s.box.seal("handoff", hand)
r := httptest.NewRequest("GET",
"http://other.example.com"+RoutePrefix+"/finish?token="+url.QueryEscape(tok), nil)
if w := do(s.Routes(), r); w.Code != http.StatusForbidden {
t.Fatalf("got %d, want 403", w.Code)
}
}
func TestStartReturnsJSON401ForNonBrowsers(t *testing.T) {
s := newTestServer(t, "http://unused.invalid")
r := httptest.NewRequest("GET", "http://"+appHost+RoutePrefix+"/start", nil)
r.Header.Set("Accept", "application/json")
w := do(s.Routes(), r)
if w.Code != http.StatusUnauthorized {
t.Fatalf("got %d, want 401", w.Code)
}
if ct := w.Header().Get("Content-Type"); !strings.Contains(ct, "application/json") {
t.Fatalf("content type = %q", ct)
}
}
func TestExpiredSessionRejected(t *testing.T) {
s := newTestServer(t, "http://unused.invalid")
sess := sessionClaims{Email: "a@signal.org", User: "1", Host: appHost,
Exp: time.Now().Add(-time.Minute).Unix()}
val, _ := s.box.seal("session", sess)
r := httptest.NewRequest("GET", "http://"+appHost+RoutePrefix+"/verify", nil)
r.AddCookie(&http.Cookie{Name: "_google_auth", Value: val})
if w := do(s.Routes(), r); w.Code != http.StatusUnauthorized {
t.Fatalf("got %d, want 401", w.Code)
}
}
func TestSanitizeRedirect(t *testing.T) {
cases := map[string]string{
"": "/",
"/ok": "/ok",
"/ok?a=1&b=2": "/ok?a=1&b=2",
"//evil.com/x": "/",
"https://evil.com": "/",
"/x\r\nSet-Cookie: p": "/",
"\\evil": "/",
RoutePrefix + "/start": "/", // avoid redirect loops into our own routes
"relative/no/lead/slash": "/",
"/deep/path/./is/fine": "/deep/path/./is/fine",
}
for in, want := range cases {
if got := sanitizeRedirect(in); got != want {
t.Errorf("sanitizeRedirect(%q) = %q, want %q", in, got, want)
}
}
}
func TestEmailAllowed(t *testing.T) {
s := newTestServer(t, "http://unused.invalid")
s.cfg.AllowedEmails = []string{"guest@partner.com"}
cases := map[string]bool{
"greyson@signal.org": true,
"GREYSON@SIGNAL.ORG": true,
"guest@partner.com": true,
"other@partner.com": false,
"evil@notsignal.org": false,
"greyson@signal.org.evil.c": false,
"signal.org": false,
}
for email, want := range cases {
if got := s.emailAllowed(email); got != want {
t.Errorf("emailAllowed(%q) = %v, want %v", email, got, want)
}
}
}
func TestFlexClaims(t *testing.T) {
var tok idToken
payload := `{"aud":["a","b"],"email_verified":"true","iss":"accounts.google.com","exp":99}`
if err := json.Unmarshal([]byte(payload), &tok); err != nil {
t.Fatal(err)
}
if !tok.Aud.contains("b") || !bool(tok.EmailVerified) {
t.Fatalf("flex claims parsed wrong: %+v", tok)
}
}
// mustState mints a valid state token the way handleStart would.
func mustState(t *testing.T, s *Server, host, rd string) string {
t.Helper()
tok, err := s.box.seal("state", stateClaims{
Host: host, RD: rd, Proto: "https", Nonce: randToken(),
Exp: time.Now().Add(10 * time.Minute).Unix(),
})
if err != nil {
t.Fatal(err)
}
return tok
}