package auth // Reference: M13 §4 — SSO Session Store (Redis-backed) // Tracks active SSO sessions across the system for single sign-on and single logout. // An SSO session is created when a user authenticates via SAML/OAuth and persists // until the user explicitly logs out or the session expires. // This enables: (1) SSO — one login grants access to all accounts for that user, // (2) SLO — logout terminates all active sessions across all accounts, // (3) audit — track when/where a user authenticated and from which IdP. import ( "context" "crypto/rand" "encoding/base64" "encoding/json" "fmt" "time" "github.com/redis/go-redis/v9" applogger "github.com/gochat/gochat/pkg/logger" ) // SSOSessionData holds the data stored in Redis for an active SSO session. type SSOSessionData struct { SessionID string `json:"session_id"` // unique SSO session identifier UserID uint `json:"user_id"` // user who owns this session Provider string `json:"provider"` // saml, google, github IdPEntityID string `json:"idp_entity_id"` // SAML IdP entity ID (for SAML sessions) NameID string `json:"name_id"` // SAML NameID from assertion AccountID uint `json:"account_id"` // account context for this session Role string `json:"role"` // role in the account CreatedAt int64 `json:"created_at"` // Unix timestamp of creation ExpiresAt int64 `json:"expires_at"` // Unix timestamp of expiry } // SSOSessionStore manages SSO sessions in Redis for SSO/SLO support. // Key patterns: // sso:session:{session_id} → SSOSessionData (single session) // sso:user_sessions:{user_id} → SET of session_ids (all sessions for a user) // sso:idp_sessions:{idp_entity_id} → SET of session_ids (all sessions for an IdP, for SLO) type SSOSessionStore struct { rdb redis.Cmdable ttl time.Duration // default session TTL } // NewSSOSessionStore creates a Redis-backed SSO session store. func NewSSOSessionStore(rdb redis.Cmdable, ttl time.Duration) *SSOSessionStore { if ttl == 0 { ttl = 24 * time.Hour // default: 24 hours } return &SSOSessionStore{rdb: rdb, ttl: ttl} } // SessionTTL returns the configured session TTL duration. // Used by handlers to calculate SSO session expiry timestamps. func (s *SSOSessionStore) SessionTTL() time.Duration { return s.ttl } // generateSessionID creates a random, unique SSO session identifier. func generateSSOSessionID() (string, error) { b := make([]byte, 32) if _, err := rand.Read(b); err != nil { return "", fmt.Errorf("failed to generate session ID: %w", err) } return base64.URLEncoding.EncodeToString(b), nil } // Create creates a new SSO session in Redis. // Stores the session data, adds session ID to user and IdP index sets. // Returns the session ID for reference. func (s *SSOSessionStore) Create(ctx context.Context, data *SSOSessionData) (string, error) { if data.SessionID == "" { sessionID, err := generateSSOSessionID() if err != nil { return "", err } data.SessionID = sessionID } // Set creation and expiry timestamps now := time.Now() if data.CreatedAt == 0 { data.CreatedAt = now.Unix() } if data.ExpiresAt == 0 { data.ExpiresAt = now.Add(s.ttl).Unix() } // Serialize session data dataBytes, err := json.Marshal(data) if err != nil { return "", fmt.Errorf("failed to marshal SSO session data: %w", err) } // Calculate TTL based on expires_at expiresAt := time.Unix(data.ExpiresAt, 0) ttl := time.Until(expiresAt) if ttl <= 0 { return "", fmt.Errorf("SSO session expiry is in the past") } sessionKey := fmt.Sprintf("sso:session:%s", data.SessionID) userKey := fmt.Sprintf("sso:user_sessions:%d", data.UserID) idpKey := fmt.Sprintf("sso:idp_sessions:%s", data.IdPEntityID) // Use pipeline for atomic multi-key operations pipe := s.rdb.Pipeline() pipe.Set(ctx, sessionKey, dataBytes, ttl) pipe.SAdd(ctx, userKey, data.SessionID) pipe.SAdd(ctx, idpKey, data.SessionID) // Set TTL on index sets (longer than session TTL to allow cleanup) pipe.Expire(ctx, userKey, ttl+time.Hour) pipe.Expire(ctx, idpKey, ttl+time.Hour) if _, err := pipe.Exec(ctx); err != nil { return "", fmt.Errorf("failed to store SSO session in Redis: %w", err) } applogger.L().Infof("SSO session created (id=%s, user=%d, provider=%s, idp=%s)", data.SessionID, data.UserID, data.Provider, data.IdPEntityID) return data.SessionID, nil } // Get retrieves an SSO session by session ID. func (s *SSOSessionStore) Get(ctx context.Context, sessionID string) (*SSOSessionData, error) { key := fmt.Sprintf("sso:session:%s", sessionID) dataBytes, err := s.rdb.Get(ctx, key).Bytes() if err == redis.Nil { return nil, nil // session not found or expired } if err != nil { return nil, fmt.Errorf("failed to get SSO session: %w", err) } var data SSOSessionData if err := json.Unmarshal(dataBytes, &data); err != nil { return nil, fmt.Errorf("failed to unmarshal SSO session data: %w", err) } return &data, nil } // GetByUser retrieves all active SSO sessions for a user. // Returns a list of SSOSessionData for all sessions indexed under the user. func (s *SSOSessionStore) GetByUser(ctx context.Context, userID uint) ([]*SSOSessionData, error) { userKey := fmt.Sprintf("sso:user_sessions:%d", userID) sessionIDs, err := s.rdb.SMembers(ctx, userKey).Result() if err != nil { return nil, fmt.Errorf("failed to get user SSO sessions: %w", err) } var sessions []*SSOSessionData for _, sid := range sessionIDs { data, err := s.Get(ctx, sid) if err != nil { applogger.L().Warnf("Failed to get SSO session %s for user %d: %v", sid, userID, err) continue } if data != nil { sessions = append(sessions, data) } else { // Session expired but still in index set — clean up s.rdb.SRem(ctx, userKey, sid) } } return sessions, nil } // GetByIdP retrieves all active SSO sessions for an IdP entity ID. // Used for IdP-initiated SLO: terminate all sessions associated with an IdP. func (s *SSOSessionStore) GetByIdP(ctx context.Context, idpEntityID string) ([]*SSOSessionData, error) { idpKey := fmt.Sprintf("sso:idp_sessions:%s", idpEntityID) sessionIDs, err := s.rdb.SMembers(ctx, idpKey).Result() if err != nil { return nil, fmt.Errorf("failed to get IdP SSO sessions: %w", err) } var sessions []*SSOSessionData for _, sid := range sessionIDs { data, err := s.Get(ctx, sid) if err != nil { applogger.L().Warnf("Failed to get SSO session %s for IdP %s: %v", sid, idpEntityID, err) continue } if data != nil { sessions = append(sessions, data) } else { // Session expired but still in index set — clean up s.rdb.SRem(ctx, idpKey, sid) } } return sessions, nil } // Terminate removes an SSO session and cleans up index entries. // Returns true if the session was found and terminated, false if not found. func (s *SSOSessionStore) Terminate(ctx context.Context, sessionID string) (bool, error) { // Get session data first to clean up index sets data, err := s.Get(ctx, sessionID) if err != nil { return false, fmt.Errorf("failed to get SSO session for termination: %w", err) } if data == nil { return false, nil // session not found } sessionKey := fmt.Sprintf("sso:session:%s", sessionID) userKey := fmt.Sprintf("sso:user_sessions:%d", data.UserID) idpKey := fmt.Sprintf("sso:idp_sessions:%s", data.IdPEntityID) pipe := s.rdb.Pipeline() pipe.Del(ctx, sessionKey) pipe.SRem(ctx, userKey, sessionID) pipe.SRem(ctx, idpKey, sessionID) if _, err := pipe.Exec(ctx); err != nil { return false, fmt.Errorf("failed to terminate SSO session in Redis: %w", err) } applogger.L().Infof("SSO session terminated (id=%s, user=%d, provider=%s)", sessionID, data.UserID, data.Provider) return true, nil } // TerminateUserSessions terminates all SSO sessions for a user. // Returns the number of sessions terminated. // Used for SP-initiated SLO or explicit user logout. func (s *SSOSessionStore) TerminateUserSessions(ctx context.Context, userID uint) (int, error) { sessions, err := s.GetByUser(ctx, userID) if err != nil { return 0, fmt.Errorf("failed to get user sessions for termination: %w", err) } count := 0 for _, data := range sessions { terminated, err := s.Terminate(ctx, data.SessionID) if err != nil { applogger.L().Warnf("Failed to terminate SSO session %s: %v", data.SessionID, err) continue } if terminated { count++ } } return count, nil } // TerminateIdPSessions terminates all SSO sessions for an IdP. // Returns the number of sessions terminated. // Used for IdP-initiated SLO: when the IdP sends a LogoutRequest, // we terminate all sessions for that IdP. func (s *SSOSessionStore) TerminateIdPSessions(ctx context.Context, idpEntityID string) (int, error) { sessions, err := s.GetByIdP(ctx, idpEntityID) if err != nil { return 0, fmt.Errorf("failed to get IdP sessions for termination: %w", err) } count := 0 for _, data := range sessions { terminated, err := s.Terminate(ctx, data.SessionID) if err != nil { applogger.L().Warnf("Failed to terminate SSO session %s: %v", data.SessionID, err) continue } if terminated { count++ } } return count, nil } // Refresh extends the TTL of an active SSO session. // Updates the expires_at timestamp and Redis TTL. func (s *SSOSessionStore) Refresh(ctx context.Context, sessionID string, newTTL time.Duration) error { if newTTL == 0 { newTTL = s.ttl } data, err := s.Get(ctx, sessionID) if err != nil { return fmt.Errorf("failed to get SSO session for refresh: %w", err) } if data == nil { return fmt.Errorf("SSO session not found") } // Update expires_at data.ExpiresAt = time.Now().Add(newTTL).Unix() dataBytes, err := json.Marshal(data) if err != nil { return fmt.Errorf("failed to marshal refreshed SSO session data: %w", err) } sessionKey := fmt.Sprintf("sso:session:%s", sessionID) if err := s.rdb.Set(ctx, sessionKey, dataBytes, newTTL).Err(); err != nil { return fmt.Errorf("failed to refresh SSO session in Redis: %w", err) } return nil } // CountByUser returns the number of active SSO sessions for a user. func (s *SSOSessionStore) CountByUser(ctx context.Context, userID uint) (int64, error) { userKey := fmt.Sprintf("sso:user_sessions:%d", userID) // Use SCARD for efficient count without fetching all members count, err := s.rdb.SCard(ctx, userKey).Result() if err != nil { return 0, fmt.Errorf("failed to count user SSO sessions: %w", err) } return count, nil } // Exists checks whether an SSO session exists and is not expired. func (s *SSOSessionStore) Exists(ctx context.Context, sessionID string) (bool, error) { key := fmt.Sprintf("sso:session:%s", sessionID) exists, err := s.rdb.Exists(ctx, key).Result() if err != nil { return false, fmt.Errorf("failed to check SSO session existence: %w", err) } return exists > 0, nil }