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
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-26 00:58 +0800
1"""SQL-backed user store implementation."""
3from __future__ import annotations
5from datetime import datetime
6from typing import TYPE_CHECKING, Any
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
14if TYPE_CHECKING:
15 from lexigram.contracts import DatabaseProviderProtocol
17logger = get_logger(__name__)
20@inject
21class SQLUserStore:
22 """Database-based user store (SQL-backed implementation)
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.
29 The public API mirrors :class:`UserStoreProtocol` defined in ``token_store``
30 including optional keyword arguments for backwards compatibility.
31 """
33 def __init__(self, db_provider: DatabaseProviderProtocol):
34 self.db_provider = db_provider
35 self._initialized = False
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
42 async def _user_from_row(self, row: Any) -> User:
43 """Convert database row to User object"""
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
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 )
70 def _credentials_from_row(self, row: Any) -> UserCredentials:
71 """Extract credential data from a database row."""
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
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 )
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
102 user_id = str(uuid.uuid4())
104 is_verified = bool(kwargs.get("is_verified", False))
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 )
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 }
133 await self.db_provider.execute_insert("users", payload)
135 logger.info("Created user: %s", name)
136 return user
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 = ?"
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 )
151 return await self._user_from_row(row) if row else None
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 = ?"
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 )
165 return await self._user_from_row(row) if row else None
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 }
182 await self.db_provider.execute_update(
183 "users",
184 payload,
185 "user_id = ?",
186 [user.user_id],
187 )
189 logger.info("Updated user: %s", user.name)
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])
195 logger.info("Deleted user: %s", user_id)
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 ?"
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 )
209 users = []
210 for row in normalized_rows:
211 user = await self._user_from_row(row)
212 users.append(user)
214 return users
216 async def count_users(self) -> int:
217 """Count total users"""
218 count_sql = "SELECT COUNT(*) as count FROM users"
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
229 count = row.get("count") if row else 0
230 return int(count or 0)
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)
249 async def update_credentials(self, creds: UserCredentials) -> None:
250 """Persist updated credential data for ``creds.user_id``."""
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
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)
276__all__ = ["SQLUserStore"]