@@ -74,14 +74,17 @@ def extract_user_id_from_credentials(credentials: Dict[str, Any] | None) -> str
7474 Extract user ID from credentials (static utility).
7575
7676 Priority:
77- 1. Direct 'user_id' in credentials
77+ 1. Direct 'user_id' in credentials (validated against JWT if both present)
7878 2. Extract from JWT 'sub' claim
7979
8080 Args:
8181 credentials: Dict containing 'jwt_token' and/or 'user_id'
8282
8383 Returns:
84- User ID if found, None otherwise
84+ User ID if found and validated, None otherwise
85+
86+ Raises:
87+ ValueError: If user_id and JWT user_id don't match
8588 """
8689 if not credentials :
8790 return None
@@ -92,23 +95,38 @@ def extract_user_id_from_credentials(credentials: Dict[str, Any] | None) -> str
9295
9396 # 1. Direct user_id in credentials
9497 user_id = credentials .get ("user_id" )
95- if user_id :
96- return user_id
9798
9899 # 2. Extract from JWT token
99100 jwt_token = credentials .get ("jwt_token" )
100- if not jwt_token :
101- return None
101+ jwt_user_id = None
102102
103- try :
104- payload = jwt .get_unverified_claims (jwt_token )
105- return payload .get ("sub" )
106- except JWTError as e :
107- logger .error (f"Error decoding JWT token: { e } " )
108- return None
109- except Exception as e :
110- logger .error (f"Unexpected error extracting user_id from JWT: { e } " )
111- return None
103+ if jwt_token :
104+ try :
105+ payload = jwt .get_unverified_claims (jwt_token )
106+ jwt_user_id = payload .get ("sub" )
107+ except JWTError as e :
108+ logger .error (f"Error decoding JWT token: { e } " )
109+ return None
110+ except Exception as e :
111+ logger .error (f"Unexpected error extracting user_id from JWT: { e } " )
112+ return None
113+
114+ # If both user_id and JWT token are present, validate they match
115+ if user_id and jwt_user_id :
116+ if user_id != jwt_user_id :
117+ logger .error (
118+ f"User ID mismatch: provided user_id '{ user_id } ' does not match "
119+ f"JWT token user_id '{ jwt_user_id } '"
120+ )
121+ raise ValueError (
122+ "User ID mismatch: provided user_id does not match JWT token. "
123+ "This may indicate a security issue."
124+ )
125+ logger .debug (f"User ID validated: { user_id } matches JWT token" )
126+ return user_id
127+
128+ # Return whichever is available
129+ return user_id or jwt_user_id
112130
113131
114132def get_token_hash (jwt_token : str | None ) -> str | None :
@@ -280,7 +298,18 @@ def prepare_credentials(
280298 user_id : str | None = None ,
281299 auth_token : str | None = None ,
282300 ) -> Dict [str , Any ]:
283- """Prepare Supabase-specific credentials."""
301+ """Prepare Supabase-specific credentials.
302+
303+ Args:
304+ user_id: User identifier (optional)
305+ auth_token: JWT token (optional)
306+
307+ Returns:
308+ Dict with validated user_id and jwt_token
309+
310+ Raises:
311+ ValueError: If user_id and JWT user_id don't match
312+ """
284313 credentials = {"user_id" : user_id }
285314 if auth_token :
286315 credentials ["jwt_token" ] = auth_token
@@ -289,7 +318,19 @@ def prepare_credentials(
289318 return credentials
290319
291320 def extract_user_id (self , credentials : Dict [str , Any ] | None ) -> str | None :
292- """Extract user ID from credentials."""
321+ """Extract user ID from credentials with validation.
322+
323+ If both user_id and JWT token are present, validates they match.
324+
325+ Args:
326+ credentials: Dict containing 'jwt_token' and/or 'user_id'
327+
328+ Returns:
329+ User ID if found and validated, None otherwise
330+
331+ Raises:
332+ ValueError: If user_id and JWT user_id don't match
333+ """
293334 return extract_user_id_from_credentials (credentials )
294335
295336 def _ensure_authenticated (self , credentials : Dict [str , Any ] | None ) -> bool :
0 commit comments