Coverage for src/lexigram/auth/storage/oauth_identity_store/_sqlalchemy.py: 33%
98 statements
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-26 00:58 +0800
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-26 00:58 +0800
1"""SQLAlchemy-backed OAuth identity store."""
3from __future__ import annotations
5from datetime import datetime
6from typing import TYPE_CHECKING, Any
8from lexigram.logging import get_logger
10logger = get_logger(__name__)
12if TYPE_CHECKING:
13 from lexigram.contracts.data import DatabaseProviderProtocol
16from lexigram.auth.storage.oauth_identity_store._protocol import (
17 OAuthIdentity,
18 OAuthIdentityStore,
19)
20from lexigram.di.decorators import inject
23@inject
24class SQLAlchemyOAuthIdentityStore(OAuthIdentityStore):
25 """Database-backed OAuth identity store"""
27 def __init__(self, db_provider: DatabaseProviderProtocol):
28 self.db_provider = db_provider
29 self._initialized = False
31 async def _ensure_tables(self) -> None:
32 """Ensure oauth_identities table exists."""
33 if self._initialized:
34 return
36 create_sql = """
37 CREATE TABLE IF NOT EXISTS oauth_identities (
38 user_id TEXT NOT NULL,
39 provider TEXT NOT NULL,
40 provider_user_id TEXT NOT NULL,
41 created_at DATETIME,
42 updated_at DATETIME,
43 PRIMARY KEY (provider, provider_user_id),
44 FOREIGN KEY (user_id) REFERENCES users(user_id) ON DELETE CASCADE
45 );
46 CREATE INDEX IF NOT EXISTS idx_oauth_identities_user_id ON oauth_identities(user_id);
47 """
49 async with self.db_provider.scoped_context():
50 conn = await self.db_provider.get_scoped_connection()
51 for statement in create_sql.split(";"):
52 stmt = statement.strip()
53 if stmt:
54 await conn.execute(stmt)
56 self._initialized = True
58 async def _identity_from_row(self, row: Any) -> OAuthIdentity:
59 """Convert database row to OAuthIdentity object"""
60 return OAuthIdentity(
61 user_id=row.get("user_id"),
62 provider=row.get("provider"),
63 provider_user_id=row.get("provider_user_id"),
64 created_at=row.get("created_at"),
65 updated_at=row.get("updated_at"),
66 )
68 async def create_oauth_identity(
69 self,
70 user_id: str,
71 provider: str,
72 provider_user_id: str,
73 ) -> OAuthIdentity:
74 """Create OAuth identity link"""
75 await self._ensure_tables()
77 identity = OAuthIdentity(
78 user_id=user_id,
79 provider=provider,
80 provider_user_id=provider_user_id,
81 created_at=datetime.now(),
82 updated_at=datetime.now(),
83 )
85 insert_sql = """
86 INSERT INTO oauth_identities
87 (user_id, provider, provider_user_id, created_at, updated_at)
88 VALUES (?, ?, ?, ?, ?)
89 """
91 async with self.db_provider.scoped_context():
92 conn = await self.db_provider.get_scoped_connection()
93 await conn.execute(
94 insert_sql,
95 [
96 identity.user_id,
97 identity.provider,
98 identity.provider_user_id,
99 identity.created_at,
100 identity.updated_at,
101 ],
102 )
104 logger.info("Created OAuth identity: user=%s, provider=%s", user_id, provider)
105 return identity
107 async def get_oauth_identity(
108 self,
109 provider: str,
110 provider_user_id: str,
111 ) -> OAuthIdentity | None:
112 """Get OAuth identity by provider and provider user ID"""
113 await self._ensure_tables()
115 select_sql = """
116 SELECT * FROM oauth_identities
117 WHERE provider = ? AND provider_user_id = ?
118 """
120 async with self.db_provider.scoped_context():
121 conn = await self.db_provider.get_scoped_connection()
122 result = await conn.execute(select_sql, [provider, provider_user_id])
123 row = result.rows[0] if result.rows else None
125 return await self._identity_from_row(row) if row else None
127 async def get_oauth_identities_for_user(
128 self,
129 user_id: str,
130 ) -> list[OAuthIdentity]:
131 """Get all OAuth identities for a user"""
132 await self._ensure_tables()
134 select_sql = "SELECT * FROM oauth_identities WHERE user_id = ?"
136 async with self.db_provider.scoped_context():
137 conn = await self.db_provider.get_scoped_connection()
138 result = await conn.execute(select_sql, [user_id])
139 rows = result.rows
141 identities = []
142 for row in rows:
143 identity = await self._identity_from_row(row)
144 identities.append(identity)
146 return identities
148 async def delete_oauth_identity(
149 self,
150 provider: str,
151 provider_user_id: str,
152 ) -> bool:
153 """Delete OAuth identity"""
154 await self._ensure_tables()
156 delete_sql = """
157 DELETE FROM oauth_identities
158 WHERE provider = ? AND provider_user_id = ?
159 """
161 async with self.db_provider.scoped_context():
162 conn = await self.db_provider.get_scoped_connection()
163 result = await conn.execute(delete_sql, [provider, provider_user_id])
164 deleted = bool(result.row_count > 0) # Coerce to bool for typing stability
166 if deleted:
167 logger.info(
168 "Deleted OAuth identity: provider=%s, user_id=%s",
169 provider,
170 provider_user_id,
171 )
173 return deleted
175 async def delete_oauth_identities_for_user(self, user_id: str) -> int:
176 """Delete all OAuth identities for a user"""
177 await self._ensure_tables()
179 delete_sql = "DELETE FROM oauth_identities WHERE user_id = ?"
181 async with self.db_provider.scoped_context():
182 conn = await self.db_provider.get_scoped_connection()
183 result = await conn.execute(delete_sql, [user_id])
184 deleted_count = int(result.row_count)
186 if deleted_count > 0:
187 logger.info(
188 "Deleted %d OAuth identities for user %s",
189 deleted_count,
190 user_id,
191 )
193 return int(deleted_count)
195 async def get_user_by_oauth_identity(
196 self,
197 provider: str,
198 provider_user_id: str,
199 ) -> str | None:
200 """Get local user_id by OAuth provider and external user ID."""
201 await self._ensure_tables()
203 select_sql = """
204 SELECT user_id FROM oauth_identities
205 WHERE provider = ? AND provider_user_id = ?
206 """
208 async with self.db_provider.scoped_context():
209 conn = await self.db_provider.get_scoped_connection()
210 result = await conn.execute(select_sql, [provider, provider_user_id])
211 row = result.rows[0] if result.rows else None
213 return row.get("user_id") if row else None
215 async def resolve_user_id(
216 self,
217 user_id_or_oauth_id: str,
218 provider: str = "google",
219 ) -> str | None:
220 """Resolve user_id from either UUID or OAuth external ID.
222 Resolution logic:
223 1. If user_id_or_oauth_id is a valid UUID format, check if user exists
224 2. If not a valid UUID, treat it as an OAuth provider_user_id and look up
225 """
226 import re
228 # Check if it's a valid UUID format
229 uuid_pattern = re.compile(
230 r"^[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}$",
231 re.IGNORECASE,
232 )
234 is_uuid = bool(uuid_pattern.match(user_id_or_oauth_id))
236 if is_uuid:
237 # It's a UUID - check if user exists
238 await self._ensure_tables()
239 check_user_sql = "SELECT user_id FROM users WHERE user_id = ?"
241 async with self.db_provider.scoped_context():
242 conn = await self.db_provider.get_scoped_connection()
243 result = await conn.execute(check_user_sql, [user_id_or_oauth_id])
244 row = result.rows[0] if result.rows else None
246 return row.get("user_id") if row else None
247 # Not a UUID - treat as OAuth external ID
248 return await self.get_user_by_oauth_identity(provider, user_id_or_oauth_id)
250 def resolve_user_id_sync(
251 self,
252 external_id: str,
253 provider: str = "google",
254 ) -> str | None:
255 """Synchronous resolution is not supported for database-backed store.
257 This method is required by IdentityResolverProtocol but cannot be
258 implemented safely for an async database provider.
259 """
260 logger.warning(
261 "sync_resolution_not_supported",
262 provider=provider,
263 external_id=external_id,
264 store="SQLAlchemyOAuthIdentityStore",
265 )
266 return None