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

1"""SQLAlchemy-backed OAuth identity store.""" 

2 

3from __future__ import annotations 

4 

5from datetime import datetime 

6from typing import TYPE_CHECKING, Any 

7 

8from lexigram.logging import get_logger 

9 

10logger = get_logger(__name__) 

11 

12if TYPE_CHECKING: 

13 from lexigram.contracts.data import DatabaseProviderProtocol 

14 

15 

16from lexigram.auth.storage.oauth_identity_store._protocol import ( 

17 OAuthIdentity, 

18 OAuthIdentityStore, 

19) 

20from lexigram.di.decorators import inject 

21 

22 

23@inject 

24class SQLAlchemyOAuthIdentityStore(OAuthIdentityStore): 

25 """Database-backed OAuth identity store""" 

26 

27 def __init__(self, db_provider: DatabaseProviderProtocol): 

28 self.db_provider = db_provider 

29 self._initialized = False 

30 

31 async def _ensure_tables(self) -> None: 

32 """Ensure oauth_identities table exists.""" 

33 if self._initialized: 

34 return 

35 

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 """ 

48 

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) 

55 

56 self._initialized = True 

57 

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 ) 

67 

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() 

76 

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 ) 

84 

85 insert_sql = """ 

86 INSERT INTO oauth_identities 

87 (user_id, provider, provider_user_id, created_at, updated_at) 

88 VALUES (?, ?, ?, ?, ?) 

89 """ 

90 

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 ) 

103 

104 logger.info("Created OAuth identity: user=%s, provider=%s", user_id, provider) 

105 return identity 

106 

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() 

114 

115 select_sql = """ 

116 SELECT * FROM oauth_identities 

117 WHERE provider = ? AND provider_user_id = ? 

118 """ 

119 

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 

124 

125 return await self._identity_from_row(row) if row else None 

126 

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() 

133 

134 select_sql = "SELECT * FROM oauth_identities WHERE user_id = ?" 

135 

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 

140 

141 identities = [] 

142 for row in rows: 

143 identity = await self._identity_from_row(row) 

144 identities.append(identity) 

145 

146 return identities 

147 

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() 

155 

156 delete_sql = """ 

157 DELETE FROM oauth_identities 

158 WHERE provider = ? AND provider_user_id = ? 

159 """ 

160 

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 

165 

166 if deleted: 

167 logger.info( 

168 "Deleted OAuth identity: provider=%s, user_id=%s", 

169 provider, 

170 provider_user_id, 

171 ) 

172 

173 return deleted 

174 

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() 

178 

179 delete_sql = "DELETE FROM oauth_identities WHERE user_id = ?" 

180 

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) 

185 

186 if deleted_count > 0: 

187 logger.info( 

188 "Deleted %d OAuth identities for user %s", 

189 deleted_count, 

190 user_id, 

191 ) 

192 

193 return int(deleted_count) 

194 

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() 

202 

203 select_sql = """ 

204 SELECT user_id FROM oauth_identities 

205 WHERE provider = ? AND provider_user_id = ? 

206 """ 

207 

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 

212 

213 return row.get("user_id") if row else None 

214 

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. 

221 

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 

227 

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 ) 

233 

234 is_uuid = bool(uuid_pattern.match(user_id_or_oauth_id)) 

235 

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 = ?" 

240 

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 

245 

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) 

249 

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. 

256 

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