fromozuzhouzzz
Deploy grok2api v3.0.11 to HF Spaces
f1dd159
Raw
History Blame Contribute Delete
6.98 kB
package adminauth
import (
"context"
"errors"
"path/filepath"
"sync"
"testing"
"time"
"github.com/chenyme/grok2api/backend/internal/domain/admin"
"github.com/chenyme/grok2api/backend/internal/infra/persistence/relational"
"github.com/chenyme/grok2api/backend/internal/infra/security"
"github.com/chenyme/grok2api/backend/internal/repository"
)
func TestRefreshTokenRotationAndLogout(t *testing.T) {
database, err := relational.OpenSQLite(context.Background(), filepath.Join(t.TempDir(), "auth.db"))
if err != nil {
t.Fatal(err)
}
defer database.Close()
if err := database.InitializeSchema(context.Background()); err != nil {
t.Fatal(err)
}
service := NewService(relational.NewAdminRepository(database), relational.NewAdminSessionRepository(database), security.NewTokenService("12345678901234567890123456789012"), time.Minute, time.Hour)
ctx := context.Background()
if err := service.Bootstrap(ctx, "admin", "password123"); err != nil {
t.Fatal(err)
}
_, tokens, err := service.Login(ctx, "admin", "password123", "127.0.0.1")
if err != nil {
t.Fatal(err)
}
rotated, err := service.Refresh(ctx, tokens.RefreshToken)
if err != nil {
t.Fatal(err)
}
if _, err := service.Refresh(ctx, tokens.RefreshToken); !errors.Is(err, ErrInvalidSession) {
t.Fatalf("旧 refresh token 仍可使用: %v", err)
}
if err := service.Logout(ctx, rotated.RefreshToken); err != nil {
t.Fatal(err)
}
if _, err := service.AuthenticateAccess(ctx, rotated.AccessToken); !errors.Is(err, ErrInvalidSession) {
t.Fatalf("注销后的 access token 仍可使用: %v", err)
}
if _, err := service.Refresh(ctx, rotated.RefreshToken); !errors.Is(err, ErrInvalidSession) {
t.Fatalf("注销后的 refresh token 仍可使用: %v", err)
}
}
func TestChangePasswordRevokesAllSessions(t *testing.T) {
database, err := relational.OpenSQLite(context.Background(), filepath.Join(t.TempDir(), "auth.db"))
if err != nil {
t.Fatal(err)
}
defer database.Close()
if err := database.InitializeSchema(context.Background()); err != nil {
t.Fatal(err)
}
service := NewService(relational.NewAdminRepository(database), relational.NewAdminSessionRepository(database), security.NewTokenService("12345678901234567890123456789012"), time.Minute, time.Hour)
ctx := context.Background()
if err := service.Bootstrap(ctx, "admin", "password123"); err != nil {
t.Fatal(err)
}
adminValue, tokens, err := service.Login(ctx, "admin", "password123", "127.0.0.1")
if err != nil {
t.Fatal(err)
}
if err := service.ChangePassword(ctx, adminValue.ID, "password123", "password456"); err != nil {
t.Fatal(err)
}
if _, err := service.AuthenticateAccess(ctx, tokens.AccessToken); !errors.Is(err, ErrInvalidSession) {
t.Fatalf("修改密码后的 access token 仍可使用: %v", err)
}
if _, err := service.Refresh(ctx, tokens.RefreshToken); !errors.Is(err, ErrInvalidSession) {
t.Fatalf("修改密码后的 refresh token 仍可使用: %v", err)
}
if _, _, err := service.Login(ctx, "admin", "password123", "127.0.0.1"); !errors.Is(err, ErrInvalidCredentials) {
t.Fatalf("旧密码仍可登录: %v", err)
}
if _, _, err := service.Login(ctx, "admin", "password456", "127.0.0.1"); err != nil {
t.Fatalf("新密码无法登录: %v", err)
}
}
func TestLoginRateLimiterFailureIsEnforced(t *testing.T) {
service := NewService(nil, nil, security.NewTokenService("12345678901234567890123456789012"), time.Minute, time.Hour)
service.SetLoginRateLimiter(rejectingRateLimiter{})
if _, _, err := service.Login(context.Background(), "admin", "password123", "127.0.0.1"); !errors.Is(err, ErrLoginRateLimited) {
t.Fatalf("login rate limit error = %v", err)
}
}
func TestLoginDistinguishesPersistenceFailure(t *testing.T) {
service := NewService(failingAdminRepository{}, nil, security.NewTokenService("12345678901234567890123456789012"), time.Minute, time.Hour)
if _, _, err := service.Login(context.Background(), "admin", "password123", "127.0.0.1"); !errors.Is(err, ErrRuntimeUnavailable) {
t.Fatalf("login persistence error = %v", err)
}
}
func TestConcurrentRefreshAllowsExactlyOneRotation(t *testing.T) {
database, err := relational.OpenSQLite(context.Background(), filepath.Join(t.TempDir(), "auth.db"))
if err != nil {
t.Fatal(err)
}
defer database.Close()
if err := database.InitializeSchema(context.Background()); err != nil {
t.Fatal(err)
}
baseSessions := relational.NewAdminSessionRepository(database)
sessions := newCoordinatedSessionRepository(baseSessions, 2)
service := NewService(
relational.NewAdminRepository(database),
sessions,
security.NewTokenService("12345678901234567890123456789012"),
time.Minute,
time.Hour,
)
ctx := context.Background()
if err := service.Bootstrap(ctx, "admin", "password123"); err != nil {
t.Fatal(err)
}
_, tokens, err := service.Login(ctx, "admin", "password123", "127.0.0.1")
if err != nil {
t.Fatal(err)
}
type refreshResult struct {
tokens Tokens
err error
}
results := make(chan refreshResult, 2)
start := make(chan struct{})
for range 2 {
go func() {
<-start
rotated, refreshErr := service.Refresh(ctx, tokens.RefreshToken)
results <- refreshResult{tokens: rotated, err: refreshErr}
}()
}
close(start)
var successful Tokens
successCount := 0
invalidCount := 0
for range 2 {
result := <-results
switch {
case result.err == nil:
successCount++
successful = result.tokens
case errors.Is(result.err, ErrInvalidSession):
invalidCount++
default:
t.Fatalf("unexpected refresh error: %v", result.err)
}
}
if successCount != 1 || invalidCount != 1 {
t.Fatalf("successes = %d, invalid sessions = %d", successCount, invalidCount)
}
if _, err := service.Refresh(ctx, successful.RefreshToken); err != nil {
t.Fatalf("winning refresh token is unusable: %v", err)
}
}
type coordinatedSessionRepository struct {
repository.AdminSessionRepository
mu sync.Mutex
remaining int
ready chan struct{}
}
type rejectingRateLimiter struct{}
func (rejectingRateLimiter) Allow(context.Context, string, int, time.Time) (bool, error) {
return false, nil
}
type failingAdminRepository struct{ repository.AdminRepository }
func (failingAdminRepository) GetByUsername(context.Context, string) (admin.Admin, error) {
return admin.Admin{}, errors.New("database unavailable")
}
func newCoordinatedSessionRepository(base repository.AdminSessionRepository, reads int) *coordinatedSessionRepository {
return &coordinatedSessionRepository{AdminSessionRepository: base, remaining: reads, ready: make(chan struct{})}
}
func (r *coordinatedSessionRepository) GetByTokenHash(ctx context.Context, tokenHash string) (admin.Session, error) {
session, err := r.AdminSessionRepository.GetByTokenHash(ctx, tokenHash)
if err != nil {
return admin.Session{}, err
}
r.mu.Lock()
if r.remaining > 0 {
r.remaining--
if r.remaining == 0 {
close(r.ready)
}
}
ready := r.ready
r.mu.Unlock()
<-ready
return session, nil
}