Coverage for src/lexigram/auth/storage/_sql_store.py: 68%

108 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-08-26 00:58 +0800

1"""SQL-backed user store implementation.""" 

2 

3from __future__ import annotations 

4 

5from datetime import datetime 

6from typing import TYPE_CHECKING, Any 

7 

8from lexigram import serialization as json 

9from lexigram.auth.models.user import User, UserCredentials 

10from lexigram.di.decorators import inject 

11from lexigram.logging import get_logger 

12from lexigram.serialization import loads 

13 

14if TYPE_CHECKING: 

15 from lexigram.contracts import DatabaseProviderProtocol 

16 

17logger = get_logger(__name__) 

18 

19 

20@inject 

21class SQLUserStore: 

22 """Database-based user store (SQL-backed implementation) 

23 

24 This implementation uses the `DatabaseProvider` API from `lexigram.sql` to 

25 perform SQL operations directly, removing the hard dependency on 

26 SQLAlchemy. It preserves the original public interface expected by 

27 callers/tests while delegating to the canonical database provider. 

28 

29 The public API mirrors :class:`UserStoreProtocol` defined in ``token_store`` 

30 including optional keyword arguments for backwards compatibility. 

31 """ 

32 

33 def __init__(self, db_provider: DatabaseProviderProtocol): 

34 self.db_provider = db_provider 

35 self._initialized = False 

36 

37 async def _ensure_tables(self) -> None: 

38 """Ensure user table exists (Managed by Alembic migrations).""" 

39 # We rely on Alembic migrations to manage the schema 

40 # This prevents the application from trying to create a conflicting schema 

41 

42 async def _user_from_row(self, row: Any) -> User: 

43 """Convert database row to User object""" 

44 

45 def safe_load(val: Any, default: Any) -> Any: 

46 if val is None: 

47 return default 

48 if isinstance(val, (list, dict)): 

49 return val 

50 try: 

51 return loads(val) 

52 except (ValueError, TypeError, json.JSONDecodeError): 

53 return default 

54 

55 return User( 

56 user_id=str(row.get("user_id")), 

57 name=row.get("name") or row.get("username"), 

58 email=row.get("email"), 

59 is_active=bool(row.get("is_active", True)), 

60 is_verified=bool(row.get("is_verified", False)), 

61 roles=safe_load(row.get("roles"), []), 

62 permissions=safe_load(row.get("permissions"), []), 

63 profile=safe_load(row.get("profile"), {}), 

64 created_at=row.get("created_at"), 

65 updated_at=row.get("updated_at"), 

66 last_login_at=row.get("last_login_at"), 

67 login_count=int(row.get("login_count", 0)), 

68 ) 

69 

70 def _credentials_from_row(self, row: Any) -> UserCredentials: 

71 """Extract credential data from a database row.""" 

72 

73 def safe_load(val: Any, default: Any) -> Any: 

74 if val is None: 

75 return default 

76 if isinstance(val, (list, dict)): 

77 return val 

78 try: 

79 return loads(val) 

80 except (ValueError, TypeError, json.JSONDecodeError): 

81 return default 

82 

83 return UserCredentials( 

84 user_id=str(row.get("user_id")), 

85 hashed_password=row.get("hashed_password"), 

86 previous_hashes=safe_load(row.get("previous_passwords"), []), 

87 ) 

88 

89 async def create_user( 

90 self, 

91 name: str, 

92 email: str, 

93 hashed_password: str, 

94 roles: list[str] | None = None, 

95 permissions: list[str] | None = None, 

96 profile: dict[str, Any] | None = None, 

97 **kwargs: Any, 

98 ) -> User: 

99 """Create a new user using the canonical DatabaseProvider API.""" 

100 import uuid 

101 

102 user_id = str(uuid.uuid4()) 

103 

104 is_verified = bool(kwargs.get("is_verified", False)) 

105 

106 user = User( 

107 user_id=user_id, 

108 name=name, 

109 email=email, 

110 roles=list(roles or []), 

111 permissions=list(permissions or []), 

112 profile=profile or {}, 

113 is_verified=is_verified, 

114 ) 

115 

116 # Prefer provider-level convenience method for portability and consistent return shape 

117 payload = { 

118 "user_id": user_id, 

119 "email": email, 

120 "name": name, 

121 "hashed_password": hashed_password, 

122 "is_admin": False, 

123 "is_active": True, 

124 "is_verified": is_verified, 

125 "roles": roles or [], 

126 "permissions": permissions or [], 

127 "previous_passwords": [], 

128 "profile": profile or {}, 

129 "created_at": datetime.now(), 

130 "updated_at": datetime.now(), 

131 } 

132 

133 await self.db_provider.execute_insert("users", payload) 

134 

135 logger.info("Created user: %s", name) 

136 return user 

137 

138 async def get_user_by_id(self, user_id: str) -> User | None: 

139 """Get user by ID""" 

140 select_sql = "SELECT * FROM users WHERE user_id = ?" 

141 

142 # Use provider-level execute_query for a consistent result shape 

143 result = await self.db_provider.execute_query(select_sql, [user_id]) 

144 rows = result.rows if hasattr(result, "rows") else result 

145 row = ( 

146 rows[0] 

147 if isinstance(rows, list) and rows 

148 else (rows if isinstance(rows, dict) else None) 

149 ) 

150 

151 return await self._user_from_row(row) if row else None 

152 

153 async def get_user_by_email(self, email: str) -> User | None: 

154 """Get user by email""" 

155 select_sql = "SELECT * FROM users WHERE email = ?" 

156 

157 result = await self.db_provider.execute_query(select_sql, [email]) 

158 rows = result.rows if hasattr(result, "rows") else result 

159 row = ( 

160 rows[0] 

161 if isinstance(rows, list) and rows 

162 else (rows if isinstance(rows, dict) else None) 

163 ) 

164 

165 return await self._user_from_row(row) if row else None 

166 

167 async def update_user(self, user: User) -> None: 

168 """Update non-credential user information.""" 

169 payload = { 

170 "name": user.name, 

171 "email": user.email, 

172 "is_active": user.is_active, 

173 "is_verified": user.is_verified, 

174 "roles": user.roles, 

175 "permissions": user.permissions, 

176 "profile": user.profile, 

177 "last_login_at": user.last_login_at, 

178 "login_count": user.login_count, 

179 "updated_at": datetime.now(), 

180 } 

181 

182 await self.db_provider.execute_update( 

183 "users", 

184 payload, 

185 "user_id = ?", 

186 [user.user_id], 

187 ) 

188 

189 logger.info("Updated user: %s", user.name) 

190 

191 async def delete_user(self, user_id: str) -> None: 

192 """Delete a user""" 

193 await self.db_provider.execute_delete("users", "user_id = ?", [user_id]) 

194 

195 logger.info("Deleted user: %s", user_id) 

196 

197 async def list_users(self, skip: int = 0, limit: int = 100) -> list[User]: 

198 """List users with pagination""" 

199 select_sql = "SELECT * FROM users ORDER BY created_at DESC LIMIT ? OFFSET ?" 

200 

201 result = await self.db_provider.execute_query(select_sql, [limit, skip]) 

202 rows = result.rows if hasattr(result, "rows") else result 

203 normalized_rows = ( 

204 rows 

205 if isinstance(rows, list) 

206 else ([rows] if isinstance(rows, dict) else []) 

207 ) 

208 

209 users = [] 

210 for row in normalized_rows: 

211 user = await self._user_from_row(row) 

212 users.append(user) 

213 

214 return users 

215 

216 async def count_users(self) -> int: 

217 """Count total users""" 

218 count_sql = "SELECT COUNT(*) as count FROM users" 

219 

220 result = await self.db_provider.execute_query(count_sql) 

221 rows = result.rows if hasattr(result, "rows") else result 

222 if isinstance(rows, list) and rows: 

223 row = rows[0] 

224 elif isinstance(rows, dict): 

225 row = rows 

226 else: 

227 row = None 

228 

229 count = row.get("count") if row else 0 

230 return int(count or 0) 

231 

232 async def get_credentials(self, user_id: str) -> UserCredentials | None: 

233 """Return credential data for *user_id* from the database.""" 

234 select_sql = ( 

235 "SELECT user_id, hashed_password, previous_passwords " 

236 "FROM users WHERE user_id = ?" 

237 ) 

238 result = await self.db_provider.execute_query(select_sql, [user_id]) 

239 rows = result.rows if hasattr(result, "rows") else result 

240 row = ( 

241 rows[0] 

242 if isinstance(rows, list) and rows 

243 else (rows if isinstance(rows, dict) else None) 

244 ) 

245 if not row: 

246 return None 

247 return self._credentials_from_row(row) 

248 

249 async def update_credentials(self, creds: UserCredentials) -> None: 

250 """Persist updated credential data for ``creds.user_id``.""" 

251 

252 def safe_load(val: Any, default: Any) -> Any: 

253 if val is None: 

254 return default 

255 if isinstance(val, (list, dict)): 

256 return val 

257 try: 

258 return loads(val) 

259 except (ValueError, TypeError, json.JSONDecodeError): 

260 return default 

261 

262 payload = { 

263 "hashed_password": creds.hashed_password, 

264 "previous_passwords": creds.previous_hashes, 

265 "updated_at": datetime.now(), 

266 } 

267 await self.db_provider.execute_update( 

268 "users", 

269 payload, 

270 "user_id = ?", 

271 [creds.user_id], 

272 ) 

273 logger.info("Updated credentials for user: %s", creds.user_id) 

274 

275 

276__all__ = ["SQLUserStore"]