140 lines
4.5 KiB
Go
140 lines
4.5 KiB
Go
package auth
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"strings"
|
|
"testing"
|
|
"time"
|
|
)
|
|
|
|
func TestInMemoryOAuthStore_InsertAndGetAndDelete(t *testing.T) {
|
|
st := NewInMemoryOAuthStore()
|
|
ctx := context.Background()
|
|
|
|
in := OAuthSession{
|
|
State: "state-1",
|
|
InstallID: "42",
|
|
RedirectURI: "https://example.com/cb",
|
|
ProviderState: []byte(`{"pkce_verifier":"abc"}`),
|
|
NextURL: "/me",
|
|
ExpiresAt: time.Now().Add(10 * time.Minute),
|
|
}
|
|
if err := st.Insert(ctx, in); err != nil {
|
|
t.Fatalf("Insert: %v", err)
|
|
}
|
|
|
|
out, err := st.GetAndDelete(ctx, "state-1")
|
|
if err != nil {
|
|
t.Fatalf("GetAndDelete: %v", err)
|
|
}
|
|
if out.InstallID != "42" || out.RedirectURI != in.RedirectURI || out.NextURL != "/me" {
|
|
t.Errorf("out = %+v", out)
|
|
}
|
|
if string(out.ProviderState) != `{"pkce_verifier":"abc"}` {
|
|
t.Errorf("ProviderState = %s", out.ProviderState)
|
|
}
|
|
|
|
if _, err := st.GetAndDelete(ctx, "state-1"); !errors.Is(err, ErrOAuthSessionNotFound) {
|
|
t.Errorf("expected ErrOAuthSessionNotFound, got %v", err)
|
|
}
|
|
}
|
|
|
|
func TestInMemoryOAuthStore_DefaultsAndValidation(t *testing.T) {
|
|
st := NewInMemoryOAuthStore()
|
|
ctx := context.Background()
|
|
|
|
// Missing state — error.
|
|
if err := st.Insert(ctx, OAuthSession{InstallID: "1", RedirectURI: "/x", ExpiresAt: time.Now().Add(time.Hour)}); err == nil {
|
|
t.Error("Insert should reject empty state")
|
|
}
|
|
// Missing install_id — error.
|
|
if err := st.Insert(ctx, OAuthSession{State: "s", RedirectURI: "/x", ExpiresAt: time.Now().Add(time.Hour)}); err == nil {
|
|
t.Error("Insert should reject empty install_id")
|
|
}
|
|
// Missing expires_at — error.
|
|
if err := st.Insert(ctx, OAuthSession{State: "s", InstallID: "1", RedirectURI: "/x"}); err == nil {
|
|
t.Error("Insert should reject zero expires_at")
|
|
}
|
|
// Missing redirect_uri — error.
|
|
if err := st.Insert(ctx, OAuthSession{State: "s", InstallID: "1", ExpiresAt: time.Now().Add(time.Hour)}); err == nil {
|
|
t.Error("Insert should reject empty redirect_uri")
|
|
}
|
|
|
|
// next_url and provider_state default.
|
|
if err := st.Insert(ctx, OAuthSession{State: "s2", InstallID: "1", RedirectURI: "/x", ExpiresAt: time.Now().Add(time.Hour)}); err != nil {
|
|
t.Fatalf("Insert: %v", err)
|
|
}
|
|
out, err := st.GetAndDelete(ctx, "s2")
|
|
if err != nil {
|
|
t.Fatalf("GetAndDelete: %v", err)
|
|
}
|
|
if out.NextURL != "/" {
|
|
t.Errorf("NextURL default = %q, want /", out.NextURL)
|
|
}
|
|
if string(out.ProviderState) != "{}" {
|
|
t.Errorf("ProviderState default = %s, want {}", out.ProviderState)
|
|
}
|
|
}
|
|
|
|
func TestInMemoryOAuthStore_DeleteExpired(t *testing.T) {
|
|
st := NewInMemoryOAuthStore()
|
|
ctx := context.Background()
|
|
now := time.Now()
|
|
|
|
_ = st.Insert(ctx, OAuthSession{State: "old", InstallID: "i", RedirectURI: "/", ExpiresAt: now.Add(-time.Hour)})
|
|
_ = st.Insert(ctx, OAuthSession{State: "new", InstallID: "i", RedirectURI: "/", ExpiresAt: now.Add(time.Hour)})
|
|
|
|
deleted, err := st.DeleteExpired(ctx, now)
|
|
if err != nil {
|
|
t.Fatalf("DeleteExpired: %v", err)
|
|
}
|
|
if deleted != 1 {
|
|
t.Errorf("deleted = %d, want 1", deleted)
|
|
}
|
|
if _, err := st.GetAndDelete(ctx, "new"); err != nil {
|
|
t.Errorf("'new' should still be present: %v", err)
|
|
}
|
|
}
|
|
|
|
func TestInMemoryOAuthStore_DuplicateState(t *testing.T) {
|
|
st := NewInMemoryOAuthStore()
|
|
ctx := context.Background()
|
|
s := OAuthSession{State: "dup", InstallID: "1", RedirectURI: "/", ExpiresAt: time.Now().Add(time.Hour)}
|
|
if err := st.Insert(ctx, s); err != nil {
|
|
t.Fatalf("first Insert: %v", err)
|
|
}
|
|
if err := st.Insert(ctx, s); err == nil {
|
|
t.Error("second Insert should reject duplicate state")
|
|
}
|
|
}
|
|
|
|
func TestPGOAuthStore_EncryptsCompletionTokens(t *testing.T) {
|
|
st := NewPGOAuthStore(nil, []byte("test-secret"))
|
|
completion := OAuthCompletion{
|
|
Code: "completion-code",
|
|
AccessToken: "access-token-value",
|
|
RefreshToken: "refresh-token-value",
|
|
ExpiresIn: 900,
|
|
NextURL: "/me",
|
|
ExpiresAt: time.Now().Add(time.Minute),
|
|
}
|
|
codeHash := oauthCompletionCodeHash(completion.Code)
|
|
|
|
ciphertext, err := st.encryptCompletionTokens(completion, codeHash)
|
|
if err != nil {
|
|
t.Fatalf("encryptCompletionTokens: %v", err)
|
|
}
|
|
if strings.Contains(ciphertext, completion.AccessToken) || strings.Contains(ciphertext, completion.RefreshToken) {
|
|
t.Fatalf("ciphertext includes plaintext tokens: %q", ciphertext)
|
|
}
|
|
|
|
var out OAuthCompletion
|
|
if err := st.decryptCompletionTokens(ciphertext, codeHash, &out); err != nil {
|
|
t.Fatalf("decryptCompletionTokens: %v", err)
|
|
}
|
|
if out.AccessToken != completion.AccessToken || out.RefreshToken != completion.RefreshToken {
|
|
t.Fatalf("tokens = %q/%q", out.AccessToken, out.RefreshToken)
|
|
}
|
|
}
|