feat: 实现 OIDC 服务端会话与请求刷新

使用 AES-256-GCM 保存认证中心令牌,并以随机 Cookie 哈希关联 PostgreSQL 会话。加入闲置与绝对期限、Refresh Token 轮换以及多实例刷新锁。
This commit is contained in:
2026-07-13 19:09:10 +08:00
parent a81a7b5200
commit d345c070ae
16 changed files with 1804 additions and 58 deletions
+82
View File
@@ -0,0 +1,82 @@
package oidcsession
import (
"crypto/aes"
"crypto/cipher"
"crypto/rand"
"encoding/json"
"errors"
"fmt"
"io"
)
type TokenBundle struct {
AccessToken string `json:"accessToken"`
RefreshToken string `json:"refreshToken"`
IDToken string `json:"idToken,omitempty"`
}
type Cipher struct {
aead cipher.AEAD
}
func NewCipher(key []byte) (*Cipher, error) {
if len(key) != 32 {
return nil, errors.New("OIDC session encryption key must be exactly 32 bytes")
}
block, err := aes.NewCipher(key)
if err != nil {
return nil, err
}
aead, err := cipher.NewGCM(block)
if err != nil {
return nil, err
}
return &Cipher{aead: aead}, nil
}
func (c *Cipher) EncryptBundle(bundle TokenBundle, sessionID, gatewayUserID string) ([]byte, error) {
return c.SealJSON(bundle, tokenAAD(sessionID, gatewayUserID))
}
func (c *Cipher) SealJSON(value any, aad []byte) ([]byte, error) {
plaintext, err := json.Marshal(value)
if err != nil {
return nil, err
}
nonce := make([]byte, c.aead.NonceSize())
if _, err := io.ReadFull(rand.Reader, nonce); err != nil {
return nil, err
}
return c.aead.Seal(nonce, nonce, plaintext, aad), nil
}
func (c *Cipher) DecryptBundle(encrypted []byte, sessionID, gatewayUserID string) (TokenBundle, error) {
var bundle TokenBundle
if err := c.OpenJSON(encrypted, tokenAAD(sessionID, gatewayUserID), &bundle); err != nil {
return TokenBundle{}, err
}
if bundle.AccessToken == "" || bundle.RefreshToken == "" {
return TokenBundle{}, errors.New("OIDC session token bundle is invalid")
}
return bundle, nil
}
func (c *Cipher) OpenJSON(encrypted, aad []byte, output any) error {
if len(encrypted) <= c.aead.NonceSize() {
return errors.New("OIDC session ciphertext is invalid")
}
nonce, ciphertext := encrypted[:c.aead.NonceSize()], encrypted[c.aead.NonceSize():]
plaintext, err := c.aead.Open(nil, nonce, ciphertext, aad)
if err != nil {
return errors.New("OIDC session ciphertext authentication failed")
}
if err := json.Unmarshal(plaintext, output); err != nil {
return errors.New("OIDC session ciphertext payload is invalid")
}
return nil
}
func tokenAAD(sessionID, gatewayUserID string) []byte {
return []byte(fmt.Sprintf("easyai-gateway/oidc-session/v1\x00%s\x00%s", sessionID, gatewayUserID))
}
@@ -0,0 +1,39 @@
package oidcsession
import (
"bytes"
"strings"
"testing"
)
func TestCipherEncryptsTokenBundleAndAuthenticatesContext(t *testing.T) {
key := bytes.Repeat([]byte{0x42}, 32)
cipher, err := NewCipher(key)
if err != nil {
t.Fatal(err)
}
bundle := TokenBundle{AccessToken: "plain-access", RefreshToken: "plain-refresh", IDToken: "plain-id"}
encrypted, err := cipher.EncryptBundle(bundle, "session-id", "gateway-user-id")
if err != nil {
t.Fatal(err)
}
if strings.Contains(string(encrypted), "plain-") {
t.Fatal("token bundle was stored in plaintext")
}
decoded, err := cipher.DecryptBundle(encrypted, "session-id", "gateway-user-id")
if err != nil {
t.Fatal(err)
}
if decoded != bundle {
t.Fatalf("decoded bundle = %#v", decoded)
}
if _, err := cipher.DecryptBundle(encrypted, "another-session", "gateway-user-id"); err == nil {
t.Fatal("ciphertext was accepted with different authenticated context")
}
}
func TestCipherRequiresIndependentAES256Key(t *testing.T) {
if _, err := NewCipher(make([]byte, 31)); err == nil {
t.Fatal("31-byte key was accepted")
}
}
@@ -0,0 +1,85 @@
package oidcsession
import (
"crypto/rand"
"crypto/sha256"
"encoding/base64"
"errors"
"net/url"
"strings"
"time"
)
const LoginTransactionCookieName = "easyai_gateway_oidc_login"
var loginTransactionAAD = []byte("easyai-gateway/oidc-login-transaction/v1")
type LoginTransaction struct {
State string `json:"state"`
Nonce string `json:"nonce"`
PKCEVerifier string `json:"pkceVerifier"`
ReturnTo string `json:"returnTo"`
CreatedAt time.Time `json:"createdAt"`
}
func NewLoginTransaction(returnTo string, now time.Time) (LoginTransaction, string, error) {
if !ValidReturnTo(returnTo) {
return LoginTransaction{}, "", errors.New("returnTo must be a same-origin relative path")
}
state, err := randomBase64URL(32)
if err != nil {
return LoginTransaction{}, "", err
}
nonce, err := randomBase64URL(32)
if err != nil {
return LoginTransaction{}, "", err
}
verifier, err := randomBase64URL(32)
if err != nil {
return LoginTransaction{}, "", err
}
challengeHash := sha256.Sum256([]byte(verifier))
return LoginTransaction{State: state, Nonce: nonce, PKCEVerifier: verifier, ReturnTo: returnTo, CreatedAt: now.UTC()},
base64.RawURLEncoding.EncodeToString(challengeHash[:]), nil
}
func (c *Cipher) EncodeLoginTransaction(transaction LoginTransaction) (string, error) {
payload, err := c.SealJSON(transaction, loginTransactionAAD)
if err != nil {
return "", err
}
return base64.RawURLEncoding.EncodeToString(payload), nil
}
func (c *Cipher) DecodeLoginTransaction(encoded string, now time.Time) (LoginTransaction, error) {
payload, err := base64.RawURLEncoding.DecodeString(strings.TrimSpace(encoded))
if err != nil {
return LoginTransaction{}, errors.New("OIDC login transaction is invalid")
}
var transaction LoginTransaction
if err := c.OpenJSON(payload, loginTransactionAAD, &transaction); err != nil {
return LoginTransaction{}, err
}
if transaction.State == "" || transaction.Nonce == "" || transaction.PKCEVerifier == "" || !ValidReturnTo(transaction.ReturnTo) ||
transaction.CreatedAt.IsZero() || now.Before(transaction.CreatedAt.Add(-time.Minute)) || !now.Before(transaction.CreatedAt.Add(10*time.Minute)) {
return LoginTransaction{}, errors.New("OIDC login transaction has expired or is invalid")
}
return transaction, nil
}
func ValidReturnTo(value string) bool {
value = strings.TrimSpace(value)
if value == "" || !strings.HasPrefix(value, "/") || strings.HasPrefix(value, "//") || strings.Contains(value, "\\") {
return false
}
parsed, err := url.Parse(value)
return err == nil && !parsed.IsAbs() && parsed.Host == ""
}
func randomBase64URL(size int) (string, error) {
value := make([]byte, size)
if _, err := rand.Read(value); err != nil {
return "", err
}
return base64.RawURLEncoding.EncodeToString(value), nil
}
@@ -0,0 +1,45 @@
package oidcsession
import (
"bytes"
"strings"
"testing"
"time"
)
func TestLoginTransactionIsEncryptedBoundedAndPKCES256(t *testing.T) {
now := time.Date(2026, 7, 13, 12, 0, 0, 0, time.UTC)
cipher, _ := NewCipher(bytes.Repeat([]byte{9}, 32))
transaction, challenge, err := NewLoginTransaction("/workspace?tab=wallet", now)
if err != nil {
t.Fatal(err)
}
if transaction.State == transaction.Nonce || len(challenge) != 43 || challenge == transaction.PKCEVerifier {
t.Fatal("state, nonce and PKCE values are not independent S256 material")
}
encoded, err := cipher.EncodeLoginTransaction(transaction)
if err != nil {
t.Fatal(err)
}
if strings.Contains(encoded, transaction.State) || strings.Contains(encoded, transaction.PKCEVerifier) {
t.Fatal("login transaction cookie contains plaintext security material")
}
decoded, err := cipher.DecodeLoginTransaction(encoded, now.Add(9*time.Minute))
if err != nil || decoded.ReturnTo != transaction.ReturnTo {
t.Fatalf("decode transaction=%#v err=%v", decoded, err)
}
if _, err := cipher.DecodeLoginTransaction(encoded, now.Add(10*time.Minute)); err == nil {
t.Fatal("10-minute login transaction was accepted")
}
}
func TestValidReturnToRejectsOpenRedirects(t *testing.T) {
for _, value := range []string{"https://evil.example", "//evil.example", "/\\evil", "", "workspace"} {
if ValidReturnTo(value) {
t.Fatalf("unsafe returnTo accepted: %q", value)
}
}
if !ValidReturnTo("/workspace/tasks?from=login#latest") {
t.Fatal("safe relative returnTo was rejected")
}
}
+326
View File
@@ -0,0 +1,326 @@
package oidcsession
import (
"context"
"crypto/rand"
"crypto/sha256"
"encoding/base64"
"encoding/hex"
"errors"
"strings"
"time"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
)
var (
ErrSessionInvalid = errors.New("OIDC session is invalid")
ErrSessionExpired = errors.New("OIDC session has expired")
ErrSessionRefreshUnavailable = errors.New("OIDC session refresh is unavailable")
ErrSessionStoreUnavailable = errors.New("OIDC session store is unavailable")
ErrGatewayUserDisabled = errors.New("gateway user is disabled")
)
type Repository interface {
CreateOIDCSession(context.Context, store.CreateOIDCSessionInput) (store.OIDCSession, error)
FindOIDCSessionByHash(context.Context, []byte) (store.OIDCSession, error)
TouchOIDCSession(context.Context, string, time.Time, time.Time) error
AcquireOIDCSessionRefreshLock(context.Context, string, int64, string, time.Time, time.Time) (bool, error)
CompleteOIDCSessionRefresh(context.Context, string, int64, string, []byte, time.Time) error
DeleteOIDCSessionByHash(context.Context, []byte) error
DeleteOIDCSessionByID(context.Context, string) error
CleanupExpiredOIDCSessions(context.Context, time.Time) (int64, error)
}
type TokenVerifier interface {
Verify(context.Context, string) (*auth.User, error)
}
type PublicClient interface {
Refresh(context.Context, string) (auth.OIDCTokenResponse, error)
RevokeRefreshToken(context.Context, string) error
EndSessionURL(context.Context, string) (string, error)
}
type Config struct {
IdleTTL time.Duration
AbsoluteTTL time.Duration
RefreshBefore time.Duration
RefreshLease time.Duration
RefreshWait time.Duration
}
type Service struct {
repository Repository
cipher *Cipher
verifier TokenVerifier
client PublicClient
config Config
now func() time.Time
}
func NewService(repository Repository, cipher *Cipher, verifier TokenVerifier, client PublicClient, config Config) (*Service, error) {
if repository == nil || cipher == nil || verifier == nil || client == nil {
return nil, errors.New("OIDC session dependencies are required")
}
if config.IdleTTL <= 0 || config.AbsoluteTTL <= config.IdleTTL || config.RefreshBefore <= 0 {
return nil, errors.New("OIDC session TTL configuration is invalid")
}
if config.RefreshLease <= 0 {
config.RefreshLease = 5 * time.Second
}
if config.RefreshWait <= 0 {
config.RefreshWait = 2 * time.Second
}
return &Service{repository: repository, cipher: cipher, verifier: verifier, client: client, config: config, now: time.Now}, nil
}
func (s *Service) Create(ctx context.Context, bundle TokenBundle, localUser *auth.User) (string, error) {
if localUser == nil || localUser.GatewayUserID == "" || localUser.GatewayTenantID == "" || bundle.RefreshToken == "" {
return "", ErrSessionInvalid
}
verified, err := s.verifier.Verify(ctx, bundle.AccessToken)
if err != nil || verified == nil || verified.Source != "oidc" || verified.ID == "" || verified.ID != localUser.ID {
return "", ErrSessionInvalid
}
now := s.now()
if !verified.TokenExpiresAt.After(now) {
return "", ErrSessionExpired
}
raw, hash, err := newSessionToken()
if err != nil {
return "", ErrSessionStoreUnavailable
}
aadSessionID := hex.EncodeToString(hash)
ciphertext, err := s.cipher.EncryptBundle(bundle, aadSessionID, localUser.GatewayUserID)
if err != nil {
return "", ErrSessionStoreUnavailable
}
_, err = s.repository.CreateOIDCSession(ctx, store.CreateOIDCSessionInput{
SessionTokenHash: hash, GatewayUserID: localUser.GatewayUserID, GatewayTenantID: localUser.GatewayTenantID,
TokenCiphertext: ciphertext, AccessTokenExpiresAt: verified.TokenExpiresAt,
LastSeenAt: now, IdleExpiresAt: now.Add(s.config.IdleTTL), AbsoluteExpiresAt: now.Add(s.config.AbsoluteTTL),
})
if err != nil {
return "", ErrSessionStoreUnavailable
}
return raw, nil
}
func (s *Service) Resolve(ctx context.Context, raw string) (*auth.User, error) {
hash, err := sessionTokenHash(raw)
if err != nil {
return nil, ErrSessionInvalid
}
record, err := s.repository.FindOIDCSessionByHash(ctx, hash)
if errors.Is(err, store.ErrOIDCSessionNotFound) {
return nil, ErrSessionInvalid
}
if err != nil {
return nil, ErrSessionStoreUnavailable
}
return s.resolveRecord(ctx, hash, record)
}
func (s *Service) resolveRecord(ctx context.Context, hash []byte, record store.OIDCSession) (*auth.User, error) {
now := s.now()
if record.UserDeleted || record.UserStatus != "active" {
return nil, ErrGatewayUserDisabled
}
if !record.IdleExpiresAt.After(now) || !record.AbsoluteExpiresAt.After(now) {
_ = s.repository.DeleteOIDCSessionByID(ctx, record.ID)
return nil, ErrSessionExpired
}
bundle, err := s.cipher.DecryptBundle(record.TokenCiphertext, hex.EncodeToString(hash), record.GatewayUserID)
if err != nil {
return nil, ErrSessionStoreUnavailable
}
if !record.AccessTokenExpiresAt.After(now.Add(s.config.RefreshBefore)) {
return s.refresh(ctx, hash, record, bundle)
}
user, err := s.verifySessionUser(ctx, record, bundle.AccessToken)
if err != nil {
return nil, err
}
if err := s.touch(ctx, record, now); err != nil {
return nil, err
}
return user, nil
}
func (s *Service) refresh(ctx context.Context, hash []byte, record store.OIDCSession, bundle TokenBundle) (*auth.User, error) {
now := s.now()
lockID, err := newUUID()
if err != nil {
return nil, ErrSessionStoreUnavailable
}
acquired, err := s.repository.AcquireOIDCSessionRefreshLock(ctx, record.ID, record.RefreshVersion, lockID, now.Add(s.config.RefreshLease), now)
if err != nil {
return nil, ErrSessionStoreUnavailable
}
if !acquired {
return s.waitForRefresh(ctx, hash, record, bundle)
}
refreshed, err := s.client.Refresh(ctx, bundle.RefreshToken)
if errors.Is(err, auth.ErrOIDCInvalidGrant) {
_ = s.repository.DeleteOIDCSessionByID(ctx, record.ID)
return nil, ErrSessionExpired
}
if err != nil {
// Keep the lease until it expires: an indeterminate refresh response must not
// cause the same rotating refresh token to be replayed immediately.
if record.AccessTokenExpiresAt.After(now) {
user, verifyErr := s.verifySessionUser(ctx, record, bundle.AccessToken)
if verifyErr == nil {
if touchErr := s.touch(ctx, record, now); touchErr != nil {
return nil, touchErr
}
return user, nil
}
}
return nil, ErrSessionRefreshUnavailable
}
user, err := s.verifySessionUser(ctx, record, refreshed.AccessToken)
if err != nil {
_ = s.repository.DeleteOIDCSessionByID(ctx, record.ID)
return nil, ErrSessionExpired
}
if refreshed.RefreshToken == "" {
// OAuth token responses may omit refresh_token when the issuer keeps the
// existing token valid. Persist the old value unless a rotated one exists.
refreshed.RefreshToken = bundle.RefreshToken
}
if refreshed.IDToken == "" {
refreshed.IDToken = bundle.IDToken
}
newBundle := TokenBundle{AccessToken: refreshed.AccessToken, RefreshToken: refreshed.RefreshToken, IDToken: refreshed.IDToken}
ciphertext, err := s.cipher.EncryptBundle(newBundle, hex.EncodeToString(hash), record.GatewayUserID)
if err != nil {
// The remote refresh succeeded, so the previous rotating refresh token may
// already be invalid. Destroy the stale local session instead of replaying it.
_ = s.repository.DeleteOIDCSessionByID(ctx, record.ID)
return nil, ErrSessionStoreUnavailable
}
if err := s.repository.CompleteOIDCSessionRefresh(ctx, record.ID, record.RefreshVersion, lockID, ciphertext, user.TokenExpiresAt); err != nil {
return nil, ErrSessionStoreUnavailable
}
record.AccessTokenExpiresAt = user.TokenExpiresAt
if err := s.touch(ctx, record, now); err != nil {
return nil, err
}
return user, nil
}
func (s *Service) waitForRefresh(ctx context.Context, hash []byte, original store.OIDCSession, oldBundle TokenBundle) (*auth.User, error) {
deadline := time.Now().Add(s.config.RefreshWait)
for time.Now().Before(deadline) {
select {
case <-ctx.Done():
return nil, ErrSessionRefreshUnavailable
case <-time.After(25 * time.Millisecond):
}
record, err := s.repository.FindOIDCSessionByHash(ctx, hash)
if errors.Is(err, store.ErrOIDCSessionNotFound) {
return nil, ErrSessionExpired
}
if err != nil {
return nil, ErrSessionStoreUnavailable
}
if record.RefreshVersion > original.RefreshVersion {
return s.resolveRecord(ctx, hash, record)
}
}
now := s.now()
if original.AccessTokenExpiresAt.After(now) {
user, err := s.verifySessionUser(ctx, original, oldBundle.AccessToken)
if err == nil {
if touchErr := s.touch(ctx, original, now); touchErr != nil {
return nil, touchErr
}
return user, nil
}
}
return nil, ErrSessionRefreshUnavailable
}
func (s *Service) verifySessionUser(ctx context.Context, record store.OIDCSession, accessToken string) (*auth.User, error) {
user, err := s.verifier.Verify(ctx, accessToken)
if err != nil || user == nil || user.Source != "oidc" || user.ID != record.ExternalUserID {
return nil, ErrSessionInvalid
}
return user, nil
}
func (s *Service) touch(ctx context.Context, record store.OIDCSession, now time.Time) error {
if err := s.repository.TouchOIDCSession(ctx, record.ID, now, now.Add(s.config.IdleTTL)); err != nil {
if errors.Is(err, store.ErrOIDCSessionNotFound) {
return ErrSessionExpired
}
return ErrSessionStoreUnavailable
}
return nil
}
func (s *Service) Delete(ctx context.Context, raw string) (TokenBundle, error) {
hash, err := sessionTokenHash(raw)
if err != nil {
return TokenBundle{}, nil
}
record, err := s.repository.FindOIDCSessionByHash(ctx, hash)
if errors.Is(err, store.ErrOIDCSessionNotFound) {
return TokenBundle{}, nil
}
if err != nil {
return TokenBundle{}, ErrSessionStoreUnavailable
}
bundle, decryptErr := s.cipher.DecryptBundle(record.TokenCiphertext, hex.EncodeToString(hash), record.GatewayUserID)
if err := s.repository.DeleteOIDCSessionByHash(ctx, hash); err != nil {
return TokenBundle{}, ErrSessionStoreUnavailable
}
if decryptErr != nil {
// The session row is already gone. Treat logout as successful even though
// the unusable refresh token could not be revoked at the issuer.
return TokenBundle{}, nil
}
return bundle, nil
}
func (s *Service) Cleanup(ctx context.Context) (int64, error) {
count, err := s.repository.CleanupExpiredOIDCSessions(ctx, s.now())
if err != nil {
return 0, ErrSessionStoreUnavailable
}
return count, nil
}
func newSessionToken() (string, []byte, error) {
raw := make([]byte, 32)
if _, err := rand.Read(raw); err != nil {
return "", nil, err
}
encoded := base64.RawURLEncoding.EncodeToString(raw)
hash := sha256.Sum256([]byte(encoded))
return encoded, hash[:], nil
}
func sessionTokenHash(raw string) ([]byte, error) {
raw = strings.TrimSpace(raw)
decoded, err := base64.RawURLEncoding.DecodeString(raw)
if err != nil || len(decoded) != 32 {
return nil, ErrSessionInvalid
}
hash := sha256.Sum256([]byte(raw))
return hash[:], nil
}
func newUUID() (string, error) {
value := make([]byte, 16)
if _, err := rand.Read(value); err != nil {
return "", err
}
value[6] = value[6]&0x0f | 0x40
value[8] = value[8]&0x3f | 0x80
return hex.EncodeToString(value[0:4]) + "-" + hex.EncodeToString(value[4:6]) + "-" + hex.EncodeToString(value[6:8]) + "-" + hex.EncodeToString(value[8:10]) + "-" + hex.EncodeToString(value[10:16]), nil
}
@@ -0,0 +1,419 @@
package oidcsession
import (
"bytes"
"context"
"errors"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
)
func TestServiceStoresOnlyHashedSessionAndEncryptedTokens(t *testing.T) {
now := time.Date(2026, 7, 13, 12, 0, 0, 0, time.UTC)
repository := newFakeRepository("subject-1")
verifier := fakeVerifier{users: map[string]*auth.User{
"access-token": {ID: "subject-1", Source: "oidc", Roles: []string{"user"}, TokenExpiresAt: now.Add(5 * time.Minute)},
}}
service := newTestService(t, repository, verifier, &fakePublicClient{})
service.now = func() time.Time { return now }
localUser := &auth.User{ID: "subject-1", GatewayUserID: "11111111-1111-4111-8111-111111111111", GatewayTenantID: "22222222-2222-4222-8222-222222222222"}
raw, err := service.Create(context.Background(), TokenBundle{AccessToken: "access-token", RefreshToken: "refresh-token", IDToken: "id-token"}, localUser)
if err != nil {
t.Fatal(err)
}
if len(raw) != 43 {
t.Fatalf("opaque session length = %d, want 43", len(raw))
}
record := repository.snapshot()
if bytes.Contains(record.TokenCiphertext, []byte("access-token")) || bytes.Contains(record.TokenCiphertext, []byte("refresh-token")) {
t.Fatal("repository received plaintext token material")
}
hash, _ := sessionTokenHash(raw)
if !bytes.Equal(hash, record.SessionTokenHash) || bytes.Equal([]byte(raw), record.SessionTokenHash) {
t.Fatal("repository did not receive only the SHA-256 session hash")
}
user, err := service.Resolve(context.Background(), raw)
if err != nil || user.ID != "subject-1" {
t.Fatalf("resolve user=%#v err=%v", user, err)
}
}
func TestServiceConcurrentExpiredRequestsRefreshExactlyOnce(t *testing.T) {
now := time.Date(2026, 7, 13, 12, 0, 0, 0, time.UTC)
repository := newFakeRepository("subject-1")
verifier := fakeVerifier{users: map[string]*auth.User{
"old-access": {ID: "subject-1", Source: "oidc", Roles: []string{"user"}, TokenExpiresAt: now.Add(5 * time.Minute)},
"new-access": {ID: "subject-1", Source: "oidc", Roles: []string{"user"}, TokenExpiresAt: now.Add(5 * time.Minute)},
}}
client := &fakePublicClient{refreshResponse: auth.OIDCTokenResponse{AccessToken: "new-access", RefreshToken: "rotated-refresh", ExpiresIn: 300}, delay: 40 * time.Millisecond}
service := newTestService(t, repository, verifier, client)
service.now = func() time.Time { return now }
raw, err := service.Create(context.Background(), TokenBundle{AccessToken: "old-access", RefreshToken: "old-refresh"}, &auth.User{
ID: "subject-1", GatewayUserID: "11111111-1111-4111-8111-111111111111", GatewayTenantID: "22222222-2222-4222-8222-222222222222",
})
if err != nil {
t.Fatal(err)
}
repository.expireAccessToken(now.Add(-time.Second))
var wait sync.WaitGroup
errorsFound := make(chan error, 20)
for range 20 {
wait.Add(1)
go func() {
defer wait.Done()
user, resolveErr := service.Resolve(context.Background(), raw)
if resolveErr != nil {
errorsFound <- resolveErr
return
}
if user.ID != "subject-1" {
errorsFound <- errors.New("wrong resolved subject")
}
}()
}
wait.Wait()
close(errorsFound)
for err := range errorsFound {
t.Errorf("concurrent resolve: %v", err)
}
if got := client.refreshCalls.Load(); got != 1 {
t.Fatalf("refresh calls = %d, want exactly 1", got)
}
if repository.snapshot().RefreshVersion != 2 {
t.Fatalf("refresh version = %d, want 2", repository.snapshot().RefreshVersion)
}
}
func TestServiceDoesNotRefreshExpiredIdleSession(t *testing.T) {
now := time.Date(2026, 7, 13, 12, 0, 0, 0, time.UTC)
repository := newFakeRepository("subject-1")
verifier := fakeVerifier{users: map[string]*auth.User{"access": {ID: "subject-1", Source: "oidc", TokenExpiresAt: now.Add(5 * time.Minute)}}}
client := &fakePublicClient{refreshResponse: auth.OIDCTokenResponse{AccessToken: "new", RefreshToken: "rotated"}}
service := newTestService(t, repository, verifier, client)
service.now = func() time.Time { return now }
raw, err := service.Create(context.Background(), TokenBundle{AccessToken: "access", RefreshToken: "refresh"}, &auth.User{
ID: "subject-1", GatewayUserID: "11111111-1111-4111-8111-111111111111", GatewayTenantID: "22222222-2222-4222-8222-222222222222",
})
if err != nil {
t.Fatal(err)
}
repository.expireIdle(now.Add(-time.Second))
if _, err := service.Resolve(context.Background(), raw); !errors.Is(err, ErrSessionExpired) {
t.Fatalf("Resolve() error = %v, want ErrSessionExpired", err)
}
if client.refreshCalls.Load() != 0 {
t.Fatal("expired idle session attempted a refresh")
}
}
func TestServiceKeepsExistingRefreshTokenWhenIssuerDoesNotRotate(t *testing.T) {
now := time.Date(2026, 7, 13, 12, 0, 0, 0, time.UTC)
repository := newFakeRepository("subject-1")
verifier := fakeVerifier{users: map[string]*auth.User{
"old-access": {ID: "subject-1", Source: "oidc", TokenExpiresAt: now.Add(5 * time.Minute)},
"new-access": {ID: "subject-1", Source: "oidc", TokenExpiresAt: now.Add(5 * time.Minute)},
}}
client := &fakePublicClient{refreshResponse: auth.OIDCTokenResponse{AccessToken: "new-access"}}
service := newTestService(t, repository, verifier, client)
service.now = func() time.Time { return now }
raw, err := service.Create(context.Background(), TokenBundle{AccessToken: "old-access", RefreshToken: "existing-refresh"}, testLocalUser())
if err != nil {
t.Fatal(err)
}
repository.expireAccessToken(now.Add(-time.Second))
if _, err := service.Resolve(context.Background(), raw); err != nil {
t.Fatalf("Resolve() error = %v", err)
}
bundle, err := service.Delete(context.Background(), raw)
if err != nil {
t.Fatal(err)
}
if bundle.RefreshToken != "existing-refresh" || bundle.AccessToken != "new-access" {
t.Fatal("refreshed bundle did not retain the existing refresh token")
}
}
func TestServiceInvalidGrantDeletesSession(t *testing.T) {
now := time.Date(2026, 7, 13, 12, 0, 0, 0, time.UTC)
repository := newFakeRepository("subject-1")
verifier := fakeVerifier{users: map[string]*auth.User{
"access": {ID: "subject-1", Source: "oidc", TokenExpiresAt: now.Add(5 * time.Minute)},
}}
client := &fakePublicClient{refreshError: auth.ErrOIDCInvalidGrant}
service := newTestService(t, repository, verifier, client)
service.now = func() time.Time { return now }
raw, err := service.Create(context.Background(), TokenBundle{AccessToken: "access", RefreshToken: "revoked-refresh"}, testLocalUser())
if err != nil {
t.Fatal(err)
}
repository.expireAccessToken(now.Add(-time.Second))
if _, err := service.Resolve(context.Background(), raw); !errors.Is(err, ErrSessionExpired) {
t.Fatalf("Resolve() error = %v, want ErrSessionExpired", err)
}
if !repository.isDeleted() || client.refreshCalls.Load() != 1 {
t.Fatal("invalid_grant did not delete the session after exactly one refresh")
}
}
func TestServiceUsesStillValidAccessTokenWhenRefreshIsTemporarilyUnavailable(t *testing.T) {
now := time.Date(2026, 7, 13, 12, 0, 0, 0, time.UTC)
repository := newFakeRepository("subject-1")
verifier := fakeVerifier{users: map[string]*auth.User{
"access": {ID: "subject-1", Source: "oidc", TokenExpiresAt: now.Add(30 * time.Second)},
}}
client := &fakePublicClient{refreshError: context.DeadlineExceeded}
service := newTestService(t, repository, verifier, client)
service.now = func() time.Time { return now }
raw, err := service.Create(context.Background(), TokenBundle{AccessToken: "access", RefreshToken: "refresh"}, testLocalUser())
if err != nil {
t.Fatal(err)
}
user, err := service.Resolve(context.Background(), raw)
if err != nil || user.ID != "subject-1" {
t.Fatalf("Resolve() user=%#v error=%v", user, err)
}
if client.refreshCalls.Load() != 1 || repository.isDeleted() {
t.Fatal("temporary refresh failure did not fall back to the valid access token")
}
}
func TestServiceReturnsUnavailableWhenExpiredTokenCannotRefresh(t *testing.T) {
now := time.Date(2026, 7, 13, 12, 0, 0, 0, time.UTC)
repository := newFakeRepository("subject-1")
verifier := fakeVerifier{users: map[string]*auth.User{
"access": {ID: "subject-1", Source: "oidc", TokenExpiresAt: now.Add(5 * time.Minute)},
}}
service := newTestService(t, repository, verifier, &fakePublicClient{refreshError: context.DeadlineExceeded})
service.now = func() time.Time { return now }
raw, err := service.Create(context.Background(), TokenBundle{AccessToken: "access", RefreshToken: "refresh"}, testLocalUser())
if err != nil {
t.Fatal(err)
}
repository.expireAccessToken(now.Add(-time.Second))
if _, err := service.Resolve(context.Background(), raw); !errors.Is(err, ErrSessionRefreshUnavailable) {
t.Fatalf("Resolve() error = %v, want ErrSessionRefreshUnavailable", err)
}
}
func TestServiceRejectsWrongEncryptionKeyAndDisabledUser(t *testing.T) {
now := time.Date(2026, 7, 13, 12, 0, 0, 0, time.UTC)
repository := newFakeRepository("subject-1")
verifier := fakeVerifier{users: map[string]*auth.User{
"access": {ID: "subject-1", Source: "oidc", TokenExpiresAt: now.Add(5 * time.Minute)},
}}
service := newTestService(t, repository, verifier, &fakePublicClient{})
service.now = func() time.Time { return now }
raw, err := service.Create(context.Background(), TokenBundle{AccessToken: "access", RefreshToken: "refresh"}, testLocalUser())
if err != nil {
t.Fatal(err)
}
wrongCipher, err := NewCipher(bytes.Repeat([]byte{8}, 32))
if err != nil {
t.Fatal(err)
}
wrongKeyService, err := NewService(repository, wrongCipher, verifier, &fakePublicClient{}, Config{
IdleTTL: 30 * time.Minute, AbsoluteTTL: 8 * time.Hour, RefreshBefore: time.Minute,
})
if err != nil {
t.Fatal(err)
}
wrongKeyService.now = func() time.Time { return now }
if _, err := wrongKeyService.Resolve(context.Background(), raw); !errors.Is(err, ErrSessionStoreUnavailable) {
t.Fatalf("wrong-key Resolve() error = %v, want ErrSessionStoreUnavailable", err)
}
repository.disableUser()
if _, err := service.Resolve(context.Background(), raw); !errors.Is(err, ErrGatewayUserDisabled) {
t.Fatalf("disabled-user Resolve() error = %v, want ErrGatewayUserDisabled", err)
}
}
func TestServiceDeletesSessionEvenWhenCiphertextCannotBeDecryptedDuringLogout(t *testing.T) {
now := time.Date(2026, 7, 13, 12, 0, 0, 0, time.UTC)
repository := newFakeRepository("subject-1")
verifier := fakeVerifier{users: map[string]*auth.User{
"access": {ID: "subject-1", Source: "oidc", TokenExpiresAt: now.Add(5 * time.Minute)},
}}
service := newTestService(t, repository, verifier, &fakePublicClient{})
service.now = func() time.Time { return now }
raw, err := service.Create(context.Background(), TokenBundle{AccessToken: "access", RefreshToken: "refresh"}, testLocalUser())
if err != nil {
t.Fatal(err)
}
repository.corruptCiphertext()
if _, err := service.Delete(context.Background(), raw); err != nil {
t.Fatalf("Delete() error = %v", err)
}
if !repository.isDeleted() {
t.Fatal("logout left a session with unusable ciphertext in the store")
}
}
func testLocalUser() *auth.User {
return &auth.User{
ID: "subject-1", GatewayUserID: "11111111-1111-4111-8111-111111111111",
GatewayTenantID: "22222222-2222-4222-8222-222222222222",
}
}
func newTestService(t *testing.T, repository Repository, verifier TokenVerifier, client PublicClient) *Service {
t.Helper()
cipher, err := NewCipher(bytes.Repeat([]byte{7}, 32))
if err != nil {
t.Fatal(err)
}
service, err := NewService(repository, cipher, verifier, client, Config{
IdleTTL: 30 * time.Minute, AbsoluteTTL: 8 * time.Hour, RefreshBefore: time.Minute,
RefreshLease: 5 * time.Second, RefreshWait: 2 * time.Second,
})
if err != nil {
t.Fatal(err)
}
return service
}
type fakeVerifier struct{ users map[string]*auth.User }
func (f fakeVerifier) Verify(_ context.Context, token string) (*auth.User, error) {
user := f.users[token]
if user == nil {
return nil, auth.ErrUnauthorized
}
copy := *user
return &copy, nil
}
type fakePublicClient struct {
refreshResponse auth.OIDCTokenResponse
refreshError error
delay time.Duration
refreshCalls atomic.Int64
}
func (f *fakePublicClient) Refresh(_ context.Context, _ string) (auth.OIDCTokenResponse, error) {
f.refreshCalls.Add(1)
if f.delay > 0 {
time.Sleep(f.delay)
}
return f.refreshResponse, f.refreshError
}
func (f *fakePublicClient) RevokeRefreshToken(context.Context, string) error { return nil }
func (f *fakePublicClient) EndSessionURL(context.Context, string) (string, error) {
return "https://gateway.example.com/", nil
}
type fakeRepository struct {
mu sync.Mutex
record store.OIDCSession
deleted bool
external string
}
func newFakeRepository(external string) *fakeRepository { return &fakeRepository{external: external} }
func (f *fakeRepository) CreateOIDCSession(_ context.Context, input store.CreateOIDCSessionInput) (store.OIDCSession, error) {
f.mu.Lock()
defer f.mu.Unlock()
f.record = store.OIDCSession{
ID: "33333333-3333-4333-8333-333333333333", SessionTokenHash: append([]byte(nil), input.SessionTokenHash...),
GatewayUserID: input.GatewayUserID, GatewayTenantID: input.GatewayTenantID, ExternalUserID: f.external,
UserStatus: "active", TokenCiphertext: append([]byte(nil), input.TokenCiphertext...), AccessTokenExpiresAt: input.AccessTokenExpiresAt,
LastSeenAt: input.LastSeenAt, IdleExpiresAt: input.IdleExpiresAt, AbsoluteExpiresAt: input.AbsoluteExpiresAt, RefreshVersion: 1,
}
return f.record, nil
}
func (f *fakeRepository) FindOIDCSessionByHash(_ context.Context, hash []byte) (store.OIDCSession, error) {
f.mu.Lock()
defer f.mu.Unlock()
if f.deleted || !bytes.Equal(hash, f.record.SessionTokenHash) {
return store.OIDCSession{}, store.ErrOIDCSessionNotFound
}
item := f.record
item.SessionTokenHash = append([]byte(nil), f.record.SessionTokenHash...)
item.TokenCiphertext = append([]byte(nil), f.record.TokenCiphertext...)
return item, nil
}
func (f *fakeRepository) TouchOIDCSession(_ context.Context, _ string, lastSeen, idle time.Time) error {
f.mu.Lock()
defer f.mu.Unlock()
if f.deleted || !f.record.IdleExpiresAt.After(lastSeen) || !f.record.AbsoluteExpiresAt.After(lastSeen) {
return store.ErrOIDCSessionNotFound
}
f.record.LastSeenAt, f.record.IdleExpiresAt = lastSeen, idle
return nil
}
func (f *fakeRepository) AcquireOIDCSessionRefreshLock(_ context.Context, _ string, version int64, lockID string, until, now time.Time) (bool, error) {
f.mu.Lock()
defer f.mu.Unlock()
if f.deleted || f.record.RefreshVersion != version || f.record.RefreshLockID != "" && f.record.RefreshLockUntil.After(now) {
return false, nil
}
f.record.RefreshLockID, f.record.RefreshLockUntil = lockID, until
return true, nil
}
func (f *fakeRepository) CompleteOIDCSessionRefresh(_ context.Context, _ string, version int64, lockID string, ciphertext []byte, expires time.Time) error {
f.mu.Lock()
defer f.mu.Unlock()
if f.deleted || f.record.RefreshVersion != version || f.record.RefreshLockID != lockID {
return store.ErrOIDCSessionNotFound
}
f.record.TokenCiphertext = append([]byte(nil), ciphertext...)
f.record.AccessTokenExpiresAt = expires
f.record.RefreshVersion++
f.record.RefreshLockID = ""
f.record.RefreshLockUntil = time.Time{}
return nil
}
func (f *fakeRepository) DeleteOIDCSessionByHash(context.Context, []byte) error {
f.mu.Lock()
defer f.mu.Unlock()
f.deleted = true
return nil
}
func (f *fakeRepository) DeleteOIDCSessionByID(context.Context, string) error {
f.mu.Lock()
defer f.mu.Unlock()
f.deleted = true
return nil
}
func (f *fakeRepository) CleanupExpiredOIDCSessions(context.Context, time.Time) (int64, error) {
return 0, nil
}
func (f *fakeRepository) snapshot() store.OIDCSession {
f.mu.Lock()
defer f.mu.Unlock()
item := f.record
item.TokenCiphertext = append([]byte(nil), item.TokenCiphertext...)
return item
}
func (f *fakeRepository) expireAccessToken(expiry time.Time) {
f.mu.Lock()
defer f.mu.Unlock()
f.record.AccessTokenExpiresAt = expiry
}
func (f *fakeRepository) expireIdle(expiry time.Time) {
f.mu.Lock()
defer f.mu.Unlock()
f.record.IdleExpiresAt = expiry
}
func (f *fakeRepository) disableUser() {
f.mu.Lock()
defer f.mu.Unlock()
f.record.UserStatus = "disabled"
}
func (f *fakeRepository) corruptCiphertext() {
f.mu.Lock()
defer f.mu.Unlock()
f.record.TokenCiphertext = []byte("corrupt")
}
func (f *fakeRepository) isDeleted() bool {
f.mu.Lock()
defer f.mu.Unlock()
return f.deleted
}