Serialize SQLite write connections
This commit is contained in:
@@ -13,6 +13,8 @@ from urllib.parse import urlparse
|
|||||||
import requests
|
import requests
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
|
from src.database.sqlite_connection import connect_sqlite
|
||||||
|
|
||||||
|
|
||||||
class DBManager:
|
class DBManager:
|
||||||
_init_lock = threading.Lock()
|
_init_lock = threading.Lock()
|
||||||
@@ -34,10 +36,7 @@ class DBManager:
|
|||||||
return raw
|
return raw
|
||||||
|
|
||||||
def _get_connection(self):
|
def _get_connection(self):
|
||||||
conn = sqlite3.connect(self.db_path, timeout=10)
|
return connect_sqlite(self.db_path)
|
||||||
conn.execute("PRAGMA journal_mode=WAL")
|
|
||||||
conn.execute("PRAGMA busy_timeout=5000")
|
|
||||||
return conn
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _is_sqlite_locked_error(exc: sqlite3.OperationalError) -> bool:
|
def _is_sqlite_locked_error(exc: sqlite3.OperationalError) -> bool:
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ from typing import Any, Dict, List, Optional
|
|||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
from src.database.db_manager import DBManager
|
from src.database.db_manager import DBManager
|
||||||
|
from src.database.sqlite_connection import connect_sqlite
|
||||||
|
|
||||||
STATE_STORAGE_FILE = "file"
|
STATE_STORAGE_FILE = "file"
|
||||||
STATE_STORAGE_DUAL = "dual"
|
STATE_STORAGE_DUAL = "dual"
|
||||||
@@ -60,9 +61,7 @@ class RuntimeStateDB:
|
|||||||
return cls._instance
|
return cls._instance
|
||||||
|
|
||||||
def connect(self) -> sqlite3.Connection:
|
def connect(self) -> sqlite3.Connection:
|
||||||
conn = sqlite3.connect(self.db_path)
|
return connect_sqlite(self.db_path, row_factory=sqlite3.Row)
|
||||||
conn.row_factory = sqlite3.Row
|
|
||||||
return conn
|
|
||||||
|
|
||||||
def _init_tables(self) -> None:
|
def _init_tables(self) -> None:
|
||||||
with self.connect() as conn:
|
with self.connect() as conn:
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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
|
||||||
Reference in New Issue
Block a user