refactor(db)!: finalize pooled WNDB v2 migration
This commit is contained in:
+250
-181
@@ -1,40 +1,29 @@
|
||||
import asyncio
|
||||
import logging
|
||||
from collections import OrderedDict
|
||||
from collections.abc import AsyncIterator
|
||||
from contextlib import asynccontextmanager
|
||||
from dataclasses import dataclass
|
||||
from typing import Dict
|
||||
from uuid import UUID
|
||||
|
||||
from psycopg import AsyncConnection
|
||||
from psycopg_pool import AsyncConnectionPool
|
||||
from psycopg.rows import dict_row
|
||||
from sqlalchemy.engine.url import make_url
|
||||
from sqlalchemy.ext.asyncio import (
|
||||
AsyncEngine,
|
||||
AsyncSession,
|
||||
async_sessionmaker,
|
||||
create_async_engine,
|
||||
)
|
||||
|
||||
from app.core.config import settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
_check_async_connection = AsyncConnectionPool.check_connection
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class PgEngineEntry:
|
||||
engine: AsyncEngine
|
||||
sessionmaker: async_sessionmaker[AsyncSession]
|
||||
connection_url: str
|
||||
pool_min_size: int
|
||||
pool_max_size: int
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
@dataclass
|
||||
class PoolEntry:
|
||||
pool: AsyncConnectionPool
|
||||
connection_url: str
|
||||
pool_min_size: int
|
||||
pool_max_size: int
|
||||
borrow_count: int = 0
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
@@ -45,186 +34,203 @@ class CacheKey:
|
||||
|
||||
class ProjectConnectionManager:
|
||||
def __init__(self) -> None:
|
||||
self._pg_cache: Dict[CacheKey, PgEngineEntry] = OrderedDict()
|
||||
self._ts_cache: Dict[CacheKey, PoolEntry] = OrderedDict()
|
||||
self._pg_raw_cache: Dict[CacheKey, PoolEntry] = OrderedDict()
|
||||
self._pg_lock = asyncio.Lock()
|
||||
self._retired_ts: list[tuple[CacheKey, PoolEntry]] = []
|
||||
self._retired_pg: list[tuple[CacheKey, PoolEntry]] = []
|
||||
self._ts_lock = asyncio.Lock()
|
||||
self._pg_raw_lock = asyncio.Lock()
|
||||
|
||||
def _normalize_pg_url(self, url: str) -> str:
|
||||
parsed = make_url(url)
|
||||
if parsed.drivername in {"postgresql", "postgres"}:
|
||||
parsed = parsed.set(drivername="postgresql+psycopg")
|
||||
return parsed.render_as_string(hide_password=False)
|
||||
|
||||
async def get_pg_sessionmaker(
|
||||
async def _get_timescale_pool_locked(
|
||||
self,
|
||||
project_id: UUID,
|
||||
db_role: str,
|
||||
connection_url: str,
|
||||
pool_min_size: int,
|
||||
pool_max_size: int,
|
||||
) -> async_sessionmaker[AsyncSession]:
|
||||
async with self._pg_lock:
|
||||
normalized_url = self._normalize_pg_url(connection_url)
|
||||
pool_min_size = max(1, pool_min_size)
|
||||
pool_max_size = max(pool_min_size, pool_max_size)
|
||||
|
||||
key = CacheKey(project_id=project_id, db_role=db_role)
|
||||
entry = self._pg_cache.get(key)
|
||||
if entry:
|
||||
if (
|
||||
entry.connection_url == normalized_url
|
||||
and entry.pool_min_size == pool_min_size
|
||||
and entry.pool_max_size == pool_max_size
|
||||
):
|
||||
self._pg_cache.move_to_end(key)
|
||||
return entry.sessionmaker
|
||||
|
||||
await entry.engine.dispose()
|
||||
logger.info(
|
||||
"Rebuilding PostgreSQL engine for project %s (%s) due to config change",
|
||||
project_id,
|
||||
db_role,
|
||||
)
|
||||
self._pg_cache.pop(key, None)
|
||||
|
||||
engine = create_async_engine(
|
||||
normalized_url,
|
||||
pool_size=pool_min_size,
|
||||
max_overflow=max(0, pool_max_size - pool_min_size),
|
||||
pool_pre_ping=True,
|
||||
)
|
||||
sessionmaker = async_sessionmaker(engine, expire_on_commit=False)
|
||||
self._pg_cache[key] = PgEngineEntry(
|
||||
engine=engine,
|
||||
sessionmaker=sessionmaker,
|
||||
connection_url=normalized_url,
|
||||
pool_min_size=pool_min_size,
|
||||
pool_max_size=pool_max_size,
|
||||
)
|
||||
await self._evict_pg_if_needed()
|
||||
) -> tuple[CacheKey, AsyncConnectionPool]:
|
||||
pool_min_size = max(0, pool_min_size)
|
||||
pool_max_size = max(1, pool_min_size, pool_max_size)
|
||||
key = CacheKey(project_id=project_id, db_role=db_role)
|
||||
entry = self._ts_cache.get(key)
|
||||
if entry:
|
||||
if (
|
||||
entry.connection_url == connection_url
|
||||
and entry.pool_min_size == pool_min_size
|
||||
and entry.pool_max_size == pool_max_size
|
||||
):
|
||||
self._ts_cache.move_to_end(key)
|
||||
return key, entry.pool
|
||||
logger.info(
|
||||
"Created PostgreSQL engine for project %s (%s)", project_id, db_role
|
||||
"Rebuilding TimescaleDB pool for project %s (%s) due to config change",
|
||||
project_id,
|
||||
db_role,
|
||||
)
|
||||
return sessionmaker
|
||||
|
||||
async def get_timescale_pool(
|
||||
self,
|
||||
project_id: UUID,
|
||||
db_role: str,
|
||||
connection_url: str,
|
||||
pool_min_size: int,
|
||||
pool_max_size: int,
|
||||
) -> AsyncConnectionPool:
|
||||
async with self._ts_lock:
|
||||
pool_min_size = max(1, pool_min_size)
|
||||
pool_max_size = max(pool_min_size, pool_max_size)
|
||||
|
||||
key = CacheKey(project_id=project_id, db_role=db_role)
|
||||
entry = self._ts_cache.get(key)
|
||||
if entry:
|
||||
if (
|
||||
entry.connection_url == connection_url
|
||||
and entry.pool_min_size == pool_min_size
|
||||
and entry.pool_max_size == pool_max_size
|
||||
):
|
||||
self._ts_cache.move_to_end(key)
|
||||
return entry.pool
|
||||
|
||||
pool = AsyncConnectionPool(
|
||||
conninfo=connection_url,
|
||||
min_size=pool_min_size,
|
||||
max_size=pool_max_size,
|
||||
open=False,
|
||||
kwargs={"row_factory": dict_row},
|
||||
check=_check_async_connection,
|
||||
)
|
||||
await pool.open()
|
||||
if entry is not None:
|
||||
if entry.borrow_count:
|
||||
self._retired_ts.append((key, entry))
|
||||
else:
|
||||
await entry.pool.close()
|
||||
logger.info(
|
||||
"Rebuilding TimescaleDB pool for project %s (%s) due to config change",
|
||||
project_id,
|
||||
db_role,
|
||||
)
|
||||
self._ts_cache.pop(key, None)
|
||||
self._ts_cache[key] = PoolEntry(
|
||||
pool=pool,
|
||||
connection_url=connection_url,
|
||||
pool_min_size=pool_min_size,
|
||||
pool_max_size=pool_max_size,
|
||||
)
|
||||
logger.info("Created TimescaleDB pool for project %s (%s)", project_id, db_role)
|
||||
return key, pool
|
||||
|
||||
pool = AsyncConnectionPool(
|
||||
conninfo=connection_url,
|
||||
min_size=pool_min_size,
|
||||
max_size=pool_max_size,
|
||||
open=False,
|
||||
kwargs={"row_factory": dict_row},
|
||||
)
|
||||
await pool.open()
|
||||
self._ts_cache[key] = PoolEntry(
|
||||
pool=pool,
|
||||
connection_url=connection_url,
|
||||
pool_min_size=pool_min_size,
|
||||
pool_max_size=pool_max_size,
|
||||
)
|
||||
await self._evict_ts_if_needed()
|
||||
logger.info(
|
||||
"Created TimescaleDB pool for project %s (%s)", project_id, db_role
|
||||
)
|
||||
return pool
|
||||
|
||||
async def get_pg_pool(
|
||||
async def _get_pg_pool_locked(
|
||||
self,
|
||||
project_id: UUID,
|
||||
db_role: str,
|
||||
connection_url: str,
|
||||
pool_min_size: int,
|
||||
pool_max_size: int,
|
||||
) -> AsyncConnectionPool:
|
||||
) -> tuple[CacheKey, AsyncConnectionPool]:
|
||||
pool_min_size = max(0, pool_min_size)
|
||||
pool_max_size = max(1, pool_min_size, pool_max_size)
|
||||
key = CacheKey(project_id=project_id, db_role=db_role)
|
||||
entry = self._pg_raw_cache.get(key)
|
||||
if entry:
|
||||
if (
|
||||
entry.connection_url == connection_url
|
||||
and entry.pool_min_size == pool_min_size
|
||||
and entry.pool_max_size == pool_max_size
|
||||
):
|
||||
self._pg_raw_cache.move_to_end(key)
|
||||
return key, entry.pool
|
||||
logger.info(
|
||||
"Rebuilding PostgreSQL pool for project %s (%s) due to config change",
|
||||
project_id,
|
||||
db_role,
|
||||
)
|
||||
|
||||
pool = AsyncConnectionPool(
|
||||
conninfo=connection_url,
|
||||
min_size=pool_min_size,
|
||||
max_size=pool_max_size,
|
||||
open=False,
|
||||
kwargs={"row_factory": dict_row},
|
||||
check=_check_async_connection,
|
||||
)
|
||||
await pool.open()
|
||||
if entry is not None:
|
||||
if entry.borrow_count:
|
||||
self._retired_pg.append((key, entry))
|
||||
else:
|
||||
await entry.pool.close()
|
||||
self._pg_raw_cache[key] = PoolEntry(
|
||||
pool=pool,
|
||||
connection_url=connection_url,
|
||||
pool_min_size=pool_min_size,
|
||||
pool_max_size=pool_max_size,
|
||||
)
|
||||
logger.info("Created PostgreSQL pool for project %s (%s)", project_id, db_role)
|
||||
return key, pool
|
||||
|
||||
@asynccontextmanager
|
||||
async def pg_connection(
|
||||
self,
|
||||
project_id: UUID,
|
||||
db_role: str,
|
||||
connection_url: str,
|
||||
pool_min_size: int,
|
||||
pool_max_size: int,
|
||||
) -> AsyncIterator[AsyncConnection]:
|
||||
async with self._pg_raw_lock:
|
||||
pool_min_size = max(1, pool_min_size)
|
||||
pool_max_size = max(pool_min_size, pool_max_size)
|
||||
|
||||
key = CacheKey(project_id=project_id, db_role=db_role)
|
||||
entry = self._pg_raw_cache.get(key)
|
||||
if entry:
|
||||
if (
|
||||
entry.connection_url == connection_url
|
||||
and entry.pool_min_size == pool_min_size
|
||||
and entry.pool_max_size == pool_max_size
|
||||
):
|
||||
self._pg_raw_cache.move_to_end(key)
|
||||
return entry.pool
|
||||
|
||||
await entry.pool.close()
|
||||
logger.info(
|
||||
"Rebuilding PostgreSQL pool for project %s (%s) due to config change",
|
||||
project_id,
|
||||
db_role,
|
||||
)
|
||||
self._pg_raw_cache.pop(key, None)
|
||||
|
||||
pool = AsyncConnectionPool(
|
||||
conninfo=connection_url,
|
||||
min_size=pool_min_size,
|
||||
max_size=pool_max_size,
|
||||
open=False,
|
||||
kwargs={"row_factory": dict_row},
|
||||
)
|
||||
await pool.open()
|
||||
self._pg_raw_cache[key] = PoolEntry(
|
||||
pool=pool,
|
||||
connection_url=connection_url,
|
||||
pool_min_size=pool_min_size,
|
||||
pool_max_size=pool_max_size,
|
||||
key, pool = await self._get_pg_pool_locked(
|
||||
project_id,
|
||||
db_role,
|
||||
connection_url,
|
||||
pool_min_size,
|
||||
pool_max_size,
|
||||
)
|
||||
borrowed_entry = self._pg_raw_cache[key]
|
||||
borrowed_entry.borrow_count += 1
|
||||
await self._evict_pg_raw_if_needed()
|
||||
logger.info(
|
||||
"Created PostgreSQL pool for project %s (%s)", project_id, db_role
|
||||
)
|
||||
return pool
|
||||
|
||||
async def _evict_pg_if_needed(self) -> None:
|
||||
while len(self._pg_cache) > settings.PROJECT_PG_CACHE_SIZE:
|
||||
key, entry = self._pg_cache.popitem(last=False)
|
||||
await entry.engine.dispose()
|
||||
logger.info(
|
||||
"Evicted PostgreSQL engine for project %s (%s)",
|
||||
key.project_id,
|
||||
key.db_role,
|
||||
)
|
||||
try:
|
||||
async with pool.connection() as conn:
|
||||
yield conn
|
||||
finally:
|
||||
async with self._pg_raw_lock:
|
||||
borrowed_entry.borrow_count -= 1
|
||||
if borrowed_entry.borrow_count == 0 and borrowed_entry.pool is not (
|
||||
self._pg_raw_cache.get(key).pool
|
||||
if key in self._pg_raw_cache
|
||||
else None
|
||||
):
|
||||
self._retired_pg = [
|
||||
item for item in self._retired_pg if item[1] is not borrowed_entry
|
||||
]
|
||||
await borrowed_entry.pool.close()
|
||||
await self._evict_pg_raw_if_needed()
|
||||
|
||||
async def _evict_ts_if_needed(self) -> None:
|
||||
@asynccontextmanager
|
||||
async def timescale_connection(
|
||||
self,
|
||||
project_id: UUID,
|
||||
db_role: str,
|
||||
connection_url: str,
|
||||
pool_min_size: int,
|
||||
pool_max_size: int,
|
||||
) -> AsyncIterator[AsyncConnection]:
|
||||
async with self._ts_lock:
|
||||
key, pool = await self._get_timescale_pool_locked(
|
||||
project_id,
|
||||
db_role,
|
||||
connection_url,
|
||||
pool_min_size,
|
||||
pool_max_size,
|
||||
)
|
||||
borrowed_entry = self._ts_cache[key]
|
||||
borrowed_entry.borrow_count += 1
|
||||
await self._evict_ts_if_needed()
|
||||
|
||||
try:
|
||||
async with pool.connection() as conn:
|
||||
yield conn
|
||||
finally:
|
||||
async with self._ts_lock:
|
||||
borrowed_entry.borrow_count -= 1
|
||||
if borrowed_entry.borrow_count == 0 and borrowed_entry.pool is not (
|
||||
self._ts_cache.get(key).pool
|
||||
if key in self._ts_cache
|
||||
else None
|
||||
):
|
||||
self._retired_ts = [
|
||||
item for item in self._retired_ts if item[1] is not borrowed_entry
|
||||
]
|
||||
await borrowed_entry.pool.close()
|
||||
await self._evict_ts_if_needed()
|
||||
|
||||
async def _evict_ts_if_needed(
|
||||
self, protected_key: CacheKey | None = None
|
||||
) -> None:
|
||||
while len(self._ts_cache) > settings.PROJECT_TS_CACHE_SIZE:
|
||||
key, entry = self._ts_cache.popitem(last=False)
|
||||
idle = next(
|
||||
(
|
||||
(key, entry)
|
||||
for key, entry in self._ts_cache.items()
|
||||
if entry.borrow_count == 0 and key != protected_key
|
||||
),
|
||||
None,
|
||||
)
|
||||
if idle is None:
|
||||
return
|
||||
key, entry = idle
|
||||
self._ts_cache.pop(key)
|
||||
await entry.pool.close()
|
||||
logger.info(
|
||||
"Evicted TimescaleDB pool for project %s (%s)",
|
||||
@@ -232,9 +238,22 @@ class ProjectConnectionManager:
|
||||
key.db_role,
|
||||
)
|
||||
|
||||
async def _evict_pg_raw_if_needed(self) -> None:
|
||||
async def _evict_pg_raw_if_needed(
|
||||
self, protected_key: CacheKey | None = None
|
||||
) -> None:
|
||||
while len(self._pg_raw_cache) > settings.PROJECT_PG_CACHE_SIZE:
|
||||
key, entry = self._pg_raw_cache.popitem(last=False)
|
||||
idle = next(
|
||||
(
|
||||
(key, entry)
|
||||
for key, entry in self._pg_raw_cache.items()
|
||||
if entry.borrow_count == 0 and key != protected_key
|
||||
),
|
||||
None,
|
||||
)
|
||||
if idle is None:
|
||||
return
|
||||
key, entry = idle
|
||||
self._pg_raw_cache.pop(key)
|
||||
await entry.pool.close()
|
||||
logger.info(
|
||||
"Evicted PostgreSQL pool for project %s (%s)",
|
||||
@@ -242,17 +261,61 @@ class ProjectConnectionManager:
|
||||
key.db_role,
|
||||
)
|
||||
|
||||
async def close_all(self) -> None:
|
||||
async with self._pg_lock:
|
||||
for key, entry in list(self._pg_cache.items()):
|
||||
await entry.engine.dispose()
|
||||
logger.info(
|
||||
"Closed PostgreSQL engine for project %s (%s)",
|
||||
key.project_id,
|
||||
key.db_role,
|
||||
)
|
||||
self._pg_cache.clear()
|
||||
async def close_project(
|
||||
self, project_id: UUID, db_role: str | None = None
|
||||
) -> bool:
|
||||
"""Close this worker's idle pools for a project.
|
||||
|
||||
Returns ``False`` without interrupting requests when any matching pool is
|
||||
currently borrowed. Callers may retry after those requests complete.
|
||||
"""
|
||||
closed = True
|
||||
for cache, lock, label in (
|
||||
(self._ts_cache, self._ts_lock, "TimescaleDB"),
|
||||
(self._pg_raw_cache, self._pg_raw_lock, "PostgreSQL"),
|
||||
):
|
||||
async with lock:
|
||||
keys = [
|
||||
key
|
||||
for key in cache
|
||||
if key.project_id == project_id
|
||||
and (db_role is None or key.db_role == db_role)
|
||||
]
|
||||
for key in keys:
|
||||
entry = cache[key]
|
||||
if entry.borrow_count:
|
||||
closed = False
|
||||
continue
|
||||
cache.pop(key)
|
||||
await entry.pool.close()
|
||||
logger.info(
|
||||
"Closed %s pool for project %s (%s)",
|
||||
label,
|
||||
key.project_id,
|
||||
key.db_role,
|
||||
)
|
||||
for retired, lock in (
|
||||
(self._retired_ts, self._ts_lock),
|
||||
(self._retired_pg, self._pg_raw_lock),
|
||||
):
|
||||
async with lock:
|
||||
matches = [
|
||||
item
|
||||
for item in retired
|
||||
if item[0].project_id == project_id
|
||||
and (db_role is None or item[0].db_role == db_role)
|
||||
]
|
||||
if any(entry.borrow_count for _key, entry in matches):
|
||||
closed = False
|
||||
for item in matches:
|
||||
key, entry = item
|
||||
if entry.borrow_count:
|
||||
continue
|
||||
retired.remove(item)
|
||||
await entry.pool.close()
|
||||
return closed
|
||||
|
||||
async def close_all(self) -> None:
|
||||
async with self._ts_lock:
|
||||
for key, entry in list(self._ts_cache.items()):
|
||||
await entry.pool.close()
|
||||
@@ -262,6 +325,9 @@ class ProjectConnectionManager:
|
||||
key.db_role,
|
||||
)
|
||||
self._ts_cache.clear()
|
||||
for _key, entry in self._retired_ts:
|
||||
await entry.pool.close()
|
||||
self._retired_ts.clear()
|
||||
|
||||
async with self._pg_raw_lock:
|
||||
for key, entry in list(self._pg_raw_cache.items()):
|
||||
@@ -272,6 +338,9 @@ class ProjectConnectionManager:
|
||||
key.db_role,
|
||||
)
|
||||
self._pg_raw_cache.clear()
|
||||
for _key, entry in self._retired_pg:
|
||||
await entry.pool.close()
|
||||
self._retired_pg.clear()
|
||||
|
||||
|
||||
project_connection_manager = ProjectConnectionManager()
|
||||
|
||||
Reference in New Issue
Block a user