Repository navigation
Expand file tree
/
Copy pathdatabase.py
More file actions
512 lines (401 loc) · 15.5 KB
/
Copy pathdatabase.py
File metadata and controls
512 lines (401 loc) · 15.5 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
"""
MariaDB connection module for Wikimedia Toolforge.
Provides a lightweight connection pool with retry logic, health checks,
and proper error handling for the Toolforge environment.
Toolforge ToolsDB enforces max_user_connections=20 shared across ALL
processes for the same DB user account. Default pool_size=2 keeps the
web process footprint small (2 gunicorn workers × 2 = 4 conns), leaving
room for background workers.
"""
import logging
import os
import time
from collections.abc import Callable
from contextlib import contextmanager
from queue import Empty, Full, Queue
from typing import Any, Literal, overload
import pymysql
from pymysql import InterfaceError, OperationalError
from pymysql.cursors import DictCursor
logger = logging.getLogger(__name__)
# Connection pool configuration
_pool: Queue | None = None
_pool_size: int = int(os.environ.get("WIKIVISAGE_DB_POOL_SIZE", 2))
_db_config: dict[str, Any] = {}
# Connections older than this are proactively replaced to avoid Toolforge
# killing idle connections unexpectedly (ToolsDB drops idle conns ~300s).
MAX_CONNECTION_AGE = 270 # seconds — below the 300s server-side timeout
# Retry configuration
MAX_RETRIES = 3
INITIAL_BACKOFF = 1.0 # seconds
class DatabaseError(Exception):
"""Base exception for database-related errors."""
pass
class PoolExhaustedError(DatabaseError):
"""Raised when connection pool is exhausted and timeout is reached."""
pass
class ConfigurationError(DatabaseError):
"""Raised when database configuration is invalid or missing."""
pass
def _get_db_config() -> dict[str, Any]:
"""
Retrieve database configuration from environment variables.
Returns:
Dict containing database connection parameters.
Raises:
ConfigurationError: If required environment variables are missing.
"""
user = os.environ.get("TOOL_TOOLSDB_USER")
password = os.environ.get("TOOL_TOOLSDB_PASSWORD")
database = os.environ.get("WIKIVISAGE_DB_NAME")
if not user:
raise ConfigurationError("Missing required environment variable: TOOL_TOOLSDB_USER")
if not password:
raise ConfigurationError("Missing required environment variable: TOOL_TOOLSDB_PASSWORD")
if not database:
raise ConfigurationError("Missing required environment variable: WIKIVISAGE_DB_NAME")
host = os.environ.get("TOOL_TOOLSDB_HOST", "tools.db.svc.wikimedia.cloud")
return {
"host": host,
"user": user,
"password": password,
"database": database,
"charset": "utf8mb4",
"connect_timeout": 10,
"read_timeout": 30,
"autocommit": False,
"cursorclass": DictCursor,
}
class _PooledConnection:
"""Wraps a pymysql connection with creation timestamp for age-based eviction."""
__slots__ = ("conn", "created_at")
def __init__(self, conn: pymysql.Connection):
self.conn = conn
self.created_at: float = time.monotonic()
@property
def is_expired(self) -> bool:
return (time.monotonic() - self.created_at) >= MAX_CONNECTION_AGE
def _create_connection() -> _PooledConnection:
try:
conn = pymysql.connect(**_db_config)
logger.debug("Created new database connection")
return _PooledConnection(conn)
except Exception as e:
logger.error(f"Failed to create database connection: {e}")
raise DatabaseError(f"Could not connect to database: {e}") from e
def _close_quietly(pc: _PooledConnection) -> None:
try:
pc.conn.close()
except Exception:
pass
def _is_connection_healthy(pc: _PooledConnection) -> bool:
if pc.is_expired:
return False
try:
pc.conn.ping(reconnect=False)
# Reset age after successful ping — avoids mass-expiry when
# connections created in the same batch all hit MAX_CONNECTION_AGE
# simultaneously during the next poll cycle.
pc.created_at = time.monotonic()
return True
except Exception:
return False
def _get_connection_from_pool(timeout: float = 30.0) -> _PooledConnection:
if _pool is None:
raise DatabaseError("Connection pool not initialized. Call init_db() first.")
try:
pc = _pool.get(timeout=timeout)
except Empty:
raise PoolExhaustedError(
f"Connection pool exhausted after {timeout}s timeout. Consider increasing pool size or reducing query time."
)
if _is_connection_healthy(pc):
return pc
# Connection is dead or expired — close it and create a fresh one.
# If creation fails, the pool shrinks by one slot. This is intentional:
# re-pooling a dead connection just poisons the next caller.
_close_quietly(pc)
logger.warning("Evicted dead/expired connection from pool, creating replacement")
return _create_connection()
def _try_replenish_pool() -> None:
"""Best-effort: replace a discarded connection with a fresh one.
Called after closing a bad or expired connection so the pool does not
permanently shrink under transient failures. Failures here are logged
and silently swallowed — the pool will recover on the next successful
return or on the next caller that creates a fresh connection at GET time.
"""
if _pool is None:
return
try:
fresh = _create_connection()
try:
_pool.put_nowait(fresh)
logger.debug("Replenished pool with fresh replacement connection")
except Full:
_close_quietly(fresh)
except Exception as e:
logger.warning(f"Could not replenish pool after discarding connection: {e}")
def _return_connection_to_pool(pc: _PooledConnection) -> None:
if _pool is None:
_close_quietly(pc)
return
try:
if pc.conn.open and not pc.is_expired:
try:
pc.conn.rollback()
except Exception:
_close_quietly(pc)
_try_replenish_pool()
return
try:
_pool.put_nowait(pc)
except Full:
_close_quietly(pc)
else:
_close_quietly(pc)
_try_replenish_pool()
except Exception:
_close_quietly(pc)
_try_replenish_pool()
def _execute_with_retry(func: Callable[..., Any], *args, allow_retry: bool = True, **kwargs) -> Any:
"""
Execute a function with exponential backoff retry logic.
Args:
func: The function to execute.
*args: Positional arguments to pass to func.
allow_retry: If False, execute once without retrying (for write operations
where retrying could cause duplicate inserts).
**kwargs: Keyword arguments to pass to func.
Returns:
The result of func.
Raises:
DatabaseError: If all retries fail.
"""
max_attempts = MAX_RETRIES if allow_retry else 1
last_exception: Exception | None = None
for attempt in range(max_attempts):
try:
return func(*args, **kwargs)
except (OperationalError, InterfaceError, PoolExhaustedError) as e:
last_exception = e
if attempt < max_attempts - 1:
backoff = INITIAL_BACKOFF * (2**attempt)
logger.warning(
f"Database operation failed (attempt {attempt + 1}/{max_attempts}): {e}. Retrying in {backoff}s..."
)
time.sleep(backoff)
else:
logger.error(f"Database operation failed after {max_attempts} attempts: {e}")
if last_exception is not None:
raise DatabaseError(f"Database operation failed: {last_exception}") from last_exception
raise DatabaseError("Retry logic failed without capturing an exception")
@contextmanager
def get_connection(timeout: float = 30.0):
"""
Context manager to get a database connection from the pool.
Automatically returns the connection to the pool after use.
Handles rollback on exceptions.
Args:
timeout: Maximum time to wait for a connection (seconds).
Yields:
A database connection.
Example:
with get_connection() as conn:
cursor = conn.cursor()
cursor.execute("SELECT * FROM users")
results = cursor.fetchall()
"""
pc: _PooledConnection | None = None
try:
result = _execute_with_retry(_get_connection_from_pool, timeout)
pc = result # type: _PooledConnection
yield pc.conn
except Exception:
if pc and pc.conn.open:
try:
pc.conn.rollback()
except Exception:
_close_quietly(pc)
pc = None
raise
finally:
if pc:
_return_connection_to_pool(pc)
@overload
def execute_query(
sql: str, params: tuple | dict | None = None, fetch: Literal[True] = True
) -> list[dict[str, Any]]: ...
@overload
def execute_query(sql: str, params: tuple | dict | None = None, *, fetch: Literal[False]) -> int: ...
def execute_query(
sql: str, params: tuple | dict | None = None, fetch: bool = True
) -> list[dict[str, Any]] | int | None:
"""
Execute a SQL query with automatic connection and cursor management.
Args:
sql: The SQL query to execute.
params: Parameters to bind to the query (tuple or dict).
fetch: If True, fetch and return results. If False, return affected row count.
Returns:
If fetch=True: List of result rows as dictionaries.
If fetch=False: Number of affected rows.
Raises:
DatabaseError: If query execution fails.
Example:
# SELECT query
users = execute_query("SELECT * FROM users WHERE id = %s", (user_id,))
# INSERT/UPDATE query
affected = execute_query(
"UPDATE users SET name = %s WHERE id = %s",
("Alice", 123),
fetch=False
)
"""
def _execute() -> list[dict[str, Any]] | int | None:
with get_connection() as conn, conn.cursor() as cursor:
cursor.execute(sql, params)
if fetch:
results = list(cursor.fetchall())
logger.debug(f"Query returned {len(results)} rows")
return results
else:
conn.commit()
affected = cursor.rowcount
logger.debug(f"Query affected {affected} rows")
return affected
try:
return _execute_with_retry(_execute)
except Exception as e:
logger.error(f"Query execution failed: {sql[:100]}... Error: {e}")
raise DatabaseError(f"Query execution failed: {e}") from e
def execute_insert(sql: str, params: tuple | dict | None = None) -> int:
"""
Execute an INSERT query and return the auto-generated row ID.
Uses cursor.lastrowid which is connection-local and race-free,
unlike SELECT ... ORDER BY id DESC LIMIT 1.
Args:
sql: The INSERT SQL query to execute.
params: Parameters to bind to the query.
Returns:
The auto-increment ID of the inserted row.
Raises:
DatabaseError: If query execution fails.
"""
def _execute() -> int:
with get_connection() as conn, conn.cursor() as cursor:
cursor.execute(sql, params)
conn.commit()
return cursor.lastrowid
try:
return _execute_with_retry(_execute, allow_retry=False)
except Exception as e:
logger.error(f"Insert execution failed: {sql[:100]}... Error: {e}")
raise DatabaseError(f"Insert execution failed: {e}") from e
def execute_transaction(
operations: Callable[[Any, Any], Any],
) -> Any:
"""
Execute multiple queries in a single database transaction.
The callable receives (connection, cursor) and should execute all
queries on that cursor. The transaction is committed on success
or rolled back on failure.
Args:
operations: A callable(conn, cursor) that performs all DB work.
Returns:
Whatever the callable returns.
Raises:
DatabaseError: If the transaction fails.
Example:
def do_work(conn, cursor):
cursor.execute("INSERT INTO ...", (...,))
new_id = cursor.lastrowid
cursor.execute("UPDATE ...", (...,))
return new_id
result = execute_transaction(do_work)
"""
def _execute() -> Any:
with get_connection() as conn, conn.cursor() as cursor:
result = operations(conn, cursor)
conn.commit()
return result
try:
return _execute_with_retry(_execute, allow_retry=False)
except Exception as e:
logger.error(f"Transaction execution failed: {e}")
raise DatabaseError(f"Transaction execution failed: {e}") from e
def init_db(pool_size: int | None = None) -> None:
"""
Initialize the database connection pool and verify connectivity.
This must be called before using any database functions.
Args:
pool_size: Number of connections to maintain in the pool.
Defaults to WIKIVISAGE_DB_POOL_SIZE env var or 2.
Raises:
ConfigurationError: If database configuration is invalid.
DatabaseError: If initial connection test fails.
Example:
init_db(pool_size=10)
"""
global _pool, _pool_size, _db_config
if pool_size is not None:
_pool_size = pool_size
logger.info(f"Initializing database connection pool (size={_pool_size})")
# Get and validate configuration
_db_config = _get_db_config()
# Read db_name directly from env to avoid CodeQL taint from the config dict
# that also contains password fields (py/clear-text-logging-sensitive-data).
db_name = os.environ.get("WIKIVISAGE_DB_NAME", "")
# Create the pool
_pool = Queue(maxsize=_pool_size)
# Pre-populate with connections
for i in range(_pool_size):
try:
pc = _create_connection()
_pool.put_nowait(pc)
logger.debug(f"Created connection {i + 1}/{_pool_size}")
except Exception as e:
logger.error(f"Failed to create initial connection {i + 1}/{_pool_size}: {e}")
# Clean up any connections created so far
close_pool()
raise DatabaseError(f"Failed to initialize connection pool: {e}") from e
# Test connectivity
try:
with get_connection() as conn, conn.cursor() as cursor:
cursor.execute("SELECT 1")
result = cursor.fetchone()
if result:
logger.info(
f"Database connection pool initialized successfully. Pool size: {_pool_size}, Database: {db_name}"
)
else:
raise DatabaseError("Connectivity test failed: No result returned")
except Exception as e:
logger.error(f"Database connectivity test failed: {e}")
close_pool()
raise DatabaseError(f"Database connectivity test failed: {e}") from e
def close_pool() -> None:
"""
Close all connections in the pool and clean up resources.
Should be called during application shutdown.
Example:
try:
# Application code
pass
finally:
close_pool()
"""
global _pool
if _pool is None:
logger.debug("Connection pool already closed or not initialized")
return
logger.info("Closing database connection pool")
closed_count = 0
while not _pool.empty():
try:
pc = _pool.get_nowait()
_close_quietly(pc)
closed_count += 1
except Empty:
break
_pool = None
logger.info(f"Closed {closed_count} database connections")