From c0b050d2a72f4ee425d8fbb77e02016f0f90e848 Mon Sep 17 00:00:00 2001 From: "2569718930@qq.com" <2569718930@qq.com> Date: Mon, 15 Jun 2026 22:23:29 +0800 Subject: [PATCH] Serialize SQLite write connections --- src/database/db_manager.py | 7 +- src/database/runtime_state.py | 5 +- src/database/sqlite_connection.py | 328 +++++++++++++++++++++++++++ tests/test_sqlite_connection_lock.py | 95 ++++++++ 4 files changed, 428 insertions(+), 7 deletions(-) create mode 100644 src/database/sqlite_connection.py create mode 100644 tests/test_sqlite_connection_lock.py diff --git a/src/database/db_manager.py b/src/database/db_manager.py index 280d0fe9..f34a4f80 100644 --- a/src/database/db_manager.py +++ b/src/database/db_manager.py @@ -13,6 +13,8 @@ from urllib.parse import urlparse import requests from loguru import logger +from src.database.sqlite_connection import connect_sqlite + class DBManager: _init_lock = threading.Lock() @@ -34,10 +36,7 @@ class DBManager: return raw def _get_connection(self): - conn = sqlite3.connect(self.db_path, timeout=10) - conn.execute("PRAGMA journal_mode=WAL") - conn.execute("PRAGMA busy_timeout=5000") - return conn + return connect_sqlite(self.db_path) @staticmethod def _is_sqlite_locked_error(exc: sqlite3.OperationalError) -> bool: diff --git a/src/database/runtime_state.py b/src/database/runtime_state.py index f206c3ba..c6edcccd 100644 --- a/src/database/runtime_state.py +++ b/src/database/runtime_state.py @@ -13,6 +13,7 @@ from typing import Any, Dict, List, Optional from loguru import logger from src.database.db_manager import DBManager +from src.database.sqlite_connection import connect_sqlite STATE_STORAGE_FILE = "file" STATE_STORAGE_DUAL = "dual" @@ -60,9 +61,7 @@ class RuntimeStateDB: return cls._instance def connect(self) -> sqlite3.Connection: - conn = sqlite3.connect(self.db_path) - conn.row_factory = sqlite3.Row - return conn + return connect_sqlite(self.db_path, row_factory=sqlite3.Row) def _init_tables(self) -> None: with self.connect() as conn: diff --git a/src/database/sqlite_connection.py b/src/database/sqlite_connection.py new file mode 100644 index 00000000..14bc3496 --- /dev/null +++ b/src/database/sqlite_connection.py @@ -0,0 +1,328 @@ +from __future__ import annotations + +import os +import sqlite3 +import threading +import time +from contextlib import contextmanager +from typing import Any, Dict, Iterator, Optional + +_PROCESS_LOCKS: Dict[str, threading.RLock] = {} +_PROCESS_LOCKS_GUARD = threading.Lock() +_WAL_INITIALIZED: set[str] = set() +_WAL_INITIALIZED_GUARD = threading.Lock() + + +def _env_float(name: str, default: float) -> float: + raw = str(os.getenv(name, "") or "").strip() + if not raw: + return default + try: + return float(raw) + except ValueError: + return default + + +def _env_int(name: str, default: int) -> int: + raw = str(os.getenv(name, "") or "").strip() + if not raw: + return default + try: + return int(float(raw)) + except ValueError: + return default + + +def _write_lock_enabled() -> bool: + raw = str(os.getenv("POLYWEATHER_SQLITE_WRITE_LOCK_ENABLED", "1") or "").strip().lower() + return raw not in {"0", "false", "no", "off"} + + +def _default_sqlite_timeout_sec() -> float: + return max(1.0, _env_float("POLYWEATHER_SQLITE_TIMEOUT_SEC", 30.0)) + + +def _default_busy_timeout_ms() -> int: + return max(1000, _env_int("POLYWEATHER_SQLITE_BUSY_TIMEOUT_MS", 30000)) + + +def _default_write_lock_timeout_sec() -> float: + return max(1.0, _env_float("POLYWEATHER_SQLITE_WRITE_LOCK_TIMEOUT_SEC", 30.0)) + + +def sqlite_write_lock_path(db_path: str) -> str: + override = str(os.getenv("POLYWEATHER_SQLITE_WRITE_LOCK_PATH") or "").strip() + if override: + return os.path.abspath(override) + return f"{os.path.abspath(os.fspath(db_path))}.write.lock" + + +def _process_lock(path: str) -> threading.RLock: + with _PROCESS_LOCKS_GUARD: + lock = _PROCESS_LOCKS.get(path) + if lock is None: + lock = threading.RLock() + _PROCESS_LOCKS[path] = lock + return lock + + +def _ensure_lock_file_ready(fd: int, path: str) -> None: + try: + if os.path.getsize(path) <= 0: + os.write(fd, b"\0") + except OSError: + pass + os.lseek(fd, 0, os.SEEK_SET) + + +if os.name == "nt": + import msvcrt + + def _try_lock_fd(fd: int) -> None: + os.lseek(fd, 0, os.SEEK_SET) + try: + msvcrt.locking(fd, msvcrt.LK_NBLCK, 1) + except OSError as exc: + raise BlockingIOError(str(exc)) from exc + + def _unlock_fd(fd: int) -> None: + os.lseek(fd, 0, os.SEEK_SET) + msvcrt.locking(fd, msvcrt.LK_UNLCK, 1) + +else: + import fcntl + + def _try_lock_fd(fd: int) -> None: + fcntl.flock(fd, fcntl.LOCK_EX | fcntl.LOCK_NB) + + def _unlock_fd(fd: int) -> None: + fcntl.flock(fd, fcntl.LOCK_UN) + + +@contextmanager +def sqlite_write_lock( + db_path: str, + *, + timeout_sec: Optional[float] = None, +) -> Iterator[None]: + if not _write_lock_enabled(): + yield + return + + lock_path = sqlite_write_lock_path(db_path) + lock_dir = os.path.dirname(lock_path) + if lock_dir: + os.makedirs(lock_dir, exist_ok=True) + + timeout = _default_write_lock_timeout_sec() if timeout_sec is None else max(0.0, float(timeout_sec)) + deadline = time.monotonic() + timeout + process_lock = _process_lock(lock_path) + remaining = max(0.0, deadline - time.monotonic()) + acquired_thread_lock = process_lock.acquire(timeout=remaining) + if not acquired_thread_lock: + raise TimeoutError(f"timed out waiting for sqlite write lock: {lock_path}") + + fd = -1 + acquired_file_lock = False + try: + fd = os.open(lock_path, os.O_CREAT | os.O_RDWR, 0o666) + _ensure_lock_file_ready(fd, lock_path) + while True: + try: + _try_lock_fd(fd) + acquired_file_lock = True + break + except BlockingIOError: + if time.monotonic() >= deadline: + raise TimeoutError(f"timed out waiting for sqlite write lock: {lock_path}") + time.sleep(0.025) + yield + finally: + try: + if acquired_file_lock and fd >= 0: + _unlock_fd(fd) + finally: + if fd >= 0: + os.close(fd) + process_lock.release() + + +_WRITE_PREFIXES = { + "alter", + "attach", + "begin", + "create", + "delete", + "detach", + "drop", + "insert", + "pragma", + "reindex", + "replace", + "update", + "vacuum", +} + + +def _strip_sql_comments(sql: str) -> str: + text = str(sql or "").lstrip() + while text.startswith("--") or text.startswith("/*"): + if text.startswith("--"): + _, _, text = text.partition("\n") + text = text.lstrip() + continue + end = text.find("*/") + if end < 0: + return "" + text = text[end + 2 :].lstrip() + return text + + +def _statement_needs_write_lock(sql: str) -> bool: + text = _strip_sql_comments(sql) + if not text: + return False + first = text.split(None, 1)[0].strip().lower() + if first not in _WRITE_PREFIXES: + return False + if first == "pragma": + lower = text.lower() + return "journal_mode" in lower or "wal_checkpoint" in lower + if first == "begin": + lower = text.lower() + return "immediate" in lower or "exclusive" in lower + return True + + +def _script_needs_write_lock(sql_script: str) -> bool: + return any(_statement_needs_write_lock(statement) for statement in str(sql_script or "").split(";")) + + +class LockedSQLiteConnection(sqlite3.Connection): + def __init__(self, *args: Any, **kwargs: Any) -> None: + super().__init__(*args, **kwargs) + self._polyweather_db_path = os.fspath(args[0]) if args else "" + self._polyweather_write_lock_timeout_sec = _default_write_lock_timeout_sec() + self._polyweather_write_lock_cm = None + + def configure_polyweather_lock( + self, + *, + db_path: str, + write_lock_timeout_sec: Optional[float] = None, + ) -> "LockedSQLiteConnection": + self._polyweather_db_path = os.fspath(db_path) + if write_lock_timeout_sec is not None: + self._polyweather_write_lock_timeout_sec = max(0.0, float(write_lock_timeout_sec)) + return self + + def _ensure_write_lock(self) -> None: + if self._polyweather_write_lock_cm is not None: + return + cm = sqlite_write_lock( + self._polyweather_db_path, + timeout_sec=self._polyweather_write_lock_timeout_sec, + ) + cm.__enter__() + self._polyweather_write_lock_cm = cm + + def _release_write_lock(self) -> None: + cm = self._polyweather_write_lock_cm + if cm is None: + return + self._polyweather_write_lock_cm = None + cm.__exit__(None, None, None) + + def execute(self, sql: str, parameters: Any = (), /) -> sqlite3.Cursor: + if _statement_needs_write_lock(sql): + self._ensure_write_lock() + return super().execute(sql, parameters) + + def executemany(self, sql: str, seq_of_parameters: Any, /) -> sqlite3.Cursor: + if _statement_needs_write_lock(sql): + self._ensure_write_lock() + return super().executemany(sql, seq_of_parameters) + + def executescript(self, sql_script: str, /) -> sqlite3.Cursor: + if _script_needs_write_lock(sql_script): + self._ensure_write_lock() + return super().executescript(sql_script) + + def commit(self) -> None: + try: + super().commit() + finally: + self._release_write_lock() + + def rollback(self) -> None: + try: + super().rollback() + finally: + self._release_write_lock() + + def close(self) -> None: + try: + super().close() + finally: + self._release_write_lock() + + def __exit__(self, exc_type: Any, exc_value: Any, traceback: Any) -> Optional[bool]: + try: + return super().__exit__(exc_type, exc_value, traceback) + finally: + self._release_write_lock() + + +def _wal_cache_key(db_path: str) -> str: + return os.path.abspath(os.fspath(db_path)) + + +def _ensure_wal_mode(conn: LockedSQLiteConnection, db_path: str, timeout_sec: float) -> None: + if os.fspath(db_path) == ":memory:": + return + cache_key = _wal_cache_key(db_path) + if cache_key in _WAL_INITIALIZED: + return + with _WAL_INITIALIZED_GUARD: + if cache_key in _WAL_INITIALIZED: + return + with sqlite_write_lock(db_path, timeout_sec=timeout_sec): + sqlite3.Connection.execute(conn, "PRAGMA journal_mode=WAL") + _WAL_INITIALIZED.add(cache_key) + + +def connect_sqlite( + db_path: str, + *, + timeout: Optional[float] = None, + busy_timeout_ms: Optional[int] = None, + enable_wal: bool = True, + row_factory: Optional[Any] = None, + write_lock_timeout_sec: Optional[float] = None, + **kwargs: Any, +) -> LockedSQLiteConnection: + sqlite_timeout = _default_sqlite_timeout_sec() if timeout is None else max(0.0, float(timeout)) + lock_timeout = ( + _default_write_lock_timeout_sec() + if write_lock_timeout_sec is None + else max(0.0, float(write_lock_timeout_sec)) + ) + conn = sqlite3.connect( + db_path, + timeout=sqlite_timeout, + factory=LockedSQLiteConnection, + **kwargs, + ) + conn.configure_polyweather_lock( + db_path=db_path, + write_lock_timeout_sec=lock_timeout, + ) + sqlite3.Connection.execute( + conn, + f"PRAGMA busy_timeout={_default_busy_timeout_ms() if busy_timeout_ms is None else max(0, int(busy_timeout_ms))}", + ) + if enable_wal: + _ensure_wal_mode(conn, db_path, lock_timeout) + if row_factory is not None: + conn.row_factory = row_factory + return conn diff --git a/tests/test_sqlite_connection_lock.py b/tests/test_sqlite_connection_lock.py new file mode 100644 index 00000000..cc1c397c --- /dev/null +++ b/tests/test_sqlite_connection_lock.py @@ -0,0 +1,95 @@ +import sqlite3 +import threading +import time + +from src.database.db_manager import DBManager +from src.database.runtime_state import RuntimeStateDB +from src.database.sqlite_connection import ( + LockedSQLiteConnection, + connect_sqlite, + sqlite_write_lock, +) + + +def test_sqlite_write_lock_times_out_when_another_thread_holds_it(tmp_path): + db_path = str(tmp_path / "polyweather.db") + errors = [] + + def contender(): + try: + with sqlite_write_lock(db_path, timeout_sec=0.05): + pass + except TimeoutError as exc: + errors.append(exc) + + with sqlite_write_lock(db_path, timeout_sec=1.0): + started = time.monotonic() + thread = threading.Thread(target=contender) + thread.start() + thread.join(timeout=1.0) + elapsed = time.monotonic() - started + + assert not thread.is_alive() + assert errors + assert elapsed >= 0.05 + + +def test_locked_sqlite_connection_serializes_concurrent_writes(tmp_path): + db_path = str(tmp_path / "polyweather.db") + with connect_sqlite(db_path) as conn: + conn.execute("CREATE TABLE items (id INTEGER PRIMARY KEY, value TEXT NOT NULL)") + conn.commit() + + first_inserted = threading.Event() + release_first = threading.Event() + order = [] + + def first_writer(): + with connect_sqlite(db_path, write_lock_timeout_sec=2.0) as conn: + conn.execute("INSERT INTO items(value) VALUES ('first')") + order.append("first-inserted") + first_inserted.set() + assert release_first.wait(timeout=2.0) + conn.commit() + order.append("first-committed") + + def second_writer(): + assert first_inserted.wait(timeout=2.0) + with connect_sqlite(db_path, write_lock_timeout_sec=2.0) as conn: + conn.execute("INSERT INTO items(value) VALUES ('second')") + conn.commit() + order.append("second-committed") + + first = threading.Thread(target=first_writer) + second = threading.Thread(target=second_writer) + first.start() + second.start() + + assert first_inserted.wait(timeout=2.0) + time.sleep(0.1) + assert order == ["first-inserted"] + + release_first.set() + first.join(timeout=2.0) + second.join(timeout=2.0) + + assert not first.is_alive() + assert not second.is_alive() + assert order == ["first-inserted", "first-committed", "second-committed"] + + with connect_sqlite(db_path) as conn: + rows = conn.execute("SELECT value FROM items ORDER BY id ASC").fetchall() + assert [row[0] for row in rows] == ["first", "second"] + + +def test_database_managers_use_locked_sqlite_connections(tmp_path): + manager = DBManager(str(tmp_path / "app.db")) + with manager._get_connection() as conn: + assert isinstance(conn, LockedSQLiteConnection) + assert conn.execute("PRAGMA busy_timeout").fetchone()[0] >= 30000 + + state_db = RuntimeStateDB(str(tmp_path / "state.db")) + with state_db.connect() as conn: + assert isinstance(conn, LockedSQLiteConnection) + assert conn.row_factory is sqlite3.Row + assert conn.execute("PRAGMA busy_timeout").fetchone()[0] >= 30000