feat: 实现 OIDC 服务端会话与请求刷新
使用 AES-256-GCM 保存认证中心令牌,并以随机 Cookie 哈希关联 PostgreSQL 会话。加入闲置与绝对期限、Refresh Token 轮换以及多实例刷新锁。
This commit is contained in:
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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 ©, 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
|
||||
}
|
||||
Reference in New Issue
Block a user