refactor: match provider identity lookups via JSON subscript (#28624)

Both the OAuth and SCIM user lookups now compare the nested JSON value with SQLAlchemy's subscript operator, which emits the correct SQL for each supported database on its own. This replaces the hand-written sqlite and postgresql branches and the column-level contains() call they used.
This commit is contained in:
Classic298
2026-08-17 01:05:16 -06:00
committed by GitHub
parent 90724cdee0
commit 73c1f5806a
+6 -19
View File
@@ -27,7 +27,6 @@ from sqlalchemy import (
select,
update,
)
from sqlalchemy.dialects.postgresql import JSONB
from sqlalchemy.ext.asyncio import AsyncSession
####################
@@ -360,16 +359,10 @@ class UsersTable:
sub: str,
db: AsyncSession | None = None,
) -> UserModel | None:
"""Look up a user by OAuth provider + subject claim (dialect-aware JSON filter)."""
"""Look up a user by OAuth provider + subject claim."""
async with get_async_db_context(db) as session:
dialect = session.bind.dialect.name
query = select(User)
if dialect == 'sqlite':
oauth_match = User.oauth.contains({provider: {'sub': sub}})
query = query.where(oauth_match)
elif dialect == 'postgresql':
oauth_match = User.oauth[provider].cast(JSONB)['sub'].astext == sub
query = query.where(oauth_match)
# Subscript, never contains(): on a JSON column contains() degrades to a substring LIKE.
query = select(User).where(User.oauth[provider]['sub'].as_string() == sub)
row = (await session.execute(query)).scalars().first()
return UserModel.model_validate(row) if row else None
@@ -379,16 +372,10 @@ class UsersTable:
external_id: str,
db: AsyncSession | None = None,
) -> UserModel | None:
"""Look up a user by SCIM provider + external ID (dialect-aware JSON filter)."""
"""Look up a user by SCIM provider + external ID."""
async with get_async_db_context(db) as session:
dialect = session.bind.dialect.name
query = select(User)
if dialect == 'sqlite':
scim_match = User.scim.contains({provider: {'external_id': external_id}})
query = query.where(scim_match)
elif dialect == 'postgresql':
scim_match = User.scim[provider].cast(JSONB)['external_id'].astext == external_id
query = query.where(scim_match)
# Subscript, never contains(): on a JSON column contains() degrades to a substring LIKE.
query = select(User).where(User.scim[provider]['external_id'].as_string() == external_id)
row = (await session.execute(query)).scalars().first()
return UserModel.model_validate(row) if row else None