@@ -153,22 +153,22 @@ def extract_session_from_headers(headers: Dict[str, str]) -> Optional[str]:
153153 # Try Authorization header for Bearer token
154154 auth_header = headers .get ("authorization" ) or headers .get ("Authorization" )
155155 if auth_header and auth_header .lower ().startswith ("bearer " ):
156- # Extract bearer token and try to find associated session
157156 token = auth_header [7 :] # Remove "Bearer " prefix
157+ # Intentionally ignore empty tokens - "Bearer " with no token should not
158+ # create a session context (avoids hash collisions on empty string)
158159 if token :
159- # Look for a session that has this access token
160- # This requires scanning sessions, but bearer tokens should be unique
160+ # Use thread-safe lookup to find session by access token
161161 store = get_oauth21_session_store ()
162- for user_email , session_info in store ._sessions . items ():
163- if session_info . get ( "access_token" ) == token :
164- return session_info . get ( " session_id" ) or f"bearer_ { user_email } "
162+ session_id = store .find_session_id_for_access_token ( token )
163+ if session_id :
164+ return session_id
165165
166- # If no session found, create a temporary session ID from token hash
167- # This allows header-based authentication to work with session context
168- import hashlib
166+ # If no session found, create a temporary session ID from token hash
167+ # This allows header-based authentication to work with session context
168+ import hashlib
169169
170- token_hash = hashlib .sha256 (token .encode ()).hexdigest ()[:8 ]
171- return f"bearer_token_{ token_hash } "
170+ token_hash = hashlib .sha256 (token .encode ()).hexdigest ()[:8 ]
171+ return f"bearer_token_{ token_hash } "
172172
173173 return None
174174
@@ -325,6 +325,32 @@ def store_session(
325325 """
326326 with self ._lock :
327327 normalized_expiry = _normalize_expiry_to_naive_utc (expiry )
328+
329+ # Clean up previous session mappings for this user before storing new one
330+ old_session = self ._sessions .get (user_email )
331+ if old_session :
332+ old_mcp_session_id = old_session .get ("mcp_session_id" )
333+ old_session_id = old_session .get ("session_id" )
334+ # Remove old MCP session mapping if it differs from new one
335+ if old_mcp_session_id and old_mcp_session_id != mcp_session_id :
336+ if old_mcp_session_id in self ._mcp_session_mapping :
337+ del self ._mcp_session_mapping [old_mcp_session_id ]
338+ logger .debug (
339+ f"Removed stale MCP session mapping: { old_mcp_session_id } "
340+ )
341+ if old_mcp_session_id in self ._session_auth_binding :
342+ del self ._session_auth_binding [old_mcp_session_id ]
343+ logger .debug (
344+ f"Removed stale auth binding: { old_mcp_session_id } "
345+ )
346+ # Remove old OAuth session binding if it differs from new one
347+ if old_session_id and old_session_id != session_id :
348+ if old_session_id in self ._session_auth_binding :
349+ del self ._session_auth_binding [old_session_id ]
350+ logger .debug (
351+ f"Removed stale OAuth session binding: { old_session_id } "
352+ )
353+
328354 session_info = {
329355 "access_token" : access_token ,
330356 "refresh_token" : refresh_token ,
@@ -570,6 +596,9 @@ def remove_session(self, user_email: str):
570596 if not mcp_session_id :
571597 logger .info (f"Removed OAuth 2.1 session for { user_email } " )
572598
599+ # Clean up any orphaned mappings that may have accumulated
600+ self ._cleanup_orphaned_mappings_locked ()
601+
573602 def has_session (self , user_email : str ) -> bool :
574603 """Check if a user has an active session."""
575604 with self ._lock :
@@ -597,6 +626,71 @@ def get_stats(self) -> Dict[str, Any]:
597626 "mcp_sessions" : list (self ._mcp_session_mapping .keys ()),
598627 }
599628
629+ def find_session_id_for_access_token (self , token : str ) -> Optional [str ]:
630+ """
631+ Thread-safe lookup of session ID by access token.
632+
633+ Args:
634+ token: The access token to search for
635+
636+ Returns:
637+ Session ID if found, None otherwise
638+ """
639+ with self ._lock :
640+ for user_email , session_info in self ._sessions .items ():
641+ if session_info .get ("access_token" ) == token :
642+ return session_info .get ("session_id" ) or f"bearer_{ user_email } "
643+ return None
644+
645+ def _cleanup_orphaned_mappings_locked (self ) -> int :
646+ """Remove orphaned mappings. Caller must hold lock."""
647+ # Collect valid session IDs and mcp_session_ids from active sessions
648+ valid_session_ids = set ()
649+ valid_mcp_session_ids = set ()
650+ for session_info in self ._sessions .values ():
651+ if session_info .get ("session_id" ):
652+ valid_session_ids .add (session_info ["session_id" ])
653+ if session_info .get ("mcp_session_id" ):
654+ valid_mcp_session_ids .add (session_info ["mcp_session_id" ])
655+
656+ removed = 0
657+
658+ # Remove orphaned MCP session mappings
659+ orphaned_mcp = [
660+ sid for sid in self ._mcp_session_mapping
661+ if sid not in valid_mcp_session_ids
662+ ]
663+ for sid in orphaned_mcp :
664+ del self ._mcp_session_mapping [sid ]
665+ removed += 1
666+ logger .debug (f"Removed orphaned MCP session mapping: { sid } " )
667+
668+ # Remove orphaned auth bindings
669+ valid_bindings = valid_session_ids | valid_mcp_session_ids
670+ orphaned_bindings = [
671+ sid for sid in self ._session_auth_binding
672+ if sid not in valid_bindings
673+ ]
674+ for sid in orphaned_bindings :
675+ del self ._session_auth_binding [sid ]
676+ removed += 1
677+ logger .debug (f"Removed orphaned auth binding: { sid } " )
678+
679+ if removed > 0 :
680+ logger .info (f"Cleaned up { removed } orphaned session mappings/bindings" )
681+
682+ return removed
683+
684+ def cleanup_orphaned_mappings (self ) -> int :
685+ """
686+ Remove orphaned entries from mcp_session_mapping and session_auth_binding.
687+
688+ Returns:
689+ Number of orphaned entries removed
690+ """
691+ with self ._lock :
692+ return self ._cleanup_orphaned_mappings_locked ()
693+
600694
601695# Global instance
602696_global_store = OAuth21SessionStore ()
0 commit comments