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