Skip to content

Commit 99f50be

Browse files
committed
chore(auth): validate user_id against jwt_token
1 parent 4246f29 commit 99f50be

1 file changed

Lines changed: 58 additions & 17 deletions

File tree

langmiddle/storage/supabase_backend.py

Lines changed: 58 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -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

114132
def 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

Comments
 (0)