package agent import ( "encoding/base64" "encoding/json" "io" "net/http" "net/http/httptest" "strconv" "strings" "testing" ) const ( oauthPath = "kubernetes/namespace/repospawner/default/oauth-credentials" oauthClientID = "mediamark-client-id" existingClientSec = "existing-client-secret-from-authentik" existingCookieSec = "existing-cookie-secret-value-abcdefghij" oauthExtraKeyValue = "extra-key-secret-value" oauthVaultClientTok = "s.vaulttoken" ) // oauthVaultStub is a KV-v2 stand-in that actually stores what is written, so // read-modify-write behaviour can be asserted end to end. type oauthVaultStub struct { data map[string]any exists bool version int readStatus int writeStatus int writes []map[string]any } func newOAuthVaultStub() *oauthVaultStub { return &oauthVaultStub{readStatus: http.StatusOK, writeStatus: http.StatusOK} } // seed makes the path exist with the given fields at version 1. func (v *oauthVaultStub) seed(data map[string]any) *oauthVaultStub { v.data = data v.exists = true v.version = 1 return v } func (v *oauthVaultStub) server(t *testing.T) *httptest.Server { t.Helper() mux := http.NewServeMux() mux.HandleFunc("/v1/auth/approle/login", func(w http.ResponseWriter, r *http.Request) { var body map[string]string _ = json.NewDecoder(r.Body).Decode(&body) if _, ok := body["secret_id"]; ok { t.Errorf("secret_id must not be sent") } _, _ = io.WriteString(w, `{"auth":{"client_token":"`+oauthVaultClientTok+`"}}`) }) mux.HandleFunc("/v1/kv/data/"+oauthPath, func(w http.ResponseWriter, r *http.Request) { if got := r.Header.Get("X-Vault-Token"); got != oauthVaultClientTok { t.Errorf("X-Vault-Token = %q, want %q", got, oauthVaultClientTok) } switch r.Method { case http.MethodGet: if v.readStatus != http.StatusOK { w.WriteHeader(v.readStatus) _, _ = io.WriteString(w, `{"errors":["permission denied"]}`) return } if !v.exists { w.WriteHeader(http.StatusNotFound) _, _ = io.WriteString(w, `{"errors":[]}`) return } payload, _ := json.Marshal(map[string]any{ "data": map[string]any{"data": v.data, "metadata": map[string]any{"version": v.version}}, }) _, _ = w.Write(payload) case http.MethodPost: if v.writeStatus != http.StatusOK { w.WriteHeader(v.writeStatus) _, _ = io.WriteString(w, `{"errors":["permission denied"]}`) return } var body struct { Data map[string]any `json:"data"` } _ = json.NewDecoder(r.Body).Decode(&body) v.writes = append(v.writes, body.Data) v.data = body.Data v.exists = true v.version++ _, _ = io.WriteString(w, `{"data":{"version":`+strconv.Itoa(v.version)+`}}`) default: t.Errorf("unexpected method %s", r.Method) } }) srv := httptest.NewServer(mux) t.Cleanup(srv.Close) return srv } func oauthOpts(vaultURL string) SeedOAuthOptions { return SeedOAuthOptions{ VaultAddr: vaultURL, RoleID: "role-xyz", KVMount: DefaultKVMount, Path: oauthPath, ClientID: oauthClientID, } } // actions flattens a result into key -> action for order-independent asserts. func actions(res SeedOAuthResult) map[string]string { m := make(map[string]string, len(res.Keys)) for _, k := range res.Keys { m[k.Name] = k.Action } return m } func stringField(t *testing.T, data map[string]any, key string) string { t.Helper() s, ok := data[key].(string) if !ok { t.Fatalf("written %s = %v, want a string", key, data[key]) } return s } // assertDecodesTo32 fails unless the value is base64 of exactly 32 bytes, which // is what oauth2-proxy requires of a cookie secret. func assertDecodesTo32(t *testing.T, enc *base64.Encoding, value, name string) { t.Helper() raw, err := enc.DecodeString(value) if err != nil { t.Fatalf("%s is not valid base64: %v", name, err) } if len(raw) != oauthSecretBytes { t.Errorf("%s decodes to %d bytes, want %d", name, len(raw), oauthSecretBytes) } } func TestSeedOAuthFreshCreate(t *testing.T) { v := newOAuthVaultStub() res, err := SeedOAuth(oauthOpts(v.server(t).URL)) if err != nil { t.Fatalf("SeedOAuth: %v", err) } if !res.Changed || res.Version != 1 { t.Errorf("Changed=%v Version=%d, want a first write at version 1", res.Changed, res.Version) } for key, want := range map[string]string{ OAuthClientIDKey: ActionCreated, OAuthClientSecretKey: ActionCreated, OAuthCookieSecretKey: ActionCreated, } { if got := actions(res)[key]; got != want { t.Errorf("%s action = %q, want %q", key, got, want) } } if len(v.writes) != 1 { t.Fatalf("%d writes, want exactly 1", len(v.writes)) } w := v.writes[0] if got := stringField(t, w, OAuthClientIDKey); got != oauthClientID { t.Errorf("written client_id = %q, want %q", got, oauthClientID) } assertDecodesTo32(t, base64.StdEncoding, stringField(t, w, OAuthClientSecretKey), OAuthClientSecretKey) assertDecodesTo32(t, base64.RawURLEncoding, stringField(t, w, OAuthCookieSecretKey), OAuthCookieSecretKey) } // The mediamark case: a client_secret already issued by Authentik must survive // while the missing keys are filled in. func TestSeedOAuthPreservesExistingClientSecret(t *testing.T) { v := newOAuthVaultStub().seed(map[string]any{OAuthClientSecretKey: existingClientSec}) res, err := SeedOAuth(oauthOpts(v.server(t).URL)) if err != nil { t.Fatalf("SeedOAuth: %v", err) } got := actions(res) for key, want := range map[string]string{ OAuthClientIDKey: ActionCreated, OAuthClientSecretKey: ActionKept, OAuthCookieSecretKey: ActionCreated, } { if got[key] != want { t.Errorf("%s action = %q, want %q", key, got[key], want) } } if len(v.writes) != 1 { t.Fatalf("%d writes, want exactly 1", len(v.writes)) } if s := stringField(t, v.writes[0], OAuthClientSecretKey); s != existingClientSec { t.Errorf("client_secret was replaced, want the existing value kept") } } func TestSeedOAuthPreservesOtherKeys(t *testing.T) { v := newOAuthVaultStub().seed(map[string]any{ OAuthClientIDKey: oauthClientID, OAuthClientSecretKey: existingClientSec, "redirect_url": "https://mediamark.unkin.net/oauth2/callback", "extra": oauthExtraKeyValue, }) res, err := SeedOAuth(oauthOpts(v.server(t).URL)) if err != nil { t.Fatalf("SeedOAuth: %v", err) } got := actions(res) for _, key := range []string{"redirect_url", "extra"} { if got[key] != ActionPreserved { t.Errorf("%s action = %q, want %q", key, got[key], ActionPreserved) } } if len(v.writes) != 1 { t.Fatalf("%d writes, want exactly 1", len(v.writes)) } w := v.writes[0] if stringField(t, w, "extra") != oauthExtraKeyValue { t.Errorf("extra key was not written back unchanged") } if stringField(t, w, "redirect_url") != "https://mediamark.unkin.net/oauth2/callback" { t.Errorf("redirect_url was not written back unchanged") } } func TestSeedOAuthRotateRegenerates(t *testing.T) { v := newOAuthVaultStub().seed(map[string]any{ OAuthClientIDKey: oauthClientID, OAuthClientSecretKey: existingClientSec, OAuthCookieSecretKey: existingCookieSec, }) o := oauthOpts(v.server(t).URL) o.Rotate = true res, err := SeedOAuth(o) if err != nil { t.Fatalf("SeedOAuth: %v", err) } got := actions(res) for key, want := range map[string]string{ OAuthClientIDKey: ActionKept, OAuthClientSecretKey: ActionRotated, OAuthCookieSecretKey: ActionRotated, } { if got[key] != want { t.Errorf("%s action = %q, want %q", key, got[key], want) } } if len(v.writes) != 1 { t.Fatalf("%d writes, want exactly 1", len(v.writes)) } w := v.writes[0] if stringField(t, w, OAuthClientSecretKey) == existingClientSec { t.Errorf("client_secret unchanged under --rotate") } if stringField(t, w, OAuthCookieSecretKey) == existingCookieSec { t.Errorf("cookie_secret unchanged under --rotate") } assertDecodesTo32(t, base64.RawURLEncoding, stringField(t, w, OAuthCookieSecretKey), OAuthCookieSecretKey) } // A complete, correct secret must not produce a new KV version. func TestSeedOAuthIdempotentWritesNothing(t *testing.T) { v := newOAuthVaultStub() url := v.server(t).URL if _, err := SeedOAuth(oauthOpts(url)); err != nil { t.Fatalf("first SeedOAuth: %v", err) } res, err := SeedOAuth(oauthOpts(url)) if err != nil { t.Fatalf("second SeedOAuth: %v", err) } if res.Changed || res.Version != 0 { t.Errorf("Changed=%v Version=%d, want an unchanged result", res.Changed, res.Version) } if len(v.writes) != 1 { t.Errorf("%d writes, want the second run to write nothing", len(v.writes)) } for _, k := range res.Keys { if k.Action != ActionKept { t.Errorf("%s action = %q, want %q", k.Name, k.Action, ActionKept) } } } func TestSeedOAuthClientIDUpdated(t *testing.T) { v := newOAuthVaultStub().seed(map[string]any{ OAuthClientIDKey: "stale-client-id", OAuthClientSecretKey: existingClientSec, OAuthCookieSecretKey: existingCookieSec, }) res, err := SeedOAuth(oauthOpts(v.server(t).URL)) if err != nil { t.Fatalf("SeedOAuth: %v", err) } if got := actions(res)[OAuthClientIDKey]; got != ActionUpdated { t.Errorf("client_id action = %q, want %q", got, ActionUpdated) } if len(v.writes) != 1 || stringField(t, v.writes[0], OAuthClientIDKey) != oauthClientID { t.Errorf("writes = %v, want the new client_id written", v.writes) } } func TestSeedOAuthLoginFailure(t *testing.T) { mux := http.NewServeMux() mux.HandleFunc("/v1/auth/approle/login", func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusBadRequest) _, _ = io.WriteString(w, `{"errors":["invalid role ID"]}`) }) vs := httptest.NewServer(mux) defer vs.Close() _, err := SeedOAuth(oauthOpts(vs.URL)) if err == nil { t.Fatal("SeedOAuth() = nil, want an approle login error") } if !strings.Contains(err.Error(), "approle login failed") { t.Errorf("error = %v, want it to name the approle login", err) } } func TestSeedOAuthReadDenied(t *testing.T) { v := newOAuthVaultStub() v.readStatus = http.StatusForbidden _, err := SeedOAuth(oauthOpts(v.server(t).URL)) if err == nil { t.Fatal("SeedOAuth() = nil, want a KV read error") } msg := err.Error() if !strings.Contains(msg, oauthPath) || !strings.Contains(msg, "policy") { t.Errorf("error = %v, want it to name the path and point at the policy", err) } if len(v.writes) != 0 { t.Errorf("wrote %v, want no write when the read is denied", v.writes) } } func TestSeedOAuthWriteDenied(t *testing.T) { v := newOAuthVaultStub() v.writeStatus = http.StatusForbidden _, err := SeedOAuth(oauthOpts(v.server(t).URL)) if err == nil { t.Fatal("SeedOAuth() = nil, want a KV write error") } msg := err.Error() if !strings.Contains(msg, oauthPath) || !strings.Contains(msg, "create/update") { t.Errorf("error = %v, want it to name the path and the missing capability", err) } } // A missing path is normal (first seed), not a not-found error. func TestSeedOAuthMissingPathIsNotAnError(t *testing.T) { v := newOAuthVaultStub() if _, err := SeedOAuth(oauthOpts(v.server(t).URL)); err != nil { t.Fatalf("SeedOAuth on a missing path: %v", err) } } func TestSeedOAuthRequiresPathAndClientID(t *testing.T) { v := newOAuthVaultStub() url := v.server(t).URL for name, mutate := range map[string]func(*SeedOAuthOptions){ "no path": func(o *SeedOAuthOptions) { o.Path = "" }, "no client id": func(o *SeedOAuthOptions) { o.ClientID = "" }, } { t.Run(name, func(t *testing.T) { o := oauthOpts(url) mutate(&o) if _, err := SeedOAuth(o); err == nil { t.Fatal("SeedOAuth() = nil, want a required-input error") } }) } } // No failure path may leak stored or generated secret material. func TestSeedOAuthErrorsNeverLeakSecrets(t *testing.T) { cases := map[string]func(*oauthVaultStub){ "read denied": func(v *oauthVaultStub) { v.readStatus = http.StatusForbidden }, "write denied": func(v *oauthVaultStub) { v.writeStatus = http.StatusForbidden }, "read error": func(v *oauthVaultStub) { v.readStatus = http.StatusInternalServerError }, "write error": func(v *oauthVaultStub) { v.writeStatus = http.StatusInternalServerError }, } for name, mutate := range cases { t.Run(name, func(t *testing.T) { v := newOAuthVaultStub().seed(map[string]any{ OAuthClientSecretKey: existingClientSec, OAuthCookieSecretKey: existingCookieSec, "extra": oauthExtraKeyValue, }) mutate(v) _, err := SeedOAuth(oauthOpts(v.server(t).URL)) if err == nil { t.Fatal("SeedOAuth() = nil, want an error") } for _, secret := range []string{existingClientSec, existingCookieSec, oauthExtraKeyValue} { if strings.Contains(err.Error(), secret) { t.Errorf("error %q leaks a secret", err) } } }) } } // The successful result carries key names and a version, never values. func TestSeedOAuthResultNeverCarriesSecrets(t *testing.T) { v := newOAuthVaultStub().seed(map[string]any{OAuthClientSecretKey: existingClientSec}) res, err := SeedOAuth(oauthOpts(v.server(t).URL)) if err != nil { t.Fatalf("SeedOAuth: %v", err) } rendered := strings.Join(append(res.KeyNames(), res.Path, res.KVMount), " ") written := v.writes[0] for _, key := range []string{OAuthClientSecretKey, OAuthCookieSecretKey} { if value := stringField(t, written, key); strings.Contains(rendered, value) { t.Errorf("result leaks the %s value", key) } } }