Skip to content

Commit ba7e760

Browse files
committed
feat: improve session management with async locking and enhance WebSocket broadcast error handling
1 parent 4ef38e4 commit ba7e760

2 files changed

Lines changed: 99 additions & 87 deletions

File tree

form-flow-backend/services/ai/session_manager.py

Lines changed: 91 additions & 85 deletions
Original file line numberDiff line numberDiff line change
@@ -45,16 +45,18 @@ def __init__(self, redis_client=None):
4545
self._redis = redis_client
4646
self._local_cache: Dict[str, Dict[str, Any]] = {}
4747
self._use_redis = True
48+
self._lock = asyncio.Lock()
4849

4950
async def _get_redis(self):
5051
"""Get Redis client, falling back to local cache if unavailable."""
51-
if self._redis is None:
52-
try:
53-
self._redis = await get_redis_client()
54-
except Exception as e:
55-
logger.warning(f"Redis unavailable, using local cache: {e}")
56-
self._use_redis = False
57-
return self._redis
52+
async with self._lock:
53+
if self._redis is None:
54+
try:
55+
self._redis = await get_redis_client()
56+
except Exception as e:
57+
logger.warning(f"Redis unavailable, using local cache: {e}")
58+
self._use_redis = False
59+
return self._redis
5860

5961
async def save_session(self, session_data: Dict[str, Any]) -> bool:
6062
"""
@@ -74,29 +76,30 @@ async def save_session(self, session_data: Dict[str, Any]) -> bool:
7476
# Serialize datetime objects
7577
serialized = self._serialize_session(session_data)
7678

77-
if self._use_redis:
78-
try:
79-
redis = await self._get_redis()
80-
if redis:
81-
key = f"{self.SESSION_PREFIX}{session_id}"
82-
await redis.setex(
83-
key,
84-
timedelta(minutes=self.SESSION_TTL_MINUTES),
85-
json.dumps(serialized)
86-
)
87-
logger.debug(f"Saved session {session_id} to Redis")
88-
return True
89-
except Exception as e:
90-
logger.warning(f"Redis save failed, using local cache: {e}")
91-
self._use_redis = False
92-
93-
# Fallback to local cache
94-
self._local_cache[session_id] = {
95-
'data': serialized,
96-
'expires_at': datetime.now() + timedelta(minutes=self.SESSION_TTL_MINUTES)
97-
}
98-
logger.debug(f"Saved session {session_id} to local cache")
99-
return True
79+
async with self._lock:
80+
if self._use_redis:
81+
try:
82+
redis = await self._get_redis()
83+
if redis:
84+
key = f"{self.SESSION_PREFIX}{session_id}"
85+
await redis.setex(
86+
key,
87+
timedelta(minutes=self.SESSION_TTL_MINUTES),
88+
json.dumps(serialized)
89+
)
90+
logger.debug(f"Saved session {session_id} to Redis")
91+
return True
92+
except Exception as e:
93+
logger.warning(f"Redis save failed, using local cache: {e}")
94+
self._use_redis = False
95+
96+
# Fallback to local cache
97+
self._local_cache[session_id] = {
98+
'data': serialized,
99+
'expires_at': datetime.now() + timedelta(minutes=self.SESSION_TTL_MINUTES)
100+
}
101+
logger.debug(f"Saved session {session_id} to local cache")
102+
return True
100103

101104
async def get_session(self, session_id: str) -> Optional[Dict[str, Any]]:
102105
"""
@@ -108,67 +111,70 @@ async def get_session(self, session_id: str) -> Optional[Dict[str, Any]]:
108111
Returns:
109112
Session data dictionary or None if not found/expired
110113
"""
111-
if self._use_redis:
112-
try:
113-
redis = await self._get_redis()
114-
if redis:
115-
key = f"{self.SESSION_PREFIX}{session_id}"
116-
data = await redis.get(key)
117-
if data:
118-
session = json.loads(data)
119-
return self._deserialize_session(session)
120-
except Exception as e:
121-
logger.warning(f"Redis get failed: {e}")
122-
self._use_redis = False
123-
124-
# Check local cache
125-
cached = self._local_cache.get(session_id)
126-
if cached:
127-
if cached['expires_at'] > datetime.now():
128-
return self._deserialize_session(cached['data'])
129-
else:
130-
del self._local_cache[session_id]
131-
132-
return None
114+
async with self._lock:
115+
if self._use_redis:
116+
try:
117+
redis = await self._get_redis()
118+
if redis:
119+
key = f"{self.SESSION_PREFIX}{session_id}"
120+
data = await redis.get(key)
121+
if data:
122+
session = json.loads(data)
123+
return self._deserialize_session(session)
124+
except Exception as e:
125+
logger.warning(f"Redis get failed: {e}")
126+
self._use_redis = False
127+
128+
# Check local cache
129+
cached = self._local_cache.get(session_id)
130+
if cached:
131+
if cached['expires_at'] > datetime.now():
132+
return self._deserialize_session(cached['data'])
133+
else:
134+
del self._local_cache[session_id]
135+
136+
return None
133137

134138
async def delete_session(self, session_id: str) -> bool:
135139
"""Delete a session."""
136-
if self._use_redis:
137-
try:
138-
redis = await self._get_redis()
139-
if redis:
140-
key = f"{self.SESSION_PREFIX}{session_id}"
141-
await redis.delete(key)
142-
logger.debug(f"Deleted session {session_id} from Redis")
143-
except Exception as e:
144-
logger.warning(f"Redis delete failed: {e}")
145-
146-
# Also remove from local cache
147-
if session_id in self._local_cache:
148-
del self._local_cache[session_id]
149-
150-
return True
140+
async with self._lock:
141+
if self._use_redis:
142+
try:
143+
redis = await self._get_redis()
144+
if redis:
145+
key = f"{self.SESSION_PREFIX}{session_id}"
146+
await redis.delete(key)
147+
logger.debug(f"Deleted session {session_id} from Redis")
148+
except Exception as e:
149+
logger.warning(f"Redis delete failed: {e}")
150+
151+
# Also remove from local cache
152+
if session_id in self._local_cache:
153+
del self._local_cache[session_id]
154+
155+
return True
151156

152157
async def extend_session(self, session_id: str) -> bool:
153158
"""Extend session TTL by the standard amount."""
154-
if self._use_redis:
155-
try:
156-
redis = await self._get_redis()
157-
if redis:
158-
key = f"{self.SESSION_PREFIX}{session_id}"
159-
await redis.expire(key, timedelta(minutes=self.SESSION_TTL_MINUTES))
160-
return True
161-
except Exception as e:
162-
logger.warning(f"Redis expire failed: {e}")
163-
164-
# Extend local cache
165-
if session_id in self._local_cache:
166-
self._local_cache[session_id]['expires_at'] = (
167-
datetime.now() + timedelta(minutes=self.SESSION_TTL_MINUTES)
168-
)
169-
return True
170-
171-
return False
159+
async with self._lock:
160+
if self._use_redis:
161+
try:
162+
redis = await self._get_redis()
163+
if redis:
164+
key = f"{self.SESSION_PREFIX}{session_id}"
165+
await redis.expire(key, timedelta(minutes=self.SESSION_TTL_MINUTES))
166+
return True
167+
except Exception as e:
168+
logger.warning(f"Redis expire failed: {e}")
169+
170+
# Extend local cache
171+
if session_id in self._local_cache:
172+
self._local_cache[session_id]['expires_at'] = (
173+
datetime.now() + timedelta(minutes=self.SESSION_TTL_MINUTES)
174+
)
175+
return True
176+
177+
return False
172178

173179
async def cleanup_local_cache(self) -> int:
174180
"""Remove expired sessions from local cache. Returns count removed."""

form-flow-backend/utils/ws.py

Lines changed: 8 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -42,12 +42,18 @@ async def broadcast_json(self, message: Any, client_id: Optional[str] = None):
4242
for conns in self._connections.values():
4343
targets.extend(list(conns))
4444

45+
# Collect dead connections to disconnect after iteration
46+
dead_connections = []
4547
for ws in targets:
4648
try:
4749
await ws.send_json(message)
4850
except Exception:
49-
# Drop dead connections silently
50-
await self.disconnect(ws, client_id or "global")
51+
# Mark dead connection for removal after iteration
52+
dead_connections.append(ws)
53+
54+
# Disconnect dead connections outside the broadcast loop
55+
for ws in dead_connections:
56+
await self.disconnect(ws, client_id or "global")
5157

5258

5359
# Singleton manager

0 commit comments

Comments
 (0)