Serialize SQLite write connections

This commit is contained in:
2569718930@qq.com
2026-06-15 22:23:29 +08:00
parent 55cfb202da
commit c0b050d2a7
4 changed files with 428 additions and 7 deletions
+3 -4
View File
@@ -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:
+2 -3
View File
@@ -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:
+328
View File
@@ -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
+95
View File
@@ -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