diff --git a/.env-sample b/.env-sample index 5aff92d..867ee46 100644 --- a/.env-sample +++ b/.env-sample @@ -35,6 +35,7 @@ DATABASE_TYPE=sqlite # Options: sqlite, postgresql - If you have enabled auto-sc DATABASE_URL=username:password@hostname:port # For PostgreSQL DATABASE_PATH=data/comet.db # Only relevant for SQLite DATABASE_BATCH_SIZE=20000 # The batch size for the database import and export operations +DATABASE_READ_REPLICA_URLS='' # Optional JSON array of PostgreSQL read-only URLs, e.g. '["user:pass@replica-1/db", "user:pass@replica-2/db"]' # ============================== # # Cache Settings (Seconds) # diff --git a/CHANGELOG.md b/CHANGELOG.md index f9f03f1..ee34921 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,13 @@ # Changelog + + +## [Unreleased] + +### Features + +* add optional PostgreSQL read replica routing with transparent primary fallback + ## [2.31.0](https://github.com/g0ldyy/comet/compare/v2.30.0...v2.31.0) (2025-12-08) diff --git a/comet/core/db_router.py b/comet/core/db_router.py new file mode 100644 index 0000000..1dc1d3d --- /dev/null +++ b/comet/core/db_router.py @@ -0,0 +1,154 @@ +import asyncio +import contextvars +from contextlib import contextmanager +from typing import List, Optional, Sequence + +from databases import Database + +from comet.core.logger import logger + + +class ReplicaAwareDatabase: + """Routes read queries to replicas while keeping writes on the primary.""" + + def __init__(self, primary: Database, replicas: Optional[Sequence[Database]] = None): + self._primary = primary + self._configured_replicas = list(replicas or []) + self._active_replicas: List[Database] = [] + self._replica_index = 0 + self._transaction_depth = contextvars.ContextVar( + "comet_db_replica_tx_depth", default=0 + ) + self._force_primary_context = contextvars.ContextVar( + "comet_db_replica_force_primary", default=False + ) + + @property + def has_replicas(self) -> bool: + return bool(self._active_replicas) + + @property + def is_connected(self) -> bool: + return self._primary.is_connected + + async def connect(self): + await self._primary.connect() + + healthy_replicas: List[Database] = [] + for replica in self._configured_replicas: + try: + await replica.connect() + except Exception as exc: # pragma: no cover - defensive logging + logger.log( + "DATABASE", + f"Read replica connection failed ({getattr(replica, 'url', 'replica')}): {exc}", + ) + else: + healthy_replicas.append(replica) + + self._active_replicas = healthy_replicas + + if self._active_replicas: + logger.log( + "DATABASE", + f"Read replicas enabled ({len(self._active_replicas)} healthy)", + ) + + async def disconnect(self): + for db in [self._primary, *self._configured_replicas]: + if not db.is_connected: + continue + try: + await db.disconnect() + except Exception as exc: # pragma: no cover - defensive logging + logger.log("DATABASE", f"Error disconnecting database: {exc}") + + def transaction(self, *args, **kwargs): + primary_transaction = self._primary.transaction(*args, **kwargs) + return _ReplicaAwareTransaction(self, primary_transaction) + + async def execute(self, query, values=None): + return await self._primary.execute(query, values) + + async def execute_many(self, query, values): + return await self._primary.execute_many(query, values) + + async def fetch_all(self, query, values=None, *, force_primary: bool = False): + return await self._run_read("fetch_all", force_primary, query, values) + + async def fetch_one(self, query, values=None, *, force_primary: bool = False): + return await self._run_read("fetch_one", force_primary, query, values) + + async def fetch_val( + self, query, values=None, column: int = 0, *, force_primary: bool = False + ): + return await self._run_read("fetch_val", force_primary, query, values, column) + + def _should_use_primary(self, explicit_force: bool) -> bool: + if explicit_force or self._force_primary_context.get(): + return True + + if not self._active_replicas: + return True + + if self._transaction_depth.get() > 0: + return True + + return False + + def _next_replica(self) -> Database: + replica = self._active_replicas[self._replica_index % len(self._active_replicas)] + self._replica_index = (self._replica_index + 1) % len(self._active_replicas) + return replica + + async def _run_read(self, method_name: str, force_primary: bool, *args): + target = ( + self._primary + if self._should_use_primary(force_primary) + else self._next_replica() + ) + + method = getattr(target, method_name) + try: + return await method(*args) + except asyncio.CancelledError: # pragma: no cover - propagate cancellations + raise + except Exception as exc: + if target is not self._primary and self._primary.is_connected: + logger.log( + "DATABASE", + f"Replica {method_name} failed, retrying on primary: {exc}", + ) + fallback = getattr(self._primary, method_name) + return await fallback(*args) + raise + + @contextmanager + def force_primary(self): + token = self._force_primary_context.set(True) + try: + yield self + finally: + self._force_primary_context.reset(token) + + def __getattr__(self, item): + return getattr(self._primary, item) + + +class _ReplicaAwareTransaction: + def __init__(self, router: ReplicaAwareDatabase, transaction_cm): + self._router = router + self._transaction_cm = transaction_cm + self._token = None + + async def __aenter__(self): + current_depth = self._router._transaction_depth.get() + self._token = self._router._transaction_depth.set(current_depth + 1) + return await self._transaction_cm.__aenter__() + + async def __aexit__(self, exc_type, exc, tb): + try: + return await self._transaction_cm.__aexit__(exc_type, exc, tb) + finally: + if self._token is not None: + self._router._transaction_depth.reset(self._token) diff --git a/comet/core/models.py b/comet/core/models.py index 9036b4b..ef3f7bd 100644 --- a/comet/core/models.py +++ b/comet/core/models.py @@ -4,7 +4,7 @@ from typing import List, Optional, Union import RTN from databases import Database -from pydantic import BaseModel, field_validator +from pydantic import BaseModel, Field, field_validator from pydantic_settings import BaseSettings, SettingsConfigDict from RTN import DefaultRanking, SettingsModel from RTN.models import (AudioRankModel, CustomRank, CustomRanksConfig, @@ -12,6 +12,8 @@ from RTN.models import (AudioRankModel, CustomRank, CustomRanksConfig, OptionsConfig, QualityRankModel, ResolutionConfig, RipsRankModel) +from comet.core.db_router import ReplicaAwareDatabase + class AppSettings(BaseSettings): model_config = SettingsConfigDict( @@ -32,6 +34,7 @@ class AppSettings(BaseSettings): DATABASE_URL: Optional[str] = "username:password@hostname:port" DATABASE_PATH: Optional[str] = "data/comet.db" DATABASE_BATCH_SIZE: Optional[int] = 20000 + DATABASE_READ_REPLICA_URLS: List[str] = Field(default_factory=list) METADATA_CACHE_TTL: Optional[int] = 2592000 # 30 days TORRENT_CACHE_TTL: Optional[int] = 1296000 # 15 days LIVE_TORRENT_CACHE_TTL: Optional[int] = 1296000 # 15 days @@ -673,13 +676,31 @@ web_config = { ], } + +def _build_database_instance(raw_url: str) -> Database: + driver = "sqlite" if settings.DATABASE_TYPE == "sqlite" else "postgresql+asyncpg" + prefix = "/" if settings.DATABASE_TYPE == "sqlite" else "" + return Database(f"{driver}://{prefix}{raw_url}") + + database_url = ( settings.DATABASE_PATH if settings.DATABASE_TYPE == "sqlite" else settings.DATABASE_URL ) -database = Database( - f"{'sqlite' if settings.DATABASE_TYPE == 'sqlite' else 'postgresql+asyncpg'}://{'/' if settings.DATABASE_TYPE == 'sqlite' else ''}{database_url}" + +replica_instances: List[Database] = [] +if settings.DATABASE_TYPE != "sqlite" and settings.DATABASE_READ_REPLICA_URLS: + for replica_url in settings.DATABASE_READ_REPLICA_URLS: + if replica_url: + replica_instances.append(_build_database_instance(replica_url)) +elif settings.DATABASE_TYPE == "sqlite" and settings.DATABASE_READ_REPLICA_URLS: + logger.log( + "DATABASE", "Read replicas are ignored for sqlite deployments" + ) + +database = ReplicaAwareDatabase( + _build_database_instance(database_url), replicas=replica_instances ) trackers = [ diff --git a/comet/services/lock.py b/comet/services/lock.py index 5ab3474..358202b 100644 --- a/comet/services/lock.py +++ b/comet/services/lock.py @@ -54,6 +54,7 @@ class DistributedLock: row = await database.fetch_one( "SELECT instance_id FROM scrape_locks WHERE lock_key = :lock_key", {"lock_key": self.lock_key}, + force_primary=True, ) success = row and row["instance_id"] == self.instance_id @@ -125,6 +126,7 @@ async def is_scrape_in_progress(media_id: str): row = await database.fetch_one( "SELECT instance_id FROM scrape_locks WHERE lock_key = :lock_key", {"lock_key": media_id}, + force_primary=True, ) return row is not None except Exception as e: