vibe-proxy/backend/sdk/cliproxy/auth/session_cache.go
2026-08-24 00:10:41 +02:00

353 lines
8.5 KiB
Go

package auth
import (
"strings"
"sync"
"time"
)
const maxStableSessionAliases = 64
// sessionEntry stores an auth binding, its identifier aliases, and expiration.
type sessionEntry struct {
authID string
expiresAt time.Time
aliases []string
}
// SessionCache provides TTL-based session to auth mapping with automatic cleanup.
type SessionCache struct {
mu sync.RWMutex
entries map[string]sessionEntry
ttl time.Duration
stopCh chan struct{}
stopOnce sync.Once
}
// NewSessionCache creates a cache with the specified TTL.
// A background goroutine periodically cleans expired entries.
func NewSessionCache(ttl time.Duration) *SessionCache {
if ttl <= 0 {
ttl = 30 * time.Minute
}
c := &SessionCache{
entries: make(map[string]sessionEntry),
ttl: ttl,
stopCh: make(chan struct{}),
}
go c.cleanupLoop()
return c
}
// Get retrieves the auth ID bound to a session, if still valid.
// Does NOT refresh the TTL on access.
func (c *SessionCache) Get(sessionID string) (string, bool) {
if sessionID == "" {
return "", false
}
now := time.Now()
c.mu.RLock()
entry, ok := c.entries[sessionID]
if ok && now.Before(entry.expiresAt) {
c.mu.RUnlock()
return entry.authID, true
}
c.mu.RUnlock()
if !ok {
return "", false
}
c.mu.Lock()
defer c.mu.Unlock()
entry, ok = c.entries[sessionID]
if !ok {
return "", false
}
if time.Now().Before(entry.expiresAt) {
return entry.authID, true
}
c.removeAliasGroupLocked(entry)
return "", false
}
// GetAndRefresh retrieves the auth ID bound to a session and refreshes the TTL
// for every identifier known to represent the same logical session.
func (c *SessionCache) GetAndRefresh(sessionID string) (string, bool) {
if sessionID == "" {
return "", false
}
now := time.Now()
c.mu.Lock()
defer c.mu.Unlock()
entry, ok := c.entries[sessionID]
if !ok {
return "", false
}
if !now.Before(entry.expiresAt) {
c.removeAliasGroupLocked(entry)
return "", false
}
aliases := compactSessionAliases(mergeSessionAliases([]string{sessionID}, entry.aliases...))
c.replaceAliasGroupsLocked(entry.authID, now.Add(c.ttl), aliases, entry)
return entry.authID, true
}
// Set binds a session to an auth ID with TTL refresh. Existing aliases for the
// same logical session remain attached when the binding is refreshed or moved.
func (c *SessionCache) Set(sessionID, authID string) {
c.SetAliases(authID, sessionID)
}
// SetAliases binds multiple identifiers for one logical session to an auth ID.
func (c *SessionCache) SetAliases(authID string, sessionIDs ...string) {
if authID == "" {
return
}
now := time.Now()
c.mu.Lock()
defer c.mu.Unlock()
aliases := mergeSessionAliases(nil, sessionIDs...)
previousGroups := make([]sessionEntry, 0, len(sessionIDs))
for _, sessionID := range sessionIDs {
entry, ok := c.entries[sessionID]
if !ok {
continue
}
if !now.Before(entry.expiresAt) {
c.removeAliasGroupLocked(entry)
continue
}
previousGroups = append(previousGroups, entry)
aliases = mergeSessionAliases(aliases, entry.aliases...)
}
aliases = compactSessionAliases(aliases)
if len(aliases) == 0 {
return
}
c.replaceAliasGroupsLocked(authID, now.Add(c.ttl), aliases, previousGroups...)
}
func (c *SessionCache) replaceAliasGroupsLocked(authID string, expiresAt time.Time, aliases []string, previousGroups ...sessionEntry) {
for _, previous := range previousGroups {
c.removeAliasGroupLocked(previous)
}
entry := sessionEntry{authID: authID, expiresAt: expiresAt, aliases: aliases}
for _, alias := range aliases {
c.entries[alias] = entry
}
}
func (c *SessionCache) removeAliasGroupLocked(entry sessionEntry) {
for _, alias := range entry.aliases {
current, ok := c.entries[alias]
if !ok || current.authID != entry.authID || !current.expiresAt.Equal(entry.expiresAt) ||
!equalSessionAliases(current.aliases, entry.aliases) {
continue
}
delete(c.entries, alias)
}
}
func compactSessionAliases(aliases []string) []string {
return compactSessionAliasesWith(aliases, isLocalPromptCacheSessionAlias)
}
func compactHomeSessionAliases(aliases []string) []string {
return compactSessionAliasesWith(aliases, func(alias string) bool {
return strings.HasPrefix(alias, "pck:")
})
}
func compactSessionAliasesWith(aliases []string, isPromptCacheAlias func(string) bool) []string {
compacted := make([]string, 0, len(aliases))
hasPromptCacheKey := false
stableAliases := 0
for _, alias := range aliases {
if isPromptCacheAlias(alias) {
if hasPromptCacheKey {
continue
}
hasPromptCacheKey = true
} else {
if stableAliases >= maxStableSessionAliases {
continue
}
stableAliases++
}
compacted = append(compacted, alias)
}
return compacted
}
func isLocalPromptCacheSessionAlias(alias string) bool {
if strings.HasPrefix(alias, "pck:") {
return true
}
_, sessionAndModel, ok := strings.Cut(alias, "::")
return ok && strings.HasPrefix(sessionAndModel, "pck:")
}
func equalSessionAliases(left, right []string) bool {
if len(left) != len(right) {
return false
}
for index := range left {
if left[index] != right[index] {
return false
}
}
return true
}
func mergeSessionAliases(existing []string, candidates ...string) []string {
aliases := make([]string, 0, len(existing)+len(candidates))
seen := make(map[string]struct{}, cap(aliases))
add := func(alias string) {
if alias == "" {
return
}
if _, ok := seen[alias]; ok {
return
}
seen[alias] = struct{}{}
aliases = append(aliases, alias)
}
for _, alias := range existing {
add(alias)
}
for _, alias := range candidates {
add(alias)
}
return aliases
}
// Touch refreshes the expiration for a session binding if it currently matches expectedAuthID.
func (c *SessionCache) Touch(sessionID, expectedAuthID string) bool {
if sessionID == "" || expectedAuthID == "" {
return false
}
now := time.Now()
c.mu.Lock()
defer c.mu.Unlock()
entry, ok := c.entries[sessionID]
if !ok || entry.authID != expectedAuthID || !now.Before(entry.expiresAt) {
return false
}
aliases := compactSessionAliases(mergeSessionAliases([]string{sessionID}, entry.aliases...))
c.replaceAliasGroupsLocked(expectedAuthID, now.Add(c.ttl), aliases, entry)
return true
}
// CompareAndDelete removes the session binding only if it is currently bound to expectedAuthID.
func (c *SessionCache) CompareAndDelete(sessionID, expectedAuthID string) bool {
if sessionID == "" || expectedAuthID == "" {
return false
}
c.mu.Lock()
defer c.mu.Unlock()
entry, ok := c.entries[sessionID]
if !ok || entry.authID != expectedAuthID {
return false
}
delete(c.entries, sessionID)
for _, alias := range entry.aliases {
if alias == sessionID {
continue
}
current, exists := c.entries[alias]
if !exists || current.authID != entry.authID {
continue
}
filtered := make([]string, 0, len(current.aliases))
for _, candidate := range current.aliases {
if candidate != sessionID {
filtered = append(filtered, candidate)
}
}
current.aliases = filtered
c.entries[alias] = current
}
return true
}
// Invalidate removes a specific session binding without allowing another alias
// in the same group to recreate it on its next refresh.
func (c *SessionCache) Invalidate(sessionID string) {
if sessionID == "" {
return
}
c.mu.Lock()
entry, ok := c.entries[sessionID]
delete(c.entries, sessionID)
if ok {
for _, alias := range entry.aliases {
if alias == sessionID {
continue
}
current, exists := c.entries[alias]
if !exists || current.authID != entry.authID {
continue
}
filtered := make([]string, 0, len(current.aliases))
for _, candidate := range current.aliases {
if candidate != sessionID {
filtered = append(filtered, candidate)
}
}
current.aliases = filtered
c.entries[alias] = current
}
}
c.mu.Unlock()
}
// InvalidateAuth removes all sessions bound to a specific auth ID.
// Used when an auth becomes unavailable.
func (c *SessionCache) InvalidateAuth(authID string) {
if authID == "" {
return
}
c.mu.Lock()
for sid, entry := range c.entries {
if entry.authID == authID {
delete(c.entries, sid)
}
}
c.mu.Unlock()
}
// Stop terminates the background cleanup goroutine.
func (c *SessionCache) Stop() {
if c == nil {
return
}
c.stopOnce.Do(func() {
close(c.stopCh)
})
}
func (c *SessionCache) cleanupLoop() {
ticker := time.NewTicker(c.ttl / 2)
defer ticker.Stop()
for {
select {
case <-c.stopCh:
return
case <-ticker.C:
c.cleanup()
}
}
}
func (c *SessionCache) cleanup() {
now := time.Now()
c.mu.Lock()
for sid, entry := range c.entries {
if !now.Before(entry.expiresAt) {
delete(c.entries, sid)
}
}
c.mu.Unlock()
}