Files
TJWaterServerBinary/app/native/wndb/core/projects.py
T

268 lines
9.7 KiB
Python

from collections.abc import Iterable
from contextlib import contextmanager
import re
from uuid import uuid4
from psycopg import sql
from psycopg.rows import dict_row
from app.core.config import settings
from app.infra.db.project_routing import (
get_project_database_name,
get_project_template_database_name,
)
from .connection import (
admin_connection,
close_project_pool,
)
_SERVER_DATABASES = frozenset({"template0", "template1", "postgres", "project"})
_TEMPORARY_DATABASE_PREFIX = "tjw_tmp_"
def _protected_databases() -> frozenset[str]:
return _SERVER_DATABASES | {
settings.METADATA_DB_NAME,
settings.WNDB_TEMPLATE_DB_NAME,
}
def _validate_project_database(name: str, *, allow_template_source: bool = False) -> None:
if not name:
raise ValueError("Project database name must not be empty")
protected = {database.casefold() for database in _protected_databases()}
is_template = name.casefold().endswith("_template")
if (
allow_template_source
and name.casefold() == settings.WNDB_TEMPLATE_DB_NAME.casefold()
):
return
if name.casefold() in protected or is_template:
raise ValueError(f"Database {name!r} is protected and cannot be managed as a project")
def list_project() -> list[str]:
ps = []
with admin_connection() as conn:
with conn.cursor(row_factory=dict_row) as cur:
for p in cur.execute(
"select datname from pg_database where datname <> all(%s) order by datname",
(list(_protected_databases()),),
):
if not str(p["datname"]).casefold().endswith("_template"):
ps.append(p["datname"])
return ps
@contextmanager
def _database_locks(cur, *database_names: str):
"""Serialize physical database lifecycle operations across server workers."""
lock_names = sorted(set(database_names), key=str.casefold)
for database_name in lock_names:
cur.execute(
"select pg_advisory_lock(hashtextextended(%s, 0))",
(f"tjwater:wndb:{database_name}",),
)
try:
yield
finally:
for database_name in reversed(lock_names):
cur.execute(
"select pg_advisory_unlock(hashtextextended(%s, 0))",
(f"tjwater:wndb:{database_name}",),
)
@contextmanager
def _temporary_database_capacity(cur, database_name: str):
"""Serialize temporary creation and enforce a server-wide hard limit."""
if not database_name.startswith(_TEMPORARY_DATABASE_PREFIX):
yield
return
lock_name = "tjwater:wndb:temporary-database-capacity"
cur.execute("select pg_advisory_lock(hashtextextended(%s, 0))", (lock_name,))
try:
cur.execute(
"select count(*) as count from pg_database where datname like %s",
(f"{_TEMPORARY_DATABASE_PREFIX}%",),
)
row = cur.fetchone()
count = int(row["count"] if row is not None else 0)
limit = max(1, settings.WNDB_TEMP_DB_MAX_COUNT)
if count >= limit:
raise RuntimeError(
f"Temporary database limit reached ({count}/{limit}); "
"retry after an active analysis completes"
)
yield
finally:
cur.execute("select pg_advisory_unlock(hashtextextended(%s, 0))", (lock_name,))
def _database_allows_connections(cur, database_name: str) -> bool:
cur.execute(
"select datallowconn from pg_database where datname = %s",
(database_name,),
)
row = cur.fetchone()
if row is None:
raise ValueError(f"Database {database_name!r} does not exist")
return bool(row["datallowconn"])
def _set_database_connections(cur, database_name: str, *, allowed: bool) -> None:
cur.execute(
"update pg_database set datallowconn = %s where datname = %s",
(allowed, database_name),
)
def temporary_project_name(project: str, purpose: str) -> str:
"""Return a collision-resistant physical database name for one run."""
physical_name = get_project_database_name(project)
safe_purpose = re.sub(r"[^a-z0-9_]+", "_", purpose.casefold()).strip("_")
safe_project = re.sub(r"[^a-z0-9_]+", "_", physical_name.casefold()).strip("_")
prefix = (
f"{_TEMPORARY_DATABASE_PREFIX}{safe_purpose or 'run'}_"
f"{safe_project or 'project'}"
)[:29]
return f"{prefix}_{uuid4().hex}"
@contextmanager
def temporary_project_database(project: str, purpose: str):
"""Clone one project's runnable model into an isolated temporary database."""
temporary_name = temporary_project_name(project, purpose)
try:
copy_project(get_project_template_database_name(project), temporary_name)
# Import lazily to keep physical database lifecycle independent from
# model-copy implementation details at module import time.
from .database import refresh_materialized_views_after_commit
from .model_replace import replace_project_model
replace_project_model(
temporary_name,
project,
copy_source_scada=True,
)
refresh_materialized_views_after_commit(temporary_name)
yield temporary_name
finally:
if have_project(temporary_name):
delete_project(temporary_name)
@contextmanager
def temporary_template_database(name_hint: str, purpose: str):
"""Create an empty schema-only temporary database from the fixed template."""
temporary_name = temporary_project_name(name_hint, purpose)
try:
copy_project(get_project_template_database_name(name_hint), temporary_name)
yield temporary_name
finally:
if have_project(temporary_name):
delete_project(temporary_name)
def have_project(name: str) -> bool:
database_name = get_project_database_name(name)
with admin_connection() as conn:
with conn.cursor() as cur:
cur.execute("select 1 from pg_database where datname = %s", (database_name,))
return cur.fetchone() is not None
def copy_project(source: str, new: str) -> None:
physical_source = get_project_database_name(source)
physical_new = get_project_database_name(new)
_validate_project_database(physical_source, allow_template_source=True)
_validate_project_database(physical_new)
close_project_pool(source)
close_project_pool(new)
with admin_connection() as admin_conn:
with admin_conn.cursor() as cur:
with _temporary_database_capacity(cur, physical_new):
with _database_locks(cur, physical_source, physical_new):
source_allowed = _database_allows_connections(cur, physical_source)
if source_allowed:
_set_database_connections(
cur,
physical_source,
allowed=False,
)
try:
cur.execute(
"select pg_terminate_backend(pid) from pg_stat_activity "
"where datname = %s and pid <> pg_backend_pid()",
(physical_source,),
)
cur.execute(
sql.SQL("create database {} with template = {}").format(
sql.Identifier(physical_new),
sql.Identifier(physical_source),
)
)
finally:
if source_allowed:
_set_database_connections(
cur,
physical_source,
allowed=True,
)
def create_project(name: str) -> None:
return copy_project(get_project_template_database_name(name), name)
def delete_project(name: str) -> None:
database_name = get_project_database_name(name)
_validate_project_database(database_name)
close_project_pool(name)
with admin_connection() as conn:
with conn.cursor() as cur:
with _database_locks(cur, database_name):
was_allowed = _database_allows_connections(cur, database_name)
if was_allowed:
_set_database_connections(cur, database_name, allowed=False)
try:
cur.execute(
"select pg_terminate_backend(pid) from pg_stat_activity "
"where datname = %s and pid <> pg_backend_pid()",
(database_name,),
)
cur.execute(
sql.SQL("drop database {}").format(
sql.Identifier(database_name)
)
)
except Exception:
if was_allowed:
_set_database_connections(cur, database_name, allowed=True)
raise
def clean_project(projects: Iterable[str]) -> None:
"""Delete only the explicitly supplied project databases."""
targets = list(dict.fromkeys(projects))
physical_targets = [get_project_database_name(project) for project in targets]
for database_name in physical_targets:
_validate_project_database(database_name)
if not targets:
return
with admin_connection() as conn:
with conn.cursor(row_factory=dict_row) as cur:
row = cur.execute("select current_database()").fetchone()
current_db = row["current_database"] if row is not None else None
if current_db in physical_targets:
raise ValueError(f"Cannot delete the current database {current_db!r}")
for project in targets:
delete_project(project)