package authn
import (
"context"
"crypto/sha512"
"errors"
"testing"
"time"
"git.sr.ht/~sircmpwn/core-go/auth"
"go.bigb.es/sourcehut-dolt/core"
)
func TestResolveBasic_ValidToken(t *testing.T) {
withStubBackend(t, &stubBackend{users: map[string]auth.AuthContext{
"bigbes": sampleUser(1, "bigbes", auth.USER_TYPE_USER),
}})
pat := forgePAT("bigbes", "", time.Now().Add(time.Hour))
ac, err := ResolveBasic(testCtx(), "bigbes", pat)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if ac.Username != "bigbes" || ac.UserID != 1 {
t.Fatalf("wrong user resolved: %+v", ac)
}
if ac.AuthMethod != auth.AUTH_OAUTH2 {
t.Fatalf("AuthMethod = %q, want %q", ac.AuthMethod, auth.AUTH_OAUTH2)
}
if ac.TokenHash != sha512.Sum512([]byte(pat)) {
t.Fatal("TokenHash must be sha512(password)")
}
}
func TestResolveBasic_UsernameTildeAndCaseInsensitive(t *testing.T) {
withStubBackend(t, &stubBackend{users: map[string]auth.AuthContext{
"bigbes": sampleUser(1, "bigbes", auth.USER_TYPE_USER),
}})
pat := forgePAT("bigbes", "", time.Now().Add(time.Hour))
// Presented username differs by a leading "~" and case; must still match.
if _, err := ResolveBasic(testCtx(), "~BigBes", pat); err != nil {
t.Fatalf("expected match for ~BigBes, got %v", err)
}
}
func TestResolveBasic_UsernameMismatch(t *testing.T) {
withStubBackend(t, &stubBackend{users: map[string]auth.AuthContext{
"bigbes": sampleUser(1, "bigbes", auth.USER_TYPE_USER),
"mallory": sampleUser(2, "mallory", auth.USER_TYPE_USER),
}})
pat := forgePAT("bigbes", "", time.Now().Add(time.Hour))
_, err := ResolveBasic(testCtx(), "mallory", pat)
if !errors.Is(err, ErrInvalidToken) {
t.Fatalf("username mismatch must wrap ErrInvalidToken, got %v", err)
}
}
func TestResolveBasic_ExpiredToken(t *testing.T) {
withStubBackend(t, &stubBackend{users: map[string]auth.AuthContext{
"bigbes": sampleUser(1, "bigbes", auth.USER_TYPE_USER),
}})
pat := forgePAT("bigbes", "", time.Now().Add(-time.Minute))
_, err := ResolveBasic(testCtx(), "bigbes", pat)
if !errors.Is(err, ErrInvalidToken) {
t.Fatalf("expired token must wrap ErrInvalidToken, got %v", err)
}
}
func TestResolveBasic_GarbageToken(t *testing.T) {
withStubBackend(t, &stubBackend{users: map[string]auth.AuthContext{}})
_, err := ResolveBasic(testCtx(), "bigbes", "this-is-not-a-token")
if !errors.Is(err, ErrInvalidToken) {
t.Fatalf("garbage token must wrap ErrInvalidToken, got %v", err)
}
}
func TestResolveBasic_Revoked(t *testing.T) {
pat := forgePAT("bigbes", "", time.Now().Add(time.Hour))
hash := sha512.Sum512([]byte(pat))
withStubBackend(t, &stubBackend{
users: map[string]auth.AuthContext{"bigbes": sampleUser(1, "bigbes", auth.USER_TYPE_USER)},
revoked: map[[64]byte]bool{hash: true},
})
_, err := ResolveBasic(testCtx(), "bigbes", pat)
if !errors.Is(err, ErrInvalidToken) {
t.Fatalf("revoked token must wrap ErrInvalidToken, got %v", err)
}
}
func TestResolveBasic_BackendDownIsTransient(t *testing.T) {
withStubBackend(t, &stubBackend{lookupErr: errBackendDown})
pat := forgePAT("bigbes", "", time.Now().Add(time.Hour))
_, err := ResolveBasic(testCtx(), "bigbes", pat)
if err == nil {
t.Fatal("expected an error")
}
if errors.Is(err, ErrInvalidToken) {
t.Fatalf("backend failure must NOT be a permanent rejection, got %v", err)
}
}
func TestResolveBasic_RevocationBackendDownIsTransient(t *testing.T) {
withStubBackend(t, &stubBackend{
users: map[string]auth.AuthContext{"bigbes": sampleUser(1, "bigbes", auth.USER_TYPE_USER)},
revokeErr: errBackendDown,
})
pat := forgePAT("bigbes", "", time.Now().Add(time.Hour))
_, err := ResolveBasic(testCtx(), "bigbes", pat)
if err == nil {
t.Fatal("expected an error")
}
if errors.Is(err, ErrInvalidToken) {
t.Fatalf("revocation backend failure must be transient, got %v", err)
}
}
// countingBackend records how many times LookupUser is called, to observe caching.
type countingBackend struct {
*stubBackend
lookups int
}
func (c *countingBackend) LookupUser(ctx context.Context, username string, out *auth.AuthContext) error {
c.lookups++
return c.stubBackend.LookupUser(ctx, username, out)
}
func TestResolveBasic_PositiveCache(t *testing.T) {
cb := &countingBackend{stubBackend: &stubBackend{
users: map[string]auth.AuthContext{"bigbes": sampleUser(1, "bigbes", auth.USER_TYPE_USER)},
}}
withStubBackend(t, cb)
pat := forgePAT("bigbes", "", time.Now().Add(time.Hour))
for i := 0; i < 3; i++ {
if _, err := ResolveBasic(testCtx(), "bigbes", pat); err != nil {
t.Fatalf("call %d: %v", i, err)
}
}
if cb.lookups != 1 {
t.Fatalf("expected 1 backend lookup (rest cached), got %d", cb.lookups)
}
}
func TestResolveBasic_CacheExpires(t *testing.T) {
cb := &countingBackend{stubBackend: &stubBackend{
users: map[string]auth.AuthContext{"bigbes": sampleUser(1, "bigbes", auth.USER_TYPE_USER)},
}}
withStubBackend(t, cb)
base := time.Now()
nowFn = func() time.Time { return base }
pat := forgePAT("bigbes", "", base.Add(time.Hour))
if _, err := ResolveBasic(testCtx(), "bigbes", pat); err != nil {
t.Fatal(err)
}
// Advance past the TTL; the entry must be re-resolved.
nowFn = func() time.Time { return base.Add(tokenCacheTTL + time.Second) }
if _, err := ResolveBasic(testCtx(), "bigbes", pat); err != nil {
t.Fatal(err)
}
if cb.lookups != 2 {
t.Fatalf("expected 2 lookups across the TTL boundary, got %d", cb.lookups)
}
}
func TestResolveBasic_NegativeNotCached(t *testing.T) {
cb := &countingBackend{stubBackend: &stubBackend{
users: map[string]auth.AuthContext{"bigbes": sampleUser(1, "bigbes", auth.USER_TYPE_USER)},
}}
withStubBackend(t, cb)
// A token for an unknown user: LookupUser errors each time (uncached).
pat := forgePAT("ghost", "", time.Now().Add(time.Hour))
for i := 0; i < 2; i++ {
if _, err := ResolveBasic(testCtx(), "ghost", pat); err == nil {
t.Fatalf("call %d: expected error for unknown user", i)
}
}
if cb.lookups != 2 {
t.Fatalf("negative results must not be cached, got %d lookups", cb.lookups)
}
}
func TestTokenGrantsAllow(t *testing.T) {
ctx := testCtx()
mustGrants := func(s string) auth.Grants {
g, err := auth.DecodeGrants(ctx, s)
if err != nil {
t.Fatalf("DecodeGrants(%q): %v", s, err)
}
return g
}
patAC := func(grants string) *auth.AuthContext {
return &auth.AuthContext{
AuthMethod: auth.AUTH_OAUTH2,
BearerToken: &auth.BearerToken{},
Grants: mustGrants(grants),
}
}
cases := []struct {
name string
ac *auth.AuthContext
mode core.AccessMode
want bool
}{
{"nil caller passes", nil, core.AccessRW, true},
{"cookie (no bearer) passes", &auth.AuthContext{AuthMethod: auth.AUTH_COOKIE}, core.AccessRW, true},
{"empty grants RO", patAC(""), core.AccessRO, true},
{"empty grants RW", patAC(""), core.AccessRW, true},
{"repos:RO allows read", patAC("dolt.sr.ht/repos:RO"), core.AccessRO, true},
{"repos:RO denies write", patAC("dolt.sr.ht/repos:RO"), core.AccessRW, false},
{"repos:RW allows read", patAC("dolt.sr.ht/repos:RW"), core.AccessRO, true},
{"repos:RW allows write", patAC("dolt.sr.ht/repos:RW"), core.AccessRW, true},
{"unrelated scope denies read", patAC("git.sr.ht/repos:RW"), core.AccessRO, false},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
if got := TokenGrantsAllow(tc.ac, tc.mode); got != tc.want {
t.Fatalf("TokenGrantsAllow = %v, want %v", got, tc.want)
}
})
}
}