diff --git a/.env.example b/.env.example index c46ebf3..8300125 100644 --- a/.env.example +++ b/.env.example @@ -47,9 +47,25 @@ PROJECT_PG_POOL_SIZE="4" PROJECT_PG_MAX_OVERFLOW="2" PROJECT_TS_POOL_MIN_SIZE="0" PROJECT_TS_POOL_MAX_SIZE="4" -WNDB_TEMPLATE_DB_NAME="tjwater_v2_template" +WNDB_SCHEMA_TEMPLATE_DB_NAME="tjwater_v2_schema_template" +TIMESCALEDB_SCHEMA_TEMPLATE_DB_NAME="tjwater_v2_timescale_template" WNDB_TEMP_DB_MAX_COUNT="8" +# ============================================ +# GeoServer 项目供应 +# ============================================ +GEOSERVER_URL="http://localhost:8080/geoserver" +GEOSERVER_USERNAME="admin" +GEOSERVER_PASSWORD="password" +# 留空时分别复用 DB_HOST / DB_PORT / DB_USER / DB_PASSWORD。 +# GeoServer 在容器内运行时,GEOSERVER_DB_HOST 通常应填写数据库服务名。 +GEOSERVER_DB_HOST="" +GEOSERVER_DB_PORT="" +GEOSERVER_DB_USER="" +GEOSERVER_DB_PASSWORD="" +# 30 天客户端缓存,单位为秒。 +GEOSERVER_CLIENT_CACHE_SECONDS="2592000" + # ============================================ # Keycloak JWT (可选) # ============================================ diff --git a/README.md b/README.md index 9f552e4..0a109a2 100644 --- a/README.md +++ b/README.md @@ -69,10 +69,22 @@ docker compose -f infra/docker/docker-compose.yml config 项目级 REST 请求通过 `X-Project-Id` 解析元数据中的数据库配置: -- `biz_data` DSN 用于管网业务数据;版本模板固定由 `WNDB_TEMPLATE_DB_NAME` 配置(当前为 `tjwater_v2_template`),模拟临时库沿用该 DSN 的主机、端口与凭据,仅替换数据库名。 +- `biz_data` DSN 用于管网业务数据。每个物理业务库使用同名 `_template` 数据库,通过逻辑订阅只同步 `network` schema;模拟临时库从该项目模板克隆。`WNDB_SCHEMA_TEMPLATE_DB_NAME`(当前为 `tjwater_v2_schema_template`)仅用于创建空业务库和 INP 导入暂存库。 - `iot_data` DSN 用于 TimescaleDB,始终使用元数据配置的完整 DSN,不再从项目代码推导数据库名。 - 元数据、业务库和 TimescaleDB 可以部署在同一主机,也可以分别部署。 +完整新建供水项目使用 `POST /api/v1/admin/project-provisions`,以 +`multipart/form-data` 同时提交 `name`、小写 `code`、可选的 +`description`、`gs_workspace`、`map_zoom` 和 INP `file`。工作流会按顺序完成: + +1. EPANET 校验 INP,并从空结构模板创建业务库、导入模型; +2. 创建同名 `_template`,复制 32 张 `network` 表并建立逻辑订阅; +3. 从 `TIMESCALEDB_SCHEMA_TEMPLATE_DB_NAME` 创建空时序库; +4. 创建 GeoServer 工作空间、PostGIS 数据存储和 7 个 GIS 图层,将客户端缓存设为 `GEOSERVER_CLIENT_CACHE_SECONDS`; +5. 最后在一个元数据事务中写入项目、两条加密数据库路由和创建者成员关系,并将项目设为 `active`。 + +基础设施任一步失败时按 GeoServer、时序库、管网模板、业务库的逆序清理;元数据提交失败也执行同样清理。旧 `POST /admin/projects` 仅保留给已经由外部流程创建好的资源登记使用,并已标记为 deprecated。 + 使用模板复制或临时方案库的模拟功能时,`biz_data` 账号必须具备数据库创建和删除权限;只有显式删除项目时才会终止该项目的现有数据库会话,普通复制不会主动中断复制源会话。 ## 测试与发布 diff --git a/app/api/v1/endpoints/admin_metadata.py b/app/api/v1/endpoints/admin_metadata.py index b8bc55b..81c9a1c 100644 --- a/app/api/v1/endpoints/admin_metadata.py +++ b/app/api/v1/endpoints/admin_metadata.py @@ -252,6 +252,12 @@ async def list_admin_projects( "/admin/projects", response_model=AdminProjectResponse, status_code=status.HTTP_201_CREATED, + deprecated=True, + summary="仅登记已有项目元数据", + description=( + "仅用于登记已经由外部流程完整创建的资源。新项目应调用 " + "POST /admin/project-provisions。" + ), ) async def create_admin_project( payload: AdminProjectCreateRequest, diff --git a/app/api/v1/endpoints/model_import.py b/app/api/v1/endpoints/model_import.py index b62e6a8..d43b2dc 100644 --- a/app/api/v1/endpoints/model_import.py +++ b/app/api/v1/endpoints/model_import.py @@ -1,3 +1,4 @@ +import json from pathlib import Path from tempfile import NamedTemporaryFile from uuid import UUID, uuid4 @@ -6,12 +7,14 @@ from fastapi import ( APIRouter, Depends, File, + Form, HTTPException, Path as ApiPath, Request, UploadFile, status, ) +from sqlalchemy.exc import IntegrityError from starlette.concurrency import run_in_threadpool from app.auth.metadata_dependencies import ( @@ -23,10 +26,21 @@ from app.auth.project_dependencies import ( resolve_project_business_routing, ) from app.core.audit import AuditAction, log_audit_event +from app.core.encryption import is_database_encryption_configured +from app.domain.schemas.admin_metadata import ( + AdminProjectResponse, + ProjectProvisionResponse, +) from app.infra.db.metadb.repositories.metadata_repository import MetadataRepository from app.infra.db.project_routing import activate_project_routing from app.native.wndb.core.database import MaterializedViewRefreshAfterCommitError from app.services.network_import import network_update +from app.services.project_provisioning import ( + ProjectProvisioningError, + ProvisionedProjectInfrastructure, + provision_project_infrastructure, + validate_project_code, +) from app.services.tjnetwork import run_inp router = APIRouter() @@ -157,6 +171,166 @@ async def _apply_model_update(content: bytes, project_code: str) -> None: ) from exc +def _provision_from_inp_sync( + content: bytes, + *, + code: str, + workspace: str, +) -> ProvisionedProjectInfrastructure: + temp_path: Path | None = None + try: + with NamedTemporaryFile(suffix=".inp", delete=False) as temp_file: + temp_file.write(content) + temp_path = Path(temp_file.name) + return provision_project_infrastructure( + code=code, + workspace=workspace, + inp_path=temp_path, + ) + finally: + if temp_path is not None: + temp_path.unlink(missing_ok=True) + + +@router.post( + "/admin/project-provisions", + response_model=ProjectProvisionResponse, + status_code=status.HTTP_201_CREATED, + summary="创建完整供水项目", +) +async def provision_project( + request: Request, + name: str = Form(..., min_length=1, max_length=100), + code: str = Form(..., min_length=1, max_length=50), + description: str | None = Form(default=None), + gs_workspace: str | None = Form(default=None, max_length=100), + map_zoom: int = Form(default=14, ge=1, le=22), + file: UploadFile = File(..., description="EPANET INP 模型文件"), + current_user=Depends(get_current_metadata_admin), + metadata_repo: MetadataRepository = Depends(get_metadata_repository), +) -> ProjectProvisionResponse: + try: + normalized_code = validate_project_code(code) + except ValueError as exc: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=str(exc), + ) from exc + workspace = gs_workspace or normalized_code + if await metadata_repo.get_project_by_code(normalized_code) is not None: + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, + detail="Project code already exists", + ) + if not is_database_encryption_configured(): + raise HTTPException( + status_code=status.HTTP_503_SERVICE_UNAVAILABLE, + detail="DATABASE_ENCRYPTION_KEY is not configured", + ) + + content, filename = await _read_upload(file) + validation_result = await _run_uploaded_inp(content) + try: + validation_payload = json.loads(validation_result) + except (TypeError, json.JSONDecodeError) as exc: + raise HTTPException( + status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, + detail="EPANET validation returned an invalid response", + ) from exc + if validation_payload.get("simulation_result") != "successful": + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail="EPANET model validation failed", + ) + try: + infrastructure = await run_in_threadpool( + _provision_from_inp_sync, + content, + code=normalized_code, + workspace=workspace, + ) + except ProjectProvisioningError as exc: + if isinstance(exc.cause, ValueError): + response_status = status.HTTP_409_CONFLICT + elif exc.stage == "preflight": + response_status = status.HTTP_503_SERVICE_UNAVAILABLE + else: + response_status = status.HTTP_500_INTERNAL_SERVER_ERROR + raise HTTPException( + status_code=response_status, + detail={ + "stage": exc.stage, + "message": str(exc.cause), + "cleanup_errors": exc.cleanup_errors, + }, + ) from exc + + map_extent = {"bbox": list(infrastructure.map_bbox), "zoom": map_zoom} + try: + project = await metadata_repo.create_provisioned_project( + name=name, + code=normalized_code, + description=description, + gs_workspace=workspace, + map_extent=map_extent, + creator_user_id=current_user.id, + business_dsn=infrastructure.business_dsn, + timescale_dsn=infrastructure.timescale_dsn, + pool_min_size=1, + pool_max_size=4, + ) + except Exception as exc: + await metadata_repo.session.rollback() + cleanup_errors = await run_in_threadpool(infrastructure.cleanup) + if isinstance(exc, IntegrityError): + response_status = status.HTTP_409_CONFLICT + detail = "Project code or workspace conflicts with an existing project" + else: + response_status = status.HTTP_503_SERVICE_UNAVAILABLE + detail = f"Metadata database error: {exc}" + if cleanup_errors: + detail = f"{detail}; cleanup failures: {', '.join(cleanup_errors)}" + raise HTTPException(status_code=response_status, detail=detail) from exc + + await log_audit_event( + action=AuditAction.CREATE, + user_id=current_user.id, + project_id=project.id, + resource_type="project_provision", + resource_id=str(project.id), + request_data={ + "name": name, + "code": normalized_code, + "filename": filename, + "gs_workspace": workspace, + "layers": list(infrastructure.layers), + }, + ip_address=request.client.host if request.client else None, + request_method=request.method, + request_path=request.url.path, + response_status=status.HTTP_201_CREATED, + session=metadata_repo.session, + ) + return ProjectProvisionResponse( + project=AdminProjectResponse( + project_id=project.id, + name=project.name, + code=project.code, + description=project.description, + gs_workspace=project.gs_workspace, + map_extent=project.map_extent, + status=project.status, + created_at=project.created_at, + updated_at=project.updated_at, + ), + business_database=normalized_code, + model_template_database=infrastructure.model_template, + timescale_database=normalized_code, + geoserver_workspace=workspace, + geoserver_layers=list(infrastructure.layers), + ) + + @router.post( "/admin/projects/{project_id}/model-imports", summary="导入桌面端水力模型", diff --git a/app/core/config.py b/app/core/config.py index 8a00719..273e612 100644 --- a/app/core/config.py +++ b/app/core/config.py @@ -43,9 +43,20 @@ class Settings(BaseSettings): PROJECT_PG_MAX_OVERFLOW: int = 2 PROJECT_TS_POOL_MIN_SIZE: int = 0 PROJECT_TS_POOL_MAX_SIZE: int = 4 - WNDB_TEMPLATE_DB_NAME: str = "tjwater_v2_template" + WNDB_SCHEMA_TEMPLATE_DB_NAME: str = "tjwater_v2_schema_template" + TIMESCALEDB_SCHEMA_TEMPLATE_DB_NAME: str = "tjwater_v2_timescale_template" WNDB_TEMP_DB_MAX_COUNT: int = 8 + # GeoServer project provisioning + GEOSERVER_URL: str = "http://localhost:8080/geoserver" + GEOSERVER_USERNAME: str = "" + GEOSERVER_PASSWORD: str = "" + GEOSERVER_DB_HOST: str = "" + GEOSERVER_DB_PORT: str = "" + GEOSERVER_DB_USER: str = "" + GEOSERVER_DB_PASSWORD: str = "" + GEOSERVER_CLIENT_CACHE_SECONDS: int = 30 * 24 * 60 * 60 + # Keycloak access token verification KEYCLOAK_PUBLIC_KEY: str = "" KEYCLOAK_ALGORITHM: str = "RS256" diff --git a/app/domain/schemas/admin_metadata.py b/app/domain/schemas/admin_metadata.py index 6ebfc25..584cc4b 100644 --- a/app/domain/schemas/admin_metadata.py +++ b/app/domain/schemas/admin_metadata.py @@ -99,6 +99,15 @@ class AdminProjectResponse(BaseModel): updated_at: datetime +class ProjectProvisionResponse(BaseModel): + project: AdminProjectResponse + business_database: str + model_template_database: str + timescale_database: str + geoserver_workspace: str + geoserver_layers: list[str] + + class ProjectDatabaseUpsertRequest(BaseModel): db_role: ProjectDbRole dsn: str | None = Field(default=None, min_length=1) diff --git a/app/infra/db/metadb/repositories/metadata_repository.py b/app/infra/db/metadb/repositories/metadata_repository.py index f14e1d2..469683f 100644 --- a/app/infra/db/metadb/repositories/metadata_repository.py +++ b/app/infra/db/metadb/repositories/metadata_repository.py @@ -241,6 +241,67 @@ class MetadataRepository: await self.session.refresh(project) return project + async def create_provisioned_project( + self, + *, + name: str, + code: str, + description: str | None, + gs_workspace: str, + map_extent: dict, + creator_user_id: UUID, + business_dsn: str, + timescale_dsn: str, + pool_min_size: int = 1, + pool_max_size: int = 4, + ) -> models.Project: + """Atomically expose a fully provisioned project and both DB routes.""" + business_secret = _encrypt_database_secret(business_dsn) + timescale_secret = _encrypt_database_secret(timescale_dsn) + project = models.Project( + id=uuid4(), + name=name, + code=code, + description=description, + gs_workspace=gs_workspace, + map_extent=map_extent, + status="active", + created_at=_utcnow(), + updated_at=_utcnow(), + ) + records = ( + project, + models.ProjectDatabase( + id=uuid4(), + project_id=project.id, + db_role="biz_data", + db_type="postgresql", + dsn_encrypted=business_secret, + pool_min_size=pool_min_size, + pool_max_size=pool_max_size, + ), + models.ProjectDatabase( + id=uuid4(), + project_id=project.id, + db_role="iot_data", + db_type="timescaledb", + dsn_encrypted=timescale_secret, + pool_min_size=pool_min_size, + pool_max_size=pool_max_size, + ), + models.UserProjectMembership( + id=uuid4(), + user_id=creator_user_id, + project_id=project.id, + project_role="member", + ), + ) + for record in records: + self.session.add(record) + await self.session.commit() + await self.session.refresh(project) + return project + async def update_project( self, project_id: UUID, diff --git a/app/infra/db/postgresql/project_spatial.py b/app/infra/db/postgresql/project_spatial.py new file mode 100644 index 0000000..53bb460 --- /dev/null +++ b/app/infra/db/postgresql/project_spatial.py @@ -0,0 +1,30 @@ +from __future__ import annotations + +from app.native.wndb.core.connection import project_connection + + +def get_project_map_bbox(project: str) -> tuple[float, float, float, float]: + """Return the combined published node/link bounds in EPSG:3857.""" + with project_connection(project) as conn, conn.cursor() as cur: + cur.execute( + """ + with project_geometries as ( + select geom from gis.junctions + union all + select geom from gis.pipes + ), bounds as ( + select ST_Extent(geom) as extent from project_geometries + ) + select ST_XMin(extent) as minx, + ST_YMin(extent) as miny, + ST_XMax(extent) as maxx, + ST_YMax(extent) as maxy + from bounds + """ + ) + row = cur.fetchone() + if row is None or any( + row[key] is None for key in ("minx", "miny", "maxx", "maxy") + ): + raise RuntimeError("Imported model has no publishable GIS geometry") + return tuple(float(row[key]) for key in ("minx", "miny", "maxx", "maxy")) diff --git a/app/infra/db/project_routing.py b/app/infra/db/project_routing.py index 9054fde..85c99e4 100644 --- a/app/infra/db/project_routing.py +++ b/app/infra/db/project_routing.py @@ -62,12 +62,18 @@ def get_project_database_name(name: str) -> str: def get_project_template_database_name(name: str | None = None) -> str: - """Return the configured immutable template for the WNDB schema version. + """Return the network-data template paired with one physical BizDB. - The template belongs to the database schema version, not to an individual - logical project or its temporary physical database name. + Project routing is resolved before the suffix is applied, so a logical + project such as ``tjwater_v2`` uses ``tjwater_v2_template``. """ - return settings.WNDB_TEMPLATE_DB_NAME + database_name = get_project_database_name(name or settings.DB_NAME) + return f"{database_name}_template" + + +def get_schema_template_database_name() -> str: + """Return the immutable, data-free template used to create BizDB schemas.""" + return settings.WNDB_SCHEMA_TEMPLATE_DB_NAME def get_project_pgconn_string(db_name: str | None = None) -> str: diff --git a/app/infra/db/timescaledb/lifecycle.py b/app/infra/db/timescaledb/lifecycle.py new file mode 100644 index 0000000..c2d4dee --- /dev/null +++ b/app/infra/db/timescaledb/lifecycle.py @@ -0,0 +1,165 @@ +from __future__ import annotations + +from contextlib import contextmanager +import re +from threading import RLock +from typing import Iterator + +from psycopg import Connection, sql +from psycopg.rows import dict_row +from psycopg_pool import ConnectionPool + +from app.core.config import get_timescaledb_pgconn_string, settings + +from .sync_pool import close_timescale_pool + + +_DATABASE_NAME = re.compile(r"^[a-z][a-z0-9_]{0,49}$") +_SERVER_DATABASES = frozenset({"template0", "template1", "postgres"}) +_admin_pool: ConnectionPool | None = None +_admin_conninfo: str | None = None +_lock = RLock() + + +def validate_timescale_database_name(name: str, *, allow_template: bool = False) -> str: + if not _DATABASE_NAME.fullmatch(name): + raise ValueError( + "TimescaleDB database name must start with a lowercase letter and " + "contain only lowercase letters, digits, and underscores" + ) + protected = {*_SERVER_DATABASES, settings.TIMESCALEDB_SCHEMA_TEMPLATE_DB_NAME} + if name in protected and not ( + allow_template and name == settings.TIMESCALEDB_SCHEMA_TEMPLATE_DB_NAME + ): + raise ValueError(f"TimescaleDB database {name!r} is protected") + return name + + +def _get_admin_pool() -> ConnectionPool: + global _admin_pool, _admin_conninfo + conninfo = get_timescaledb_pgconn_string(db_name="postgres") + with _lock: + if ( + _admin_pool is not None + and not _admin_pool.closed + and _admin_conninfo == conninfo + ): + return _admin_pool + if _admin_pool is not None and not _admin_pool.closed: + _admin_pool.close() + _admin_pool = ConnectionPool( + conninfo=conninfo, + min_size=0, + max_size=2, + kwargs={"autocommit": True, "row_factory": dict_row}, + check=ConnectionPool.check_connection, + open=True, + ) + _admin_conninfo = conninfo + return _admin_pool + + +@contextmanager +def timescale_admin_connection() -> Iterator[Connection]: + with _get_admin_pool().connection() as conn: + yield conn + + +def timescale_database_exists(name: str) -> bool: + with timescale_admin_connection() as conn, conn.cursor() as cur: + cur.execute("select 1 from pg_database where datname = %s", (name,)) + return cur.fetchone() is not None + + +def require_timescale_schema_template() -> str: + template = settings.TIMESCALEDB_SCHEMA_TEMPLATE_DB_NAME + validate_timescale_database_name(template, allow_template=True) + if not timescale_database_exists(template): + raise RuntimeError( + f"TimescaleDB schema template {template!r} does not exist" + ) + return template + + +def create_timescale_database(name: str) -> None: + validate_timescale_database_name(name) + template = require_timescale_schema_template() + close_timescale_pool(name) + with timescale_admin_connection() as conn, conn.cursor() as cur: + cur.execute( + "select pg_advisory_lock(hashtextextended(%s, 0))", + (f"tjwater:timescaledb:{name}",), + ) + try: + cur.execute("select 1 from pg_database where datname = %s", (name,)) + if cur.fetchone() is not None: + raise ValueError(f"TimescaleDB database {name!r} already exists") + cur.execute( + "select datallowconn from pg_database where datname = %s", + (template,), + ) + row = cur.fetchone() + if row is None: + raise RuntimeError( + f"TimescaleDB schema template {template!r} does not exist" + ) + template_allowed = bool(row["datallowconn"]) + if template_allowed: + cur.execute( + "update pg_database set datallowconn = false where datname = %s", + (template,), + ) + try: + cur.execute( + "select pg_terminate_backend(pid) from pg_stat_activity " + "where datname = %s and pid <> pg_backend_pid()", + (template,), + ) + cur.execute( + sql.SQL("create database {} with template = {}").format( + sql.Identifier(name), + sql.Identifier(template), + ) + ) + finally: + if template_allowed: + cur.execute( + "update pg_database set datallowconn = true where datname = %s", + (template,), + ) + finally: + cur.execute( + "select pg_advisory_unlock(hashtextextended(%s, 0))", + (f"tjwater:timescaledb:{name}",), + ) + + +def delete_timescale_database(name: str) -> None: + validate_timescale_database_name(name) + close_timescale_pool(name) + with timescale_admin_connection() as conn, conn.cursor() as cur: + cur.execute( + "select pg_advisory_lock(hashtextextended(%s, 0))", + (f"tjwater:timescaledb:{name}",), + ) + try: + cur.execute("select 1 from pg_database where datname = %s", (name,)) + if cur.fetchone() is None: + return + cur.execute( + "update pg_database set datallowconn = false where datname = %s", + (name,), + ) + cur.execute( + "select pg_terminate_backend(pid) from pg_stat_activity " + "where datname = %s and pid <> pg_backend_pid()", + (name,), + ) + cur.execute( + sql.SQL("drop database {}").format(sql.Identifier(name)) + ) + finally: + cur.execute( + "select pg_advisory_unlock(hashtextextended(%s, 0))", + (f"tjwater:timescaledb:{name}",), + ) diff --git a/app/infra/geoserver/__init__.py b/app/infra/geoserver/__init__.py new file mode 100644 index 0000000..d41e931 --- /dev/null +++ b/app/infra/geoserver/__init__.py @@ -0,0 +1 @@ +"""GeoServer administration adapters.""" diff --git a/app/infra/geoserver/client.py b/app/infra/geoserver/client.py new file mode 100644 index 0000000..0efaff6 --- /dev/null +++ b/app/infra/geoserver/client.py @@ -0,0 +1,224 @@ +from __future__ import annotations + +from dataclasses import dataclass +import math +import time +from xml.etree import ElementTree + +import httpx + +from app.core.config import settings + + +PROJECT_LAYER_NAMES = ( + "junctions", + "pipes", + "pumps", + "reservoirs", + "scada_devices", + "tanks", + "valves", +) + + +class GeoServerProvisioningError(RuntimeError): + pass + + +@dataclass(frozen=True) +class GeoServerDatabaseConfig: + host: str + port: str + user: str + password: str + + +def _longitude(web_mercator_x: float) -> float: + return web_mercator_x * 180.0 / 20037508.342789244 + + +def _latitude(web_mercator_y: float) -> float: + degrees = web_mercator_y * 180.0 / 20037508.342789244 + return 180.0 / math.pi * ( + 2.0 * math.atan(math.exp(degrees * math.pi / 180.0)) - math.pi / 2.0 + ) + + +class GeoServerAdminClient: + def __init__(self, client: httpx.Client | None = None) -> None: + if not settings.GEOSERVER_USERNAME or not settings.GEOSERVER_PASSWORD: + raise GeoServerProvisioningError( + "GEOSERVER_USERNAME and GEOSERVER_PASSWORD must be configured" + ) + self._owns_client = client is None + self._client = client or httpx.Client( + base_url=settings.GEOSERVER_URL.rstrip("/"), + auth=(settings.GEOSERVER_USERNAME, settings.GEOSERVER_PASSWORD), + timeout=30.0, + ) + + def __enter__(self) -> "GeoServerAdminClient": + return self + + def __exit__(self, *_args) -> None: + if self._owns_client: + self._client.close() + + @staticmethod + def database_config() -> GeoServerDatabaseConfig: + return GeoServerDatabaseConfig( + host=settings.GEOSERVER_DB_HOST or settings.DB_HOST, + port=settings.GEOSERVER_DB_PORT or settings.DB_PORT, + user=settings.GEOSERVER_DB_USER or settings.DB_USER, + password=settings.GEOSERVER_DB_PASSWORD or settings.DB_PASSWORD, + ) + + def _request(self, method: str, path: str, **kwargs) -> httpx.Response: + response = self._client.request(method, path, **kwargs) + if response.is_error: + raise GeoServerProvisioningError( + f"GeoServer {method} {path} failed with HTTP {response.status_code}" + ) + return response + + def check_ready(self) -> None: + self._request("GET", "/rest/about/version.json") + + def workspace_exists(self, workspace: str) -> bool: + response = self._client.get(f"/rest/workspaces/{workspace}.json") + if response.status_code == 404: + return False + if response.is_error: + raise GeoServerProvisioningError( + f"GeoServer workspace check failed with HTTP {response.status_code}" + ) + return True + + def create_project_workspace( + self, + *, + workspace: str, + database_name: str, + map_bbox: tuple[float, float, float, float], + ) -> tuple[str, ...]: + if self.workspace_exists(workspace): + raise ValueError(f"GeoServer workspace {workspace!r} already exists") + self._request( + "POST", + "/rest/workspaces", + json={"workspace": {"name": workspace}}, + ) + database = self.database_config() + entries = { + "dbtype": "postgis", + "host": database.host, + "port": database.port, + "database": database_name, + "schema": "gis", + "user": database.user, + "passwd": database.password, + "namespace": workspace, + "Expose primary keys": "true", + "Estimated extends": "true", + "validate connections": "true", + "min connections": "1", + "max connections": "10", + "Connection timeout": "20", + } + self._request( + "POST", + f"/rest/workspaces/{workspace}/datastores", + json={ + "dataStore": { + "name": workspace, + "description": f"{workspace} GIS materialized views", + "type": "PostGIS", + "enabled": True, + "connectionParameters": { + "entry": [ + {"@key": key, "$": value} + for key, value in entries.items() + ] + }, + } + }, + ) + minx, miny, maxx, maxy = map_bbox + geographic_bbox = { + "minx": _longitude(minx), + "miny": _latitude(miny), + "maxx": _longitude(maxx), + "maxy": _latitude(maxy), + "crs": "EPSG:4326", + } + native_bbox = { + "minx": minx, + "miny": miny, + "maxx": maxx, + "maxy": maxy, + "crs": "EPSG:3857", + } + for layer in PROJECT_LAYER_NAMES: + self._request( + "POST", + f"/rest/workspaces/{workspace}/datastores/{workspace}/featuretypes", + json={ + "featureType": { + "name": layer, + "nativeName": layer, + "title": layer, + "srs": "EPSG:3857", + "projectionPolicy": "FORCE_DECLARED", + "enabled": True, + "advertised": True, + "nativeBoundingBox": native_bbox, + "latLonBoundingBox": geographic_bbox, + } + }, + ) + self._set_client_cache(workspace, layer) + return PROJECT_LAYER_NAMES + + def _set_client_cache(self, workspace: str, layer: str) -> None: + path = f"/gwc/rest/layers/{workspace}:{layer}.xml" + response: httpx.Response | None = None + for _ in range(20): + response = self._client.get(path) + if response.status_code == 200: + break + if response.status_code != 404: + raise GeoServerProvisioningError( + f"GeoWebCache layer check failed with HTTP {response.status_code}" + ) + time.sleep(0.1) + if response is None or response.status_code != 200: + raise GeoServerProvisioningError( + f"GeoWebCache layer {workspace}:{layer} was not registered" + ) + root = ElementTree.fromstring(response.content) + expiry = root.find("expireClients") + if expiry is None: + expiry = ElementTree.SubElement(root, "expireClients") + expiry.text = str(settings.GEOSERVER_CLIENT_CACHE_SECONDS) + self._request( + "PUT", + path, + content=ElementTree.tostring( + root, + encoding="utf-8", + xml_declaration=True, + ), + headers={"Content-Type": "application/xml"}, + ) + + def delete_workspace(self, workspace: str) -> None: + response = self._client.delete( + f"/rest/workspaces/{workspace}", + params={"recurse": "true"}, + ) + if response.status_code == 404: + return + if response.is_error: + raise GeoServerProvisioningError( + f"GeoServer workspace cleanup failed with HTTP {response.status_code}" + ) diff --git a/app/native/wndb/core/model_replace.py b/app/native/wndb/core/model_replace.py index e69e740..825cdde 100644 --- a/app/native/wndb/core/model_replace.py +++ b/app/native/wndb/core/model_replace.py @@ -89,6 +89,56 @@ def _table_columns( return [row["column_name"] for row in cur.fetchall()] +def _geometry_srids( + conn: Connection, schema_name: str, table_name: str +) -> dict[str, int]: + with conn.cursor() as cur: + cur.execute( + """ + select f_geometry_column as column_name, srid + from geometry_columns + where f_table_schema = %s and f_table_name = %s + order by f_geometry_column + """, + (schema_name, table_name), + ) + return {row["column_name"]: int(row["srid"]) for row in cur.fetchall()} + + +def _copy_out_statement( + schema_name: str, + table_name: str, + columns: list[str], + *, + source_geometry_srids: dict[str, int], + target_geometry_srids: dict[str, int], +) -> sql.Composed: + relation = sql.Identifier(schema_name, table_name) + column_list = sql.SQL(", ").join(map(sql.Identifier, columns)) + srid_changes = { + column_name: target_geometry_srids[column_name] + for column_name, source_srid in source_geometry_srids.items() + if target_geometry_srids[column_name] != source_srid + } + if not srid_changes: + return sql.SQL("copy {} ({}) to stdout").format(relation, column_list) + + select_list = sql.SQL(", ").join( + sql.SQL("st_setsrid({}, {}) as {}").format( + sql.Identifier(column_name), + sql.Literal(srid_changes[column_name]), + sql.Identifier(column_name), + ) + if column_name in srid_changes + else sql.Identifier(column_name) + for column_name in columns + ) + return sql.SQL("copy (select {} from {}) to stdout").format( + select_list, + relation, + ) + + def _external_references(conn: Connection) -> set[tuple[str, str, str, str]]: with conn.cursor() as cur: cur.execute( @@ -128,7 +178,20 @@ def _copy_table( ) -> None: relation = sql.Identifier(schema_name, table_name) column_list = sql.SQL(", ").join(map(sql.Identifier, columns)) - copy_out = sql.SQL("copy {} ({}) to stdout").format(relation, column_list) + source_geometry_srids = _geometry_srids(source_conn, schema_name, table_name) + target_geometry_srids = _geometry_srids(target_conn, schema_name, table_name) + if source_geometry_srids.keys() != target_geometry_srids.keys(): + raise RuntimeError( + "Source/target geometry columns differ for " + f"{schema_name}.{table_name}" + ) + copy_out = _copy_out_statement( + schema_name, + table_name, + columns, + source_geometry_srids=source_geometry_srids, + target_geometry_srids=target_geometry_srids, + ) copy_in = sql.SQL("copy {} ({}) from stdin").format(relation, column_list) with source_conn.cursor().copy(copy_out) as source_copy: with target_conn.cursor().copy(copy_in) as target_copy: @@ -136,6 +199,47 @@ def _copy_table( target_copy.write(chunk) +def copy_network_tables(source_project: str, target_project: str) -> None: + """Copy only the immutable simulation input tables into a project template.""" + with project_connection(source_project) as source_conn, source_conn.transaction(): + with source_conn.cursor() as cur: + cur.execute("set transaction isolation level repeatable read, read only") + source_tables = [ + table for table in _model_tables(source_conn) if table[0] == "network" + ] + source_columns = { + table: _table_columns(source_conn, *table) for table in source_tables + } + copy_order = _copy_order(source_conn, source_tables) + + with project_transaction(target_project) as target_conn: + target_tables = [ + table for table in _model_tables(target_conn) if table[0] == "network" + ] + if set(target_tables) != set(source_tables): + raise RuntimeError("Source/target network schemas differ") + for table, columns in source_columns.items(): + if _table_columns(target_conn, *table) != columns: + raise RuntimeError( + f"Source/target columns differ for {table[0]}.{table[1]}" + ) + with target_conn.cursor() as cur: + for schema_name, table_name in reversed(copy_order): + cur.execute( + sql.SQL("delete from {}").format( + sql.Identifier(schema_name, table_name) + ) + ) + for schema_name, table_name in copy_order: + _copy_table( + source_conn, + target_conn, + schema_name, + table_name, + source_columns[(schema_name, table_name)], + ) + + def replace_project_model( target_project: str, source_project: str, diff --git a/app/native/wndb/core/project_templates.py b/app/native/wndb/core/project_templates.py new file mode 100644 index 0000000..6ad4c5d --- /dev/null +++ b/app/native/wndb/core/project_templates.py @@ -0,0 +1,191 @@ +from __future__ import annotations + +import time + +from psycopg import sql +from psycopg.conninfo import make_conninfo + +from app.core.config import settings +from app.infra.db.project_routing import get_project_template_database_name + +from .connection import admin_connection, close_project_pool, project_connection +from .model_replace import copy_network_tables +from .projects import ( + _ensure_project_model_template_ready, + copy_project, + delete_project, + have_project, +) + + +PUBLICATION_NAME = "wndb_network_pub" +SUBSCRIPTION_NAME = "wndb_network_sub" + + +def _slot_name(project: str) -> str: + return f"{project}_network_slot" + + +def _publisher_conninfo(project: str) -> str: + return make_conninfo( + dbname=project, + host=settings.DB_HOST, + port=settings.DB_PORT, + user=settings.DB_USER, + password=settings.DB_PASSWORD, + ) + + +def ensure_replication_worker_capacity() -> None: + """Keep one worker slot spare while admitting one new project subscription.""" + with admin_connection() as conn, conn.cursor() as cur: + cur.execute( + """ + select current_setting('max_worker_processes')::integer as maximum, + count(*) filter ( + where backend_type in ( + 'logical replication launcher', + 'logical replication worker' + ) + ) as used + from pg_stat_activity + """ + ) + row = cur.fetchone() + maximum = int(row["maximum"]) + used = int(row["used"]) + if used + 1 >= maximum: + raise RuntimeError( + "PostgreSQL has no reserved logical-replication worker capacity " + f"(used={used}, max_worker_processes={maximum}); increase " + "max_worker_processes before provisioning another project" + ) + + +def _create_publication_and_slot(project: str) -> None: + slot_name = _slot_name(project) + with project_connection(project) as conn, conn.cursor() as cur: + cur.execute( + "select n.nspname as schema_name, c.relname as table_name " + "from pg_class c join pg_namespace n on n.oid = c.relnamespace " + "where n.nspname = 'network' and c.relkind in ('r', 'p') " + "order by c.relname" + ) + tables = [(row["schema_name"], row["table_name"]) for row in cur.fetchall()] + if not tables: + raise RuntimeError("Project database has no network tables to publish") + cur.execute( + sql.SQL("create publication {} for table {}").format( + sql.Identifier(PUBLICATION_NAME), + sql.SQL(", ").join( + sql.Identifier(schema_name, table_name) + for schema_name, table_name in tables + ), + ) + ) + cur.execute( + "select slot_name from pg_create_logical_replication_slot(%s, 'pgoutput')", + (slot_name,), + ) + + +def _create_subscription(project: str, template: str) -> None: + with project_connection(template) as conn, conn.cursor() as cur: + cur.execute( + sql.SQL( + "create subscription {} connection {} publication {} " + "with (create_slot = false, copy_data = false, slot_name = {})" + ).format( + sql.Identifier(SUBSCRIPTION_NAME), + sql.Literal(_publisher_conninfo(project)), + sql.Identifier(PUBLICATION_NAME), + sql.Literal(_slot_name(project)), + ) + ) + + +def _wait_until_ready(template: str) -> None: + last_error: RuntimeError | None = None + for _ in range(50): + try: + _ensure_project_model_template_ready(template) + return + except RuntimeError as exc: + last_error = exc + time.sleep(0.1) + raise last_error or RuntimeError( + f"Project model template {template!r} did not become ready" + ) + + +def create_project_model_template(project: str) -> str: + template = get_project_template_database_name(project) + if have_project(template): + raise ValueError(f"Project model template {template!r} already exists") + template_created = False + try: + _create_publication_and_slot(project) + copy_project( + settings.WNDB_SCHEMA_TEMPLATE_DB_NAME, + template, + allow_template_source=True, + allow_template_target=True, + ) + template_created = True + copy_network_tables(project, template) + _create_subscription(project, template) + _wait_until_ready(template) + return template + except Exception: + if template_created: + delete_project_model_template(project) + else: + _drop_publisher_objects(project) + raise + + +def _drop_subscription(template: str) -> None: + if not have_project(template): + return + try: + with project_connection(template) as conn, conn.cursor() as cur: + cur.execute( + sql.SQL("drop subscription if exists {}").format( + sql.Identifier(SUBSCRIPTION_NAME) + ) + ) + finally: + close_project_pool(template) + + +def _drop_publisher_objects( + project: str, + *, + drop_publication: bool = True, + drop_slot: bool = True, +) -> None: + if not have_project(project): + return + with project_connection(project) as conn, conn.cursor() as cur: + if drop_slot: + cur.execute( + "select pg_drop_replication_slot(slot_name) " + "from pg_replication_slots where slot_name = %s and active = false", + (_slot_name(project),), + ) + if drop_publication: + cur.execute( + sql.SQL("drop publication if exists {}").format( + sql.Identifier(PUBLICATION_NAME) + ) + ) + + +def delete_project_model_template(project: str) -> None: + template = get_project_template_database_name(project) + try: + _drop_subscription(template) + finally: + _drop_publisher_objects(project) + if have_project(template): + delete_project(template, allow_template=True) diff --git a/app/native/wndb/core/projects.py b/app/native/wndb/core/projects.py index 0526796..0b0ffcb 100644 --- a/app/native/wndb/core/projects.py +++ b/app/native/wndb/core/projects.py @@ -9,22 +9,24 @@ from psycopg.rows import dict_row from app.core.config import settings from app.infra.db.project_routing import ( get_project_database_name, + get_schema_template_database_name, get_project_template_database_name, ) from .connection import ( admin_connection, close_project_pool, + project_connection, ) -_SERVER_DATABASES = frozenset({"template0", "template1", "postgres", "project"}) +_SERVER_DATABASES = frozenset({"template0", "template1", "postgres"}) _TEMPORARY_DATABASE_PREFIX = "tjw_tmp_" def _protected_databases() -> frozenset[str]: return _SERVER_DATABASES | { settings.METADATA_DB_NAME, - settings.WNDB_TEMPLATE_DB_NAME, + settings.WNDB_SCHEMA_TEMPLATE_DB_NAME, } @@ -34,10 +36,7 @@ def _validate_project_database(name: str, *, allow_template_source: bool = False 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() - ): + if allow_template_source and is_template: return if name.casefold() in protected or is_template: raise ValueError(f"Database {name!r} is protected and cannot be managed as a project") @@ -120,6 +119,72 @@ def _set_database_connections(cur, database_name: str, *, allowed: bool) -> None ) +def _is_project_model_template(database_name: str) -> bool: + return ( + database_name.casefold().endswith("_template") + and database_name.casefold() + != get_schema_template_database_name().casefold() + ) + + +def _ensure_project_model_template_ready(database_name: str) -> None: + """Reject cloning a missing, disabled, or initially syncing subscription.""" + if not _is_project_model_template(database_name): + return + + try: + with project_connection(database_name) as conn: + with conn.cursor() as cur: + cur.execute( + """ + select count(distinct subscription_row.oid) as subscriptions, + count(distinct subscription_row.oid) + filter (where subscription_row.subenabled) + as enabled_subscriptions, + count(subscription_status.pid) + filter (where subscription_status.pid is not null) + as active_workers + from pg_subscription subscription_row + left join pg_stat_subscription subscription_status + on subscription_status.subid = subscription_row.oid + where subscription_row.subdbid = ( + select oid from pg_database + where datname = current_database() + ) + """ + ) + status = cur.fetchone() + cur.execute( + """ + select count(*) as relations, + count(*) filter (where srsubstate <> 'r') + as pending_relations + from pg_subscription_rel + """ + ) + relations = cur.fetchone() + finally: + close_project_pool(database_name) + + if ( + status is None + or int(status["subscriptions"]) != 1 + or int(status["enabled_subscriptions"]) != 1 + or int(status["active_workers"]) < 1 + ): + raise RuntimeError( + f"Project model template {database_name!r} subscription is not active" + ) + if ( + relations is None + or int(relations["relations"]) < 1 + or int(relations["pending_relations"]) > 0 + ): + raise RuntimeError( + f"Project model template {database_name!r} is still synchronizing" + ) + + 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) @@ -137,18 +202,11 @@ 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( + copy_project( + get_project_template_database_name(project), temporary_name, - project, - copy_source_scada=True, + allow_template_source=True, ) - refresh_materialized_views_after_commit(temporary_name) yield temporary_name finally: if have_project(temporary_name): @@ -160,7 +218,11 @@ 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) + copy_project( + get_schema_template_database_name(), + temporary_name, + allow_template_source=True, + ) yield temporary_name finally: if have_project(temporary_name): @@ -175,11 +237,28 @@ def have_project(name: str) -> bool: return cur.fetchone() is not None -def copy_project(source: str, new: str) -> None: +def copy_project( + source: str, + new: str, + *, + allow_template_source: bool = False, + allow_template_target: bool = False, +) -> 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) + _validate_project_database( + physical_source, + allow_template_source=allow_template_source, + ) + if physical_new.casefold() == get_schema_template_database_name().casefold(): + raise ValueError( + f"Database {physical_new!r} is the protected schema template" + ) + _validate_project_database( + physical_new, + allow_template_source=allow_template_target, + ) + _ensure_project_model_template_ready(physical_source) close_project_pool(source) close_project_pool(new) @@ -216,12 +295,23 @@ def copy_project(source: str, new: str) -> None: def create_project(name: str) -> None: - return copy_project(get_project_template_database_name(name), name) + return copy_project( + get_schema_template_database_name(), + name, + allow_template_source=True, + ) -def delete_project(name: str) -> None: +def delete_project(name: str, *, allow_template: bool = False) -> None: database_name = get_project_database_name(name) - _validate_project_database(database_name) + if database_name.casefold() == get_schema_template_database_name().casefold(): + raise ValueError( + f"Database {database_name!r} is the protected schema template" + ) + _validate_project_database( + database_name, + allow_template_source=allow_template, + ) close_project_pool(name) with admin_connection() as conn: with conn.cursor() as cur: diff --git a/app/native/wndb/inp/importer.py b/app/native/wndb/inp/importer.py index 1bff391..cf062d6 100644 --- a/app/native/wndb/inp/importer.py +++ b/app/native/wndb/inp/importer.py @@ -12,7 +12,7 @@ from ..core.projects import ( temporary_project_name, temporary_template_database, ) -from app.infra.db.project_routing import get_project_template_database_name +from app.infra.db.project_routing import get_schema_template_database_name from ..core.connection import project_transaction from ..core.model_replace import replace_project_model from ..core.database import ( @@ -404,7 +404,11 @@ def read_inp(project: str, inp: str, version: str = "3") -> bool: staging_project = temporary_project_name(project, "model_import") replacement_committed = False try: - copy_project(get_project_template_database_name(project), staging_project) + copy_project( + get_schema_template_database_name(), + staging_project, + allow_template_source=True, + ) with project_transaction(staging_project): parse_file(staging_project, inp, version) replace_project_model(project, staging_project) diff --git a/app/native/wndb/model/options_v3.py b/app/native/wndb/model/options_v3.py index 9e00b8a..d638bfe 100644 --- a/app/native/wndb/model/options_v3.py +++ b/app/native/wndb/model/options_v3.py @@ -5,7 +5,10 @@ from .options import get_option_schema, get_option_v3_schema, generate_v2, gener def _parse_v2(v2_lines: list[str]) -> dict[str, str]: cs_v2 = g_update_prefix | { 'type' : 'option' } for s in v2_lines: - tokens = s.split() + stripped = s.strip() + if not stripped or stripped.startswith(';'): + continue + tokens = stripped.split() if tokens[0].upper() == 'PATTERN': # can not upper id value = tokens[1] if len(tokens) > 1 else '' cs_v2 |= { 'PATTERN' : value } @@ -23,6 +26,22 @@ def _parse_v2(v2_lines: list[str]) -> dict[str, str]: return cs_v2 +def _option_changes(option_type: str, *change_sets: ChangeSet) -> ChangeSet: + values: dict[str, str] = {} + for change_set in change_sets: + for operation in change_set.operations: + values.update( + { + key: str(value) + for key, value in operation.items() + if key not in {'operation', 'type'} + } + ) + if not values: + return ChangeSet() + return ChangeSet(g_update_prefix | {'type': option_type} | values) + + def _inp_in_option_v3(section: list[str]) -> ChangeSet: if len(section) <= 0: return ChangeSet() @@ -30,10 +49,11 @@ def _inp_in_option_v3(section: list[str]) -> ChangeSet: cs_v3 = g_update_prefix | { 'type' : 'option_v3' } v2_lines = [] for s in section: - if s.startswith(';'): + stripped = s.strip() + if not stripped or stripped.startswith(';'): continue - tokens = s.strip().split() + tokens = stripped.split() key = tokens[0] if key in get_option_v3_schema('').keys(): value = '' @@ -43,14 +63,17 @@ def _inp_in_option_v3(section: list[str]) -> ChangeSet: value = ' '.join(tokens[1:]) cs_v3 |= { key : value } else: - v2_lines.append(s.strip()) + v2_lines.append(stripped) - # unlikely... cs_v2 = _parse_v2(v2_lines) + direct_v2 = _option_changes('option', ChangeSet(cs_v2)) + direct_v3 = _option_changes('option_v3', ChangeSet(cs_v3)) + generated_v2 = generate_v2(direct_v3) if direct_v3.operations else ChangeSet() + generated_v3 = generate_v3(direct_v2) if direct_v2.operations else ChangeSet() - result = ChangeSet(cs_v3) - result.merge(generate_v3(ChangeSet(cs_v2))) - result.merge(generate_v2(result)) + result = ChangeSet() + result.merge(_option_changes('option', generated_v2, direct_v2)) + result.merge(_option_changes('option_v3', generated_v3, direct_v3)) return result diff --git a/app/services/burst_location.py b/app/services/burst_location.py index cab29f4..5d9c5d4 100644 --- a/app/services/burst_location.py +++ b/app/services/burst_location.py @@ -1,6 +1,5 @@ from __future__ import annotations -import os from datetime import datetime, timedelta from typing import Any from uuid import UUID @@ -10,7 +9,7 @@ import pandas as pd from app.algorithms.burst_localization import run_burst_location from app.infra.db.postgresql.scada import get_all_scada_info from app.infra.db.timescaledb.internal_queries import InternalQueries -from app.native.wndb.inp.exporter import dump_inp +from app.services.project_inp import temporary_project_inp from app.services.scheme_management import ( get_analysis_run, store_scheme_info, @@ -300,20 +299,20 @@ def run_burst_location_by_network( burst_flow_samples = 1 if burst_flow_series is not None else 0 normal_flow_samples = 1 if normal_flow_series is not None else 0 - inp_path = _prepare_burst_inp(network) - result = run_burst_location( - wn_inp_path=inp_path, - pressure_scada_ids=selected_pressure_ids, - burst_pressure=burst_pressure_series, - normal_pressure=normal_pressure_series, - burst_leakage=burst_leakage, - flow_scada_ids=selected_flow_ids, - burst_flow=burst_flow_series, - normal_flow=normal_flow_series, - min_dpressure=min_dpressure, - basic_pressure=basic_pressure, - visualize_partition=False, - ) + with temporary_project_inp(network, purpose="burst-location") as inp_path: + result = run_burst_location( + wn_inp_path=str(inp_path), + pressure_scada_ids=selected_pressure_ids, + burst_pressure=burst_pressure_series, + normal_pressure=normal_pressure_series, + burst_leakage=burst_leakage, + flow_scada_ids=selected_flow_ids, + burst_flow=burst_flow_series, + normal_flow=normal_flow_series, + min_dpressure=min_dpressure, + basic_pressure=basic_pressure, + visualize_partition=False, + ) payload: dict[str, Any] = { **result, @@ -775,16 +774,3 @@ def _normalize_timeseries_by_id( def _to_datetime(value: datetime | str) -> datetime: return parse_utc_time(value) - - -def _prepare_burst_inp(network: str) -> str: - project_root = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..")) - db_inp_dir = os.path.join(project_root, "db_inp") - os.makedirs(db_inp_dir, exist_ok=True) - inp_path = os.path.join(db_inp_dir, f"{network}.burst.inp") - if os.path.isfile(inp_path) and os.path.getsize(inp_path) > 0: - return inp_path - dump_inp(network, inp_path, "2") - if not os.path.isfile(inp_path) or os.path.getsize(inp_path) <= 0: - raise ValueError(f"爆管定位 INP 文件无效: {inp_path}") - return inp_path diff --git a/app/services/dma_leakage_estimation.py b/app/services/dma_leakage_estimation.py index 8110208..b350212 100644 --- a/app/services/dma_leakage_estimation.py +++ b/app/services/dma_leakage_estimation.py @@ -14,7 +14,7 @@ from app.native.wndb.gis.network_views import ( get_network_link_nodes, get_network_node_coords, ) -from app.native.wndb.inp.exporter import dump_inp +from app.services.project_inp import temporary_project_inp from app.services.scheme_management import store_analysis_run_with_result from app.domain.time import parse_utc_time, utc_now @@ -42,8 +42,6 @@ def run_leakage_identification( sensor_nodes: list[str] | None = None, scheme_name: str | None = None, ) -> dict[str, Any]: - inp_path = _prepare_leakage_inp(network) - selected_sensor_nodes = ( list(dict.fromkeys([node for node in (sensor_nodes or []) if node])) if sensor_nodes @@ -52,7 +50,7 @@ def run_leakage_identification( if not selected_sensor_nodes: raise ValueError("未提供有效传感器节点,且系统未识别到可用压力传感器。") - area_map, areas, node_coords = _build_area_map_by_topology( + area_map, areas, _ = _build_area_map_by_topology( network, selected_sensor_nodes, dma_count ) @@ -73,26 +71,25 @@ def run_leakage_identification( observed_df = observed_pressure_data q_sum_m3s = DmaLeakageOptimizer._flow_to_m3s(q_sum, q_sum_unit) - identifier = DmaLeakageOptimizer( - inp_path=inp_path, - sensor_nodes=selected_sensor_nodes, - area_map=area_map, - start_time=start_time, - duration=duration, - timestep=timestep, - q_sum=q_sum_m3s, - ) - result_df = identifier.run_identification( - observed_pressure_data=observed_df, - pop_size=pop_size, - max_gen=max_gen, - n_workers=n_workers, - output_flow_unit=output_flow_unit, - save_result=False, - ) + with temporary_project_inp(network, purpose="dma-leakage") as inp_path: + identifier = DmaLeakageOptimizer( + inp_path=str(inp_path), + sensor_nodes=selected_sensor_nodes, + area_map=area_map, + start_time=start_time, + duration=duration, + timestep=timestep, + q_sum=q_sum_m3s, + ) + result_df = identifier.run_identification( + observed_pressure_data=observed_df, + pop_size=pop_size, + max_gen=max_gen, + n_workers=n_workers, + output_flow_unit=output_flow_unit, + save_result=False, + ) rows = result_df.to_dict(orient="records") - # node_visual_payload = _build_node_visual_payload(area_map, node_coords, rows) - # drawing_payload = _build_drawing_payload(node_visual_payload) payload = { "result_path": result_df.attrs.get("result_path"), "sensor_nodes": selected_sensor_nodes, @@ -100,8 +97,6 @@ def run_leakage_identification( "area_count": len(set(area_map.values())), "node_area_map": area_map, "areas": areas, - # "node_visual_payload": node_visual_payload, - # "drawing_payload": drawing_payload, "rows": rows, } if scheme_name: @@ -326,16 +321,3 @@ def _build_observed_pressure_from_scada( def _to_datetime(value: datetime | str) -> datetime: return parse_utc_time(value) - - -def _prepare_leakage_inp(network: str) -> str: - project_root = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..")) - db_inp_dir = os.path.join(project_root, "db_inp") - os.makedirs(db_inp_dir, exist_ok=True) - inp_path = os.path.join(db_inp_dir, f"{network}.leakage.inp") - if os.path.isfile(inp_path) and os.path.getsize(inp_path) > 0: - return inp_path - dump_inp(network, inp_path, "2") - if not os.path.isfile(inp_path) or os.path.getsize(inp_path) <= 0: - raise ValueError(f"漏损识别 INP 文件无效: {inp_path}") - return inp_path diff --git a/app/services/project_inp.py b/app/services/project_inp.py new file mode 100644 index 0000000..78fb5fd --- /dev/null +++ b/app/services/project_inp.py @@ -0,0 +1,35 @@ +from collections.abc import Iterator +from contextlib import contextmanager +from pathlib import Path +from tempfile import NamedTemporaryFile + +from app.native.wndb.inp.exporter import dump_inp + + +PROJECT_INP_DIRECTORY = Path("db_inp") + + +@contextmanager +def temporary_project_inp( + project_code: str, + *, + purpose: str, + version: str = "2", +) -> Iterator[Path]: + """Export one project model to an isolated INP and remove it afterwards.""" + PROJECT_INP_DIRECTORY.mkdir(parents=True, exist_ok=True) + with NamedTemporaryFile( + dir=PROJECT_INP_DIRECTORY, + prefix=f"{purpose}_", + suffix=".inp", + delete=False, + ) as temporary_file: + path = Path(temporary_file.name).resolve() + + try: + dump_inp(project_code, str(path), version) + if not path.is_file() or path.stat().st_size == 0: + raise ValueError(f"项目 {project_code!r} 的 INP 导出失败") + yield path + finally: + path.unlink(missing_ok=True) diff --git a/app/services/project_provisioning.py b/app/services/project_provisioning.py new file mode 100644 index 0000000..d416c6e --- /dev/null +++ b/app/services/project_provisioning.py @@ -0,0 +1,216 @@ +from __future__ import annotations + +from dataclasses import dataclass +import logging +from pathlib import Path +import re +from urllib.parse import quote + +from app.core.config import settings +from app.infra.db.timescaledb.lifecycle import ( + create_timescale_database, + delete_timescale_database, + require_timescale_schema_template, + timescale_database_exists, +) +from app.infra.db.postgresql.project_spatial import get_project_map_bbox +from app.infra.geoserver.client import GeoServerAdminClient +from app.native.wndb.core.project_templates import ( + create_project_model_template, + delete_project_model_template, + ensure_replication_worker_capacity, +) +from app.native.wndb.core.projects import create_project, delete_project, have_project +from app.services.network_import import network_update + + +logger = logging.getLogger(__name__) +_PROJECT_CODE = re.compile(r"^[a-z][a-z0-9_]{0,49}$") + + +class ProjectProvisioningError(RuntimeError): + def __init__(self, stage: str, cause: Exception, cleanup_errors: list[str]) -> None: + self.stage = stage + self.cause = cause + self.cleanup_errors = cleanup_errors + cleanup_suffix = ( + f"; cleanup failures: {', '.join(cleanup_errors)}" + if cleanup_errors + else "" + ) + super().__init__(f"Project provisioning failed at {stage}: {cause}{cleanup_suffix}") + + +def validate_project_code(code: str) -> str: + if not _PROJECT_CODE.fullmatch(code): + raise ValueError( + "Project code must start with a lowercase letter and contain only " + "lowercase letters, digits, and underscores" + ) + if code.endswith("_template"): + raise ValueError("Project code must not end with '_template'") + if code in { + "postgres", + "template0", + "template1", + settings.METADATA_DB_NAME.casefold(), + }: + raise ValueError(f"Project code {code!r} is reserved") + return code + + +def _database_url(*, timescale: bool, database_name: str) -> str: + if timescale: + host = settings.TIMESCALEDB_DB_HOST + port = settings.TIMESCALEDB_DB_PORT + user = settings.TIMESCALEDB_DB_USER + password = settings.TIMESCALEDB_DB_PASSWORD + else: + host = settings.DB_HOST + port = settings.DB_PORT + user = settings.DB_USER + password = settings.DB_PASSWORD + return ( + f"postgresql://{quote(user, safe='')}:{quote(password, safe='')}" + f"@{host}:{port}/{database_name}" + ) + + +@dataclass +class ProvisionedProjectInfrastructure: + code: str + workspace: str + model_template: str + map_bbox: tuple[float, float, float, float] + layers: tuple[str, ...] + + @property + def business_dsn(self) -> str: + return _database_url(timescale=False, database_name=self.code) + + @property + def timescale_dsn(self) -> str: + return _database_url(timescale=True, database_name=self.code) + + def cleanup(self) -> list[str]: + errors: list[str] = [] + cleanup_steps = ( + ("geoserver", self._delete_geoserver), + ("timescaledb", lambda: delete_timescale_database(self.code)), + ("model_template", lambda: delete_project_model_template(self.code)), + ("business_database", lambda: delete_project(self.code)), + ) + for name, cleanup in cleanup_steps: + try: + cleanup() + except Exception as exc: # preserve every cleanup attempt + logger.exception("Project provisioning cleanup failed at %s", name) + errors.append(name) + return errors + + def _delete_geoserver(self) -> None: + with GeoServerAdminClient() as geoserver: + geoserver.delete_workspace(self.workspace) + + +def _preflight(code: str, workspace: str, geoserver: GeoServerAdminClient) -> None: + validate_project_code(code) + if not _PROJECT_CODE.fullmatch(workspace): + raise ValueError( + "GeoServer workspace must start with a lowercase letter and contain " + "only lowercase letters, digits, and underscores" + ) + if have_project(settings.WNDB_SCHEMA_TEMPLATE_DB_NAME) is False: + raise RuntimeError( + f"Business schema template {settings.WNDB_SCHEMA_TEMPLATE_DB_NAME!r} does not exist" + ) + require_timescale_schema_template() + ensure_replication_worker_capacity() + if have_project(code): + raise ValueError(f"Business database {code!r} already exists") + model_template = f"{code}_template" + if have_project(model_template): + raise ValueError(f"Project model template {model_template!r} already exists") + if timescale_database_exists(code): + raise ValueError(f"TimescaleDB database {code!r} already exists") + geoserver.check_ready() + if geoserver.workspace_exists(workspace): + raise ValueError(f"GeoServer workspace {workspace!r} already exists") + + +def provision_project_infrastructure( + *, + code: str, + workspace: str, + inp_path: str | Path, +) -> ProvisionedProjectInfrastructure: + stage = "preflight" + business_created = False + template_attempted = False + timescale_created = False + workspace_attempted = False + map_bbox: tuple[float, float, float, float] | None = None + layers: tuple[str, ...] = () + try: + with GeoServerAdminClient() as geoserver: + _preflight(code, workspace, geoserver) + stage = "business_database" + create_project(code) + business_created = True + + stage = "model_import" + network_update(str(inp_path), code) + map_bbox = get_project_map_bbox(code) + + stage = "model_template" + template_attempted = True + model_template = create_project_model_template(code) + + stage = "timescaledb" + create_timescale_database(code) + timescale_created = True + + stage = "geoserver" + workspace_attempted = True + layers = geoserver.create_project_workspace( + workspace=workspace, + database_name=code, + map_bbox=map_bbox, + ) + except Exception as exc: + cleanup_errors: list[str] = [] + if workspace_attempted: + try: + with GeoServerAdminClient() as geoserver: + geoserver.delete_workspace(workspace) + except Exception: + logger.exception("Failed to remove GeoServer workspace %s", workspace) + cleanup_errors.append("geoserver") + if timescale_created: + try: + delete_timescale_database(code) + except Exception: + logger.exception("Failed to remove TimescaleDB database %s", code) + cleanup_errors.append("timescaledb") + if template_attempted: + try: + delete_project_model_template(code) + except Exception: + logger.exception("Failed to remove project template for %s", code) + cleanup_errors.append("model_template") + if business_created: + try: + delete_project(code) + except Exception: + logger.exception("Failed to remove business database %s", code) + cleanup_errors.append("business_database") + raise ProjectProvisioningError(stage, exc, cleanup_errors) from exc + + assert map_bbox is not None + return ProvisionedProjectInfrastructure( + code=code, + workspace=workspace, + model_template=f"{code}_template", + map_bbox=map_bbox, + layers=layers, + ) diff --git a/app/services/sensor_placement.py b/app/services/sensor_placement.py index e2046f5..9472431 100644 --- a/app/services/sensor_placement.py +++ b/app/services/sensor_placement.py @@ -1,3 +1,4 @@ +from collections.abc import Iterator from contextlib import contextmanager from datetime import datetime import fcntl @@ -16,7 +17,7 @@ import wntr from app.algorithms.pressure_sensor_placement import kmeans_placement from app.algorithms.pressure_sensor_placement import sensitivity_placement from app.infra.db.postgresql import sensor_placement as sensor_placement_repository -from app.native.wndb.inp.exporter import dump_inp +from app.services.project_inp import temporary_project_inp class SensorPlacementNotFoundError(LookupError): @@ -31,7 +32,7 @@ class SensorPlacementConflictError(RuntimeError): pass -def _sensor_inp_path(project_code: str) -> Path: +def _sensor_lock_path(project_code: str) -> Path: if ( not project_code or project_code in {".", ".."} @@ -40,14 +41,13 @@ def _sensor_inp_path(project_code: str) -> Path: or "\x00" in project_code ): raise SensorPlacementValidationError("管网名称不是有效的项目标识") - return Path("db_inp") / f"{project_code}.db.inp" + return Path("db_inp") / f"{project_code}.sensor.lock" @contextmanager -def _sensor_inp_lock(project_code: str): - inp_path = _sensor_inp_path(project_code) - inp_path.parent.mkdir(parents=True, exist_ok=True) - lock_path = inp_path.with_suffix(".sensor.lock") +def _sensor_run_lock(project_code: str) -> Iterator[None]: + lock_path = _sensor_lock_path(project_code) + lock_path.parent.mkdir(parents=True, exist_ok=True) with lock_path.open("w", encoding="utf-8") as lock_file: try: fcntl.flock(lock_file.fileno(), fcntl.LOCK_EX | fcntl.LOCK_NB) @@ -56,11 +56,23 @@ def _sensor_inp_lock(project_code: str): "当前项目已有监测点优化任务正在运行,请稍后重试" ) from exc try: - yield inp_path + yield finally: fcntl.flock(lock_file.fileno(), fcntl.LOCK_UN) +@contextmanager +def _sensor_network_model( + project_code: str, +) -> Iterator[wntr.network.WaterNetworkModel]: + with _sensor_run_lock(project_code): + with temporary_project_inp( + project_code, + purpose="sensor-placement", + ) as inp_path: + yield wntr.network.WaterNetworkModel(str(inp_path)) + + def _create_validated_placement( project_code: str, *, @@ -88,10 +100,7 @@ def optimize_sensor_placement_by_sensitivity( ) -> dict[str, Any]: """Run sensitivity placement and persist the validated result.""" - with _sensor_inp_lock(project_code): - inp_path = _sensor_inp_path(project_code) - dump_inp(project_code, str(inp_path), "2") - network_model = wntr.network.WaterNetworkModel(str(inp_path)) + with _sensor_network_model(project_code) as network_model: sensor_locations = sensitivity_placement.optimize_sensor_placement( network_model, sensor_num=sensor_count, @@ -115,9 +124,7 @@ def optimize_sensor_placement_by_kmeans( ) -> dict[str, Any]: """Export the model, run K-means placement, and persist the result.""" - with _sensor_inp_lock(project_code) as inp_path: - dump_inp(project_code, str(inp_path), "2") - network_model = wntr.network.WaterNetworkModel(str(inp_path)) + with _sensor_network_model(project_code) as network_model: sensor_locations = kmeans_placement.optimize_sensor_placement( network_model, sensor_count=sensor_count, diff --git a/contracts/manifest.json b/contracts/manifest.json index 159b90e..f704c5d 100644 --- a/contracts/manifest.json +++ b/contracts/manifest.json @@ -3,7 +3,7 @@ "contracts": { "server": { "file": "server-v1.openapi.json", - "sha256": "f07ecfa4843955bd63a4cdf0d7d336d593f41725a8abc1b0da6cce7ff4fff982" + "sha256": "fb720e3009948c3cb2f37674d0bbdbb7b6136973223eccca89880475697435bc" } } } diff --git a/contracts/server-v1.openapi.json b/contracts/server-v1.openapi.json index 8387923..47842b5 100644 --- a/contracts/server-v1.openapi.json +++ b/contracts/server-v1.openapi.json @@ -491,6 +491,65 @@ "title": "Body_patch_admin_projects_project_id_model_imports", "type": "object" }, + "Body_post_admin_project_provisions": { + "properties": { + "code": { + "maxLength": 50, + "minLength": 1, + "title": "Code", + "type": "string" + }, + "description": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Description" + }, + "file": { + "description": "EPANET INP 模型文件", + "format": "binary", + "title": "File", + "type": "string" + }, + "gs_workspace": { + "anyOf": [ + { + "maxLength": 100, + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Gs Workspace" + }, + "map_zoom": { + "default": 14, + "maximum": 22.0, + "minimum": 1.0, + "title": "Map Zoom", + "type": "integer" + }, + "name": { + "maxLength": 100, + "minLength": 1, + "title": "Name", + "type": "string" + } + }, + "required": [ + "name", + "code", + "file" + ], + "title": "Body_post_admin_project_provisions", + "type": "object" + }, "Body_post_admin_projects_project_id_model_imports": { "properties": { "file": { @@ -2014,6 +2073,46 @@ "title": "ProjectMetaResponse", "type": "object" }, + "ProjectProvisionResponse": { + "properties": { + "business_database": { + "title": "Business Database", + "type": "string" + }, + "geoserver_layers": { + "items": { + "type": "string" + }, + "title": "Geoserver Layers", + "type": "array" + }, + "geoserver_workspace": { + "title": "Geoserver Workspace", + "type": "string" + }, + "model_template_database": { + "title": "Model Template Database", + "type": "string" + }, + "project": { + "$ref": "#/components/schemas/AdminProjectResponse" + }, + "timescale_database": { + "title": "Timescale Database", + "type": "string" + } + }, + "required": [ + "project", + "business_database", + "model_template_database", + "timescale_database", + "geoserver_workspace", + "geoserver_layers" + ], + "title": "ProjectProvisionResponse", + "type": "object" + }, "ProjectSummaryResponse": { "properties": { "code": { @@ -2879,6 +2978,102 @@ ] } }, + "/api/v1/admin/project-provisions": { + "post": { + "operationId": "post_admin_project_provisions", + "requestBody": { + "content": { + "multipart/form-data": { + "schema": { + "$ref": "#/components/schemas/Body_post_admin_project_provisions" + } + } + }, + "required": true + }, + "responses": { + "201": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ProjectProvisionResponse" + } + } + }, + "description": "Successful Response" + }, + "401": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ProblemDetails" + } + } + }, + "description": "Authentication required" + }, + "403": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ProblemDetails" + } + } + }, + "description": "Insufficient permission" + }, + "404": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ProblemDetails" + } + } + }, + "description": "Resource not found" + }, + "409": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ProblemDetails" + } + } + }, + "description": "Resource conflict" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ProblemDetails" + } + } + }, + "description": "Validation error" + }, + "503": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ProblemDetails" + } + } + }, + "description": "Dependency unavailable" + } + }, + "security": [ + { + "OAuth2PasswordBearer": [] + } + ], + "summary": "创建完整供水项目", + "tags": [ + "Model Administration" + ] + } + }, "/api/v1/admin/projects": { "get": { "operationId": "get_admin_projects", @@ -2990,6 +3185,7 @@ ] }, "post": { + "description": "仅用于登记已经由外部流程完整创建的资源。新项目应调用 POST /admin/project-provisions。", "operationId": "post_admin_projects", "requestBody": { "content": { @@ -3078,7 +3274,7 @@ "OAuth2PasswordBearer": [] } ], - "summary": "Create Admin Project", + "summary": "仅登记已有项目元数据", "tags": [ "Metadata Admin" ] diff --git a/infra/docker/docker-compose.yml b/infra/docker/docker-compose.yml index 215ec73..23e5515 100644 --- a/infra/docker/docker-compose.yml +++ b/infra/docker/docker-compose.yml @@ -106,6 +106,18 @@ services: POSTGRES_DB: ${POSTGIS_DB_NAME} POSTGRES_USER: ${POSTGIS_USER} POSTGRES_PASSWORD: ${DB_PASSWORD_SHARED} + command: + - postgres + - -c + - wal_level=logical + - -c + - max_worker_processes=32 + - -c + - max_logical_replication_workers=24 + - -c + - max_replication_slots=24 + - -c + - max_wal_senders=24 ports: - "${POSTGIS_PORT}:5432" volumes: diff --git a/resources/db_v2/DATABASE_ARCHITECTURE.md b/resources/db_v2/DATABASE_ARCHITECTURE.md index 4185c96..3029f63 100644 --- a/resources/db_v2/DATABASE_ARCHITECTURE.md +++ b/resources/db_v2/DATABASE_ARCHITECTURE.md @@ -1,12 +1,12 @@ # TJWater 数据库改造说明与当前结构 -> 本文记录截至 2026-08-27 的数据库实际状态。结构、约束、行数、TimescaleDB chunk 和策略均直接读取数据库,不以仓库中的 SQL 脚本为依据。文中不包含主机、端口、账号、密码或 DSN。 +> 本文记录截至 2026-09-10 的数据库实际状态。结构、约束、行数、TimescaleDB chunk 和策略均直接读取数据库,不以仓库中的 SQL 脚本为依据。文中不包含主机、端口、账号、密码或 DSN。 ## 改造范围与当前状态 -本次改造保留原 `tjwater` 业务库和时序库,新建的隔离库完成验证后已正式重命名为 `tjwater_v2`。元数据库仍为 `system_hub`,逻辑项目 `tjwater_next` 通过 `biz_data` 和 `iot_data` 两条路由分别关联 `tjwater_v2` 业务库与时序库。项目已切换为 `active`,原数据库没有被覆盖,仍可用于对照和回退。 +本次改造保留原 `tjwater` 业务库和时序库,新建的隔离库完成验证后已正式重命名为 `tjwater_v2`。元数据库仍为 `system_hub`,逻辑项目 `tjwater_v2` 通过 `biz_data` 和 `iot_data` 两条路由分别关联 `tjwater_v2` 业务库与时序库。项目已切换为 `active`,原数据库没有被覆盖,仍可用于对照和回退。 -命名约定:正式物理库名为 `tjwater_v2`,WNDB 版本模板固定为 `tjwater_v2_template`。元数据库中的逻辑项目代码和 GeoServer 工作空间仍为 `tjwater_next`;这些是路由与图层限定名,不是物理库名。模板已创建并清空项目模型、SCADA、分析运行和物化视图数据,仅保留数据库结构、PostGIS 对象和必要配置键种子,普通连接已关闭。 +命名约定:正式物理库名和元数据项目代码均为 `tjwater_v2`,其管网镜像模板为 `tjwater_v2_template`;其他物理业务库同样使用同名 `_template`。项目模板通过逻辑订阅只同步 `network` schema。数据为空且禁止普通连接的 `tjwater_v2_schema_template` 仅用于创建业务库和 INP 暂存库。GeoServer 工作空间仍使用 `tjwater_next`,它是图层限定名,不是项目代码或物理库名。 已经完成的数据库修改包括: @@ -15,7 +15,7 @@ - `realtime` 采用冷热数据策略,72 小时后的 chunk 自动转为有序列存。 - `analysis` 按 `stored_at` 分区,入库满 24 小时的 chunk 自动转为有序列存。 - GIS 查询层改用物化视图,当前 7 张物化视图均已填充。 -- GeoServer 已建立 `tjwater_next` 工作空间和同名数据存储,数据存储连接 `tjwater_v2` 业务库的 `gis` schema 并发布 7 个图层。GeoWebCache 的服务端与客户端缓存有效期均为 300 秒。 +- GeoServer 已建立 `tjwater_next` 工作空间和同名数据存储,数据存储连接 `tjwater_v2` 业务库的 `gis` schema 并发布 7 个图层。全部 TJWater 业务工作空间的 GeoWebCache 客户端缓存有效期统一为 30 天(`2,592,000` 秒);服务端瓦片在 GIS 物化视图刷新或坐标修正后显式清理。 - 旧库中的 `operation`、`current_operation`、`batch_operation`、`operation_table`、`restore_operation` 和 `snapshot_operation` 没有进入新业务库。 - `system_hub.public` 补充了项目数据库外键、数据库路由约束、连接池约束、必要的非空约束,以及 5 张表和 44 个字段的中文数据库注释。 - 用户角色和项目角色仍是可扩展字符串,没有增加枚举检查约束。 @@ -25,14 +25,14 @@ - 后端批量元素查询读取 GIS 物化视图,模型增删改和 INP 导入提交后执行并发刷新;批量事务只刷新一次。 - `pattern_values`、`pattern_flow_samples`、`curve_points`、`demands` 和 `link_vertices` 的顺序号按所属父对象编号,主键已改为父对象 ID 与 `sequence_no` 的复合键。 -逻辑项目 `tjwater_next` 当前为 `active`。排水项目 `lingang` 已迁入 `system_hub.public`,供水和排水后端共用同一套项目、成员、数据库路由和审计表。 +逻辑项目 `tjwater_v2` 当前为 `active`。 ## 数据库总体关系 ```mermaid flowchart LR subgraph META["system_hub 元数据库"] - MP["public
统一元数据
5 个项目"] + MP["public
统一元数据
8 个项目"] end subgraph OLD["原数据库,保持不变"] @@ -50,9 +50,8 @@ flowchart LR MP -->|"tjwater 的 biz_data"| OB MP -->|"tjwater 的 iot_data"| OT - MP -->|"项目 tjwater_next 的 biz_data"| NB - MP -->|"项目 tjwater_next 的 iot_data"| NT - MP -->|"lingang 的两条数据库路由"| DRAIN["排水项目数据库"] + MP -->|"项目 tjwater_v2 的 biz_data"| NB + MP -->|"项目 tjwater_v2 的 iot_data"| NT NB -->|"gis 物化视图"| GS -->|"WFS / WMTS"| WEB ``` @@ -62,7 +61,7 @@ flowchart LR ### public:当前主元数据 -`public` 当前有 5 个项目、1 个用户、5 条成员关系和 10 条数据库路由。5 个项目均配置了一条 `biz_data` 和一条 `iot_data` 路由。审计日志数量会随接口请求持续增加,不在文档中固化行数。 +`public` 当前有 8 个项目(5 个 `active`、3 个 `inactive`)、1 个用户、5 条成员关系和 16 条数据库路由。8 个项目均配置了一条 `biz_data` 和一条 `iot_data` 路由。审计日志数量会随接口请求持续增加,不在文档中固化行数。 | 表 | 用途 | 主要字段 | | --- | --- | --- | @@ -136,7 +135,7 @@ erDiagram ### 排水元数据合并 -排水后端原先使用独立的 `hub` schema,其中有 1 个 `lingang` 项目、2 条数据库路由和 203 条审计记录,没有用户或成员关系。项目 UUID 保持不变,两条路由转换为 `public.project_databases` 使用的 Fernet 加密格式,连接池参数由 `pool_size + max_overflow` 映射为 `pool_min_size + pool_max_size`。原审计记录已迁入 `public.audit_logs`。 +排水后端原先使用独立的 `hub` schema。迁移校验后确认 `lingang` 项目属于误入数据,其项目和数据库路由已从当前 `public` 元数据中移除。 排水后端已改用 `public.projects`、`public.project_databases`、`public.user_project_membership`、`public.users` 和 `public.audit_logs`。供水和排水服务读取同一份项目状态、Keycloak 身份、项目权限及数据库路由。原 `hub` schema 已在迁移校验和连接测试通过后删除。 @@ -318,7 +317,33 @@ flowchart LR ### GeoServer 与前端图层 -`system_hub.public.projects` 中的 `tjwater_next` 项目已配置 `gs_workspace=tjwater_next`,当前状态为 `active`。GeoServer 的 `tjwater_next` 数据存储连接 `tjwater_v2` 业务库并限定到 `gis` schema,图层名称直接采用物化视图名称。7 个图层使用相同的项目管网发布边界,空图层和视口内没有要素的瓦片会返回空 MVT,不会产生越界错误。 +`system_hub.public.projects` 中的 `tjwater_v2` 项目已配置 `gs_workspace=tjwater_next`,当前状态为 `active`。GeoServer 的 `tjwater_next` 数据存储连接 `tjwater_v2` 业务库并限定到 `gis` schema,图层名称直接采用物化视图名称。7 个图层使用相同的项目管网发布边界,空图层和视口内没有要素的瓦片会返回空 MVT,不会产生越界错误。 + +GeoWebCache 对 TJWater 的 11 个业务工作空间、80 个已发布图层统一返回 `Cache-Control: max-age=2592000, must-revalidate`。GIS 数据更新后必须先刷新物化视图、重算 GeoServer 的 native/latLon 边界,再显式清理对应图层的服务端瓦片;30 天客户端缓存不会替代这套更新流程。 + +### 高铁湛江北站项目 + +`zjb` 是高铁湛江北站供水项目的逻辑项目代码、业务库、时序库、GeoServer 工作空间和数据存储名称。业务库已使用纠正版 `zhanjiangbei_water_network_pump_station_corrected.inp` 重新导入,当前有 272 个节点和 285 条连接,其中包括 6 台泵和 3 个阀门;`zjb_template` 通过逻辑订阅同步 32 张 `network` 表。时序库使用与 v2 项目一致的 `analysis`、`realtime`、`scada` 结构,目前数据为空。元数据库中项目状态为 `active`,`biz_data` 与 `iot_data` 均使用 1~4 条项目连接池。 + +源文件使用尚未取得测量定义的 DXF 工程坐标。`zjb.public.spatial_ref_sys` 将其登记为 `TJWater:990001`,原始几何数值保存在该 SRID 下;`gis.to_web_mercator` 根据四个 OSM 控制点执行二维四参数拟合,7 张 GeoServer 物化视图输出 `EPSG:3857`。控制点均方根残差约 4.62 米,适合在线地图展示,不可用于施工或测量放样,取得正式控制点后应重新拟合并刷新物化视图与瓦片。当前发布范围约为 `110.3513°E~110.3573°E、21.2185°N~21.2225°N`。 + +### 新项目供应工作流 + +后端通过 `POST /api/v1/admin/project-provisions` 接收项目元数据和 INP 文件。项目在全部外部资源通过校验后才写入元数据库并变为 `active`,因此普通项目列表不会看到创建中的半成品。 + +```mermaid +flowchart LR + INP["INP 校验"] --> BIZ["创建 BizDB
导入 network / gis"] + BIZ --> TMPL["创建 code_template
复制 network 并订阅"] + TMPL --> TS["从空模板创建 TimescaleDB"] + TS --> GS["创建 GeoServer
7 图层 + 30 天客户端缓存"] + GS --> META["元数据单事务
项目 + 2 路由 + 创建者成员"] + META --> ACTIVE["active"] +``` + +任一步失败都只清理由本次请求新建的资源,并按 GeoServer、TimescaleDB、项目模板、业务库逆序回滚。业务 PostgreSQL 必须为每个活动模板提供一个逻辑复制 worker;工作流在创建前保留一个 worker 余量,容量不足会在资源创建前失败。容器部署基线为 `max_worker_processes=32`、`max_logical_replication_workers=24`、`max_replication_slots=24` 和 `max_wal_senders=24`。旧版非活动项目的三个模板订阅已停用,不再占用 worker。 + +TimescaleDB 的空结构模板为 `tjwater_v2_timescale_template`,普通连接关闭;当前结构包含空的 `analysis`、`realtime`、`scada` 表、hypertable 与压缩策略。新时序库直接从该模板克隆,不从任一业务项目复制数据。 | 前端数据源 | GeoServer 图层 | 几何 | 当前要素数 | | --- | --- | --- | ---: | @@ -426,7 +451,7 @@ erDiagram 元数据库通过 SQLAlchemy 异步连接池访问;项目请求按 `system_hub.public.project_databases` 路由到业务库和时序库。异步业务查询和异步时序查询由项目级动态池管理,池条目记录借用数;配置更新建立新一代池,旧池只在已有借用归还后关闭。原生 WNDB 同步访问使用按数据库缓存且同样带借用计数的 `psycopg_pool.ConnectionPool`,数据库创建、复制和删除使用独立的 PostgreSQL 管理池及数据库级 advisory lock,同步 TimescaleDB 访问也使用按数据库缓存的连接池。应用目录中已没有直接调用 `psycopg.connect` 的业务代码。 -WNDB 批量修改和模拟参数准备在同一条池连接和同一事务中执行,提交后只刷新一次 GIS 物化视图。普通写入和 INP 整体替换使用同一项目级事务锁;INP 先在唯一暂存库校验,再从一致性快照事务替换当前库的 `network/gis` 模型表。临时分析库由固定模板建结构后复制当前项目模型、SCADA 映射并刷新视图。实时节点和连接结果在一个事务中执行整体先删后写,同一结果时间使用事务级锁;分析结果按 `run_id` 加事务级锁。Timescale 复合查询按节点/管段批量读取,SCADA 清洗使用单条集合更新,不再逐点往返。 +WNDB 批量修改和模拟参数准备在同一条池连接和同一事务中执行,提交后只刷新一次 GIS 物化视图。普通写入和 INP 整体替换使用同一项目级事务锁;INP 先在唯一暂存库校验,再从一致性快照事务替换当前库的 `network/gis` 模型表。临时分析库从当前项目的 `_template` 管网镜像克隆,只使用 `network` 数据修改参数、导出 INP 并运行 EPANET;GIS 和 SCADA 都不进入临时库。实时节点和连接结果在一个事务中执行整体先删后写,同一结果时间使用事务级锁;分析结果按 `run_id` 加事务级锁。Timescale 复合查询按节点/管段批量读取,SCADA 清洗使用单条集合更新,不再逐点往返。 自动化真实数据库测试分别执行 64 次业务库和 64 次时序库并发借用,查询结果一致,连接均能归还池中。嵌套 WNDB 写入和分析运行生命周期测试会在外层强制回滚,数据库没有残留记录。`DatabaseCommand` 的 pattern 新增、修改、删除也在同一池化事务中完成,并验证了五张明细表的复合主键、级联解除需求模式关联、结果变更和整体回滚。实时覆盖测试确认第二批数据替换第一批数据,外层回滚后测试记录为 0。 diff --git a/resources/db_v2/WNDB_STRUCTURE.md b/resources/db_v2/WNDB_STRUCTURE.md index 55b2ebd..e2a8c78 100644 --- a/resources/db_v2/WNDB_STRUCTURE.md +++ b/resources/db_v2/WNDB_STRUCTURE.md @@ -8,7 +8,9 @@ 本次调整了代码文件、导入关系、命令分派、项目生命周期接口和 WNDB 内部命令对象。无状态服务不再发布“打开、关闭、是否打开项目”三个旧 HTTP 操作,数据库连接在请求中按需从池借用。历史撤销日志已从数据库中移除,内部接口不再保留无效的兼容字段。真实库回归时发现五张明细表错误地把局部顺序号设成全局主键,已在 `tjwater_v2` 中改为父对象 ID 与 `sequence_no` 的复合主键。 -`tjwater_v2` 是 v2 业务库和时序库的正式物理库名。元数据库中的逻辑项目代码仍为 `tjwater_next`,由项目路由指向 `tjwater_v2`;两者不必同名。版本模板固定为 `tjwater_v2_template`,不按项目代码动态派生。目前该模板已从实际 v2 结构创建、清空项目数据、刷新空物化视图并封存,压缩后约 19 MB。 +`tjwater_v2` 是 v2 业务库和时序库的正式物理库名,元数据库中的逻辑项目代码也为 `tjwater_v2`;GeoServer 工作空间仍为 `tjwater_next`。每个物理业务库动态派生同名 `_template` 管网镜像,逻辑订阅只同步 `network` schema。独立的 `tjwater_v2_schema_template` 保持空数据并禁止普通连接,只提供统一 v2 结构。 + +五个源业务库(包括 `zjb`)都在各自的数据库命名空间内使用 publication `wndb_network_pub`,五个模板库都使用 subscription `wndb_network_sub`。PostgreSQL 实例级的 replication slot 按物理库唯一命名,例如 `tjwater_v2_network_slot`、`md_v2_network_slot` 和 `zjb_network_slot`,因此同名发布/订阅不会串库。逻辑复制只跟踪 32 张 `network` 表的 DML;GIS、SCADA、分析和其他业务数据不进入发布,DDL 与序列状态也需显式维护。 ## 当前目录 @@ -62,15 +64,17 @@ app/native/wndb/ `database.py` 提供 `ChangeSet`、`DatabaseCommand`、参数化查询和物化视图刷新。`DatabaseCommand` 只保存待执行的 SQL 和执行成功后返回给调用方的变更列表,不再生成或保存撤销 SQL。模型直接修改时按需刷新视图;批量命令在外层事务提交后只刷新一次。物化视图保留模型坐标 `x`、`y`,同时将供 GeoServer 使用的 `geom` 转换为 `EPSG:3857`,WNDB 查询不会把发布坐标误当成模型坐标。 -`projects.py` 只负责项目数据库的创建、复制、删除和异常安全的临时库上下文,不再保存“项目已打开”状态,也不混入模型查询。`postgres`、模板库、旧 WNDB 模板库 `project` 和元数据库均属于保护对象;模板复制源只能精确匹配 `WNDB_TEMPLATE_DB_NAME`,不能重新引入每项目 `_template`。批量清理不再扫描并删除服务器上的未知数据库,调用方必须显式提供每一个目标库名。数据库级 advisory lock 与 `datallowconn` 共同串行化多 worker 下的复制和删除。普通项目复制若复制源仍有其他会话会直接失败,不再主动终止正常请求。 +`projects.py` 只负责项目数据库的创建、复制、删除和异常安全的临时库上下文,不再保存“项目已打开”状态,也不混入模型查询。`postgres`、空结构模板、所有 `_template` 项目模板和元数据库属于保护对象。旧 WNDB 空模板库 `project` 已退出架构并从业务 PostgreSQL 与 TimescaleDB 实例删除。每个物理业务库对应一个同名 `_template` 管网模板,逻辑订阅只同步 `network` schema;空结构模板由 `WNDB_SCHEMA_TEMPLATE_DB_NAME` 配置,用于创建业务库和 INP 暂存库。批量清理不再扫描并删除服务器上的未知数据库,调用方必须显式提供目标。数据库级 advisory lock 与 `datallowconn` 串行化生命周期操作;克隆项目模板时仅临时禁止连接,完成后恢复,以便订阅工作进程继续同步。 -`model_replace.py` 在源库可重复读快照中读取 `network`、`gis` 基表,并按外键拓扑顺序复制到目标业务库。替换在单一事务内完成,不再删除并重建整个业务库;普通模型修改和整体替换共用同一项目级事务锁。INP 替换时,`analysis.results` 保留历史记录,只有新模型中不存在的元素引用会置空,`asset.scada_devices` 保留仍能匹配新节点或管段的设备;临时分析库则从当前项目复制模型和有效 SCADA 映射。 +`project_templates.py` 管理正式项目与其 `_template` 的一对一关系:先为源库 32 张 `network` 表创建 publication 和唯一 replication slot,再从空结构模板建立目标库、做一次一致性网络数据复制,最后以 `copy_data=false` 接续逻辑订阅。建库前检查复制 worker 余量,订阅的 worker、32 张关系状态均 ready 后才允许工作流继续。删除或失败回滚时先移除 subscription 和 slot,再删除模板库及 publication,不遗留 WAL slot。 + +`model_replace.py` 在源库可重复读快照中读取 `network`、`gis` 基表,并按外键拓扑顺序复制到目标业务库。替换在单一事务内完成,不再删除并重建整个业务库;普通模型修改和整体替换共用同一项目级事务锁。INP 暂存库与目标项目使用不同原始坐标 SRID 时,复制阶段保留坐标数值并将几何重新标记为目标列的 SRID,支持 `zjb` 等使用项目自定义工程坐标系的业务库。INP 替换时,`analysis.results` 保留历史记录,只有新模型中不存在的元素引用会置空,`asset.scada_devices` 保留仍能匹配新节点或管段的设备。临时分析库只需要订阅模板中的 `network` 数据;GIS 只是 INP 的可选展示章节,SCADA 也不是 EPANET 求解器的输入表。 ### model:管网模型和仿真配置 `model` 按业务实体命名,不再使用 `s2_junctions.py` 这类 INP 章节编号。节点、连接、模式、曲线、需求、规则和仿真设置都能从文件名直接定位。 -`options_v2.py` 和 `options_v3.py` 分别负责 EPANET V2、V3 的 `[OPTIONS]` 章节导入导出;数据库中的 `engine_version = 'legacy'` 仍表示 V2 配置,仅作为现有存储标识保留。 +`options_v2.py` 和 `options_v3.py` 分别负责 EPANET V2、V3 的 `[OPTIONS]` 章节导入导出;数据库中的 `engine_version = 'legacy'` 仍表示 V2 配置,仅作为现有存储标识保留。通过 V3 入口导入标准 EPANET INP 时会同时保存 legacy 原值和映射后的 V3 值,保证数据库再次导出的 V2 INP 不会引用模板遗留的 pattern。 每个实体模块保留三类紧密相关的函数:读取实体、生成并执行实体变更、转换该实体对应的一行或一段 INP 内容。完整文件的读取顺序、事务和项目生命周期由 `inp` 目录负责。因此,实体级编解码仍靠近实体定义,跨章节编排已经集中。 @@ -99,7 +103,7 @@ app/native/wndb/ `sections.py` 只保存 INP 章节名称和输出顺序。旧文件中混放的 `s1_title`、`s2_junction` 等命令类型常量已经移除。 -`importer.py` 负责文件分段、导入顺序、项目事务、版本转换和导入后的物化视图刷新。INP 更新先从 `tjwater_v2_template` 创建唯一暂存库并完成解析,再在当前业务库中事务替换模型表。模型提交后即使暂存库清理失败也仍会刷新物化视图;清理失败会记录日志,不再遮蔽主操作。ChangeSet 导入使用每请求唯一临时文件并在 `finally` 删除,避免同项目并发导入互相覆盖。`exporter.py` 负责按 EPANET 版本组织各章节并写出文件或 `ChangeSet`。 +`importer.py` 负责文件分段、导入顺序、项目事务、版本转换和导入后的物化视图刷新。INP 更新先从 `tjwater_v2_schema_template` 创建唯一空暂存库并完成解析,再在当前业务库中事务替换模型表。模型提交后即使暂存库清理失败也仍会刷新物化视图;清理失败会记录日志,不再遮蔽主操作。ChangeSet 导入使用每请求唯一临时文件并在 `finally` 删除,避免同项目并发导入互相覆盖。`exporter.py` 负责按 EPANET 版本组织各章节并写出文件或 `ChangeSet`。 ### commands:批量修改和级联关系 @@ -146,7 +150,9 @@ flowchart LR WNDB 当前使用同步 `psycopg` 连接池。`network/`、`components/` 和同步 EPANET 仿真接口统一声明为同步处理函数;公开 REST 路由的异步适配器把这些函数送入线程池,并把项目路由上下文传入工作线程,不会在事件循环线程上阻塞数据库或求解器。异步业务库和时序库访问使用带借用计数和代际切换的项目池:活跃旧池不会被 LRU 淘汰或强制关闭,配置变化后新请求立即使用新池,旧池在已有借用归还后关闭。元数据库保持独立 SQLAlchemy 异步池。 -临时分析库先由固定 `tjwater_v2_template` 提供结构,再从当前项目的一致性快照复制 `network/gis` 模型与 SCADA 映射并刷新物化视图。V3→V2 格式转换不需要项目模型,单独使用空模板临时库。旧 `online_Analysis.py`、restore 和 open/close 项目脚本已经删除,不再保留每项目模板与 operation 恢复入口。 +临时分析库直接从当前物理业务库的 `_template` 管网镜像创建,模板中的 `network` 数据由逻辑订阅维护。扩展仿真只在临时库修改管网参数、导出 INP 并调用 EPANET;GIS 和 SCADA 都不复制到临时库。当前 v2 项目的 SCADA 设备均为 `non_realtime`,扩展仿真使用接口显式传入的参数。V3→V2 格式转换不需要项目模型,单独使用空结构模板临时库。旧 `online_Analysis.py`、restore 和 open/close 项目脚本已经删除,不再保留 operation 恢复入口。 + +监测点选址、爆管定位和 DMA 漏损识别统一通过应用服务按请求导出唯一的临时 INP。算法运行期间文件有效,成功或异常退出时都会删除;不再复用按项目命名的固定 INP 缓存,避免模型更新后算法继续读取旧文件,也避免并发请求覆盖彼此的输入。 ## 依赖方向 @@ -196,7 +202,7 @@ WNDB 根包不再作为依赖汇聚点。上层若只需要管道查询,应直 ## 验证结果 - 本地 conda 环境单元、鉴权和 API 测试:313 项通过,2 项按条件跳过。 -- 一次性实库从 `tjwater_v2_template` 创建后,通过 INP 暂存解析和事务替换得到 11 个节点、13 条连接、11 条坐标及 9 条 junction 物化视图记录,验证后已完整删除。 +- 一次性实库从空结构模板创建后,通过 INP 暂存解析和事务替换得到 11 个节点、13 条连接、11 条坐标及 9 条 junction 物化视图记录,验证后已完整删除。 - `tjwater_v2` 统一视图覆盖 87,907 个节点和 91,054 条链路,与六个来源物化视图的合计数量一致。实测完整节点读取约 0.17 秒、完整链路读取约 0.10 秒、完整拓扑两次批量查询约 1.12 秒;耗时仅作为当前环境基线,不作为固定性能承诺。 -- `tjwater_v2` 真实数据库测试:11 项通过,覆盖业务库和时序库并发借用、失效连接自动重建、临时库模型/SCADA/视图完整克隆与清理、嵌套事务回滚、分析运行生命周期、恶意标识符转义、明细表复合主键、统一 GIS 查询视图,以及 WNDB pattern 增删改、级联解除需求关联和整体回滚。 +- `tjwater_v2` 真实数据库测试覆盖业务库和时序库并发借用、失效连接自动重建、只含 `network` 数据的临时库仿真与清理、嵌套事务回滚、分析运行生命周期、恶意标识符转义、明细表复合主键、统一 GIS 查询视图,以及 WNDB pattern 增删改、级联解除需求关联和整体回滚。 - Python 编译、未使用导入扫描、撤销字段残留扫描和 `git diff --check` 均通过。 diff --git a/tests/api/test_model_import_endpoints.py b/tests/api/test_model_import_endpoints.py index ae42670..5da68d2 100644 --- a/tests/api/test_model_import_endpoints.py +++ b/tests/api/test_model_import_endpoints.py @@ -1,6 +1,7 @@ import asyncio +from datetime import datetime, timezone from types import SimpleNamespace -from unittest.mock import AsyncMock +from unittest.mock import AsyncMock, Mock from uuid import uuid4 import pytest @@ -191,3 +192,114 @@ def test_committed_refresh_failure_is_not_wrapped_as_retryable_500( asyncio.run(model_import._apply_model_update(b"[TITLE]\n", "demo")) assert exc_info.value is error + + +def test_project_provision_creates_metadata_only_after_infrastructure(monkeypatch): + project_id = uuid4() + now = datetime(2026, 1, 1, tzinfo=timezone.utc) + project = SimpleNamespace( + id=project_id, + name="Demo", + code="demo", + description=None, + gs_workspace="demo", + map_extent={"bbox": [1.0, 2.0, 3.0, 4.0], "zoom": 15}, + status="active", + created_at=now, + updated_at=now, + ) + infrastructure = SimpleNamespace( + map_bbox=(1.0, 2.0, 3.0, 4.0), + business_dsn="postgresql://business/demo", + timescale_dsn="postgresql://timescale/demo", + model_template="demo_template", + layers=("junctions", "pipes"), + ) + repo = SimpleNamespace( + session=SimpleNamespace(rollback=AsyncMock()), + get_project_by_code=AsyncMock(return_value=None), + create_provisioned_project=AsyncMock(return_value=project), + ) + monkeypatch.setattr( + model_import, + "_run_uploaded_inp", + AsyncMock(return_value='{"simulation_result":"successful"}'), + ) + monkeypatch.setattr( + model_import, + "is_database_encryption_configured", + lambda: True, + ) + + async def fake_threadpool(function, *args, **kwargs): + assert function is model_import._provision_from_inp_sync + return infrastructure + + monkeypatch.setattr(model_import, "run_in_threadpool", fake_threadpool) + monkeypatch.setattr(model_import, "log_audit_event", AsyncMock()) + client = _client( + admin=SimpleNamespace(id=uuid4(), role="admin", is_superuser=True), + repo=repo, + ) + + response = client.post( + "/api/v1/admin/project-provisions", + data={"name": "Demo", "code": "demo", "map_zoom": "15"}, + files={"file": ("model.inp", VALID_INP)}, + ) + + assert response.status_code == 201 + assert response.json()["model_template_database"] == "demo_template" + repo.create_provisioned_project.assert_awaited_once() + assert repo.create_provisioned_project.await_args.kwargs["business_dsn"] == infrastructure.business_dsn + model_import.log_audit_event.assert_awaited_once() + + +def test_project_provision_cleans_infrastructure_when_metadata_commit_fails(monkeypatch): + cleanup = Mock(return_value=[]) + infrastructure = SimpleNamespace( + map_bbox=(1.0, 2.0, 3.0, 4.0), + business_dsn="postgresql://business/demo", + timescale_dsn="postgresql://timescale/demo", + model_template="demo_template", + layers=("junctions", "pipes"), + cleanup=cleanup, + ) + repo = SimpleNamespace( + session=SimpleNamespace(rollback=AsyncMock()), + get_project_by_code=AsyncMock(return_value=None), + create_provisioned_project=AsyncMock(side_effect=RuntimeError("metadata down")), + ) + monkeypatch.setattr( + model_import, + "_run_uploaded_inp", + AsyncMock(return_value='{"simulation_result":"successful"}'), + ) + monkeypatch.setattr( + model_import, + "is_database_encryption_configured", + lambda: True, + ) + + async def fake_threadpool(function, *args, **kwargs): + if function is model_import._provision_from_inp_sync: + return infrastructure + return function(*args, **kwargs) + + monkeypatch.setattr(model_import, "run_in_threadpool", fake_threadpool) + monkeypatch.setattr(model_import, "log_audit_event", AsyncMock()) + client = _client( + admin=SimpleNamespace(id=uuid4(), role="admin", is_superuser=True), + repo=repo, + ) + + response = client.post( + "/api/v1/admin/project-provisions", + data={"name": "Demo", "code": "demo"}, + files={"file": ("model.inp", VALID_INP)}, + ) + + assert response.status_code == 503 + repo.session.rollback.assert_awaited_once() + cleanup.assert_called_once_with() + model_import.log_audit_event.assert_not_awaited() diff --git a/tests/integration/test_database_pooling_live.py b/tests/integration/test_database_pooling_live.py index 5542cbd..d916bc7 100644 --- a/tests/integration/test_database_pooling_live.py +++ b/tests/integration/test_database_pooling_live.py @@ -7,6 +7,7 @@ import pytest from psycopg import connect from app.core.config import get_pgconn_string +from app.infra.epanet import run_project_return_dict from app.infra.db.dynamic_manager import ProjectConnectionManager from app.infra.db.postgresql.scada import ScadaInfoRepository from app.infra.db.timescaledb.sync_pool import timescale_connection @@ -52,7 +53,7 @@ def test_timeseries_pool_handles_concurrent_borrows() -> None: assert names == [PROJECT] * 64 -def test_temporary_project_clone_copies_model_scada_and_views() -> None: +def test_temporary_project_clone_runs_with_network_data_only() -> None: def counts(project: str) -> dict: with project_connection(project) as conn, conn.cursor() as cur: cur.execute( @@ -61,8 +62,8 @@ def test_temporary_project_clone_copies_model_scada_and_views() -> None: (select count(*) from network.nodes) as nodes, (select count(*) from network.links) as links, (select count(*) from asset.scada_devices) as scada, - (select count(*) from gis.junctions) as mv_junctions, - (select count(*) from gis.pipes) as mv_pipes + (select count(*) from gis.node_geometries) as node_geometries, + (select count(*) from gis.link_vertices) as link_vertices """ ) return dict(cur.fetchone()) @@ -70,7 +71,18 @@ def test_temporary_project_clone_copies_model_scada_and_views() -> None: source_counts = counts(PROJECT) temporary = None with temporary_project_database(PROJECT, "clone_validation") as temporary: - assert counts(temporary) == source_counts + temporary_counts = counts(temporary) + assert temporary_counts["nodes"] == source_counts["nodes"] + assert temporary_counts["links"] == source_counts["links"] + assert temporary_counts["scada"] == 0 + assert temporary_counts["node_geometries"] == 0 + assert temporary_counts["link_vertices"] == 0 + + result = run_project_return_dict(temporary) + output = result["output"] + assert result["simulation_result"] == "successful" + assert len(output["node_results"]) == source_counts["nodes"] + assert len(output["link_results"]) == source_counts["links"] assert temporary is not None assert have_project(temporary) is False diff --git a/tests/unit/test_geoserver_client.py b/tests/unit/test_geoserver_client.py new file mode 100644 index 0000000..2cf7791 --- /dev/null +++ b/tests/unit/test_geoserver_client.py @@ -0,0 +1,46 @@ +from xml.etree import ElementTree + +import httpx + +from app.infra.geoserver import client as geoserver_client + + +def test_project_workspace_sets_30_day_client_cache(monkeypatch): + cache_updates: list[str] = [] + + def handler(request: httpx.Request) -> httpx.Response: + path = request.url.path + if request.method == "GET" and path.endswith("/rest/workspaces/demo.json"): + return httpx.Response(404) + if request.method == "GET" and "/gwc/rest/layers/" in path: + return httpx.Response( + 200, + content=b"demo0", + ) + if request.method == "PUT" and "/gwc/rest/layers/" in path: + root = ElementTree.fromstring(request.content) + cache_updates.append(root.findtext("expireClients")) + return httpx.Response(200) + return httpx.Response(201) + + monkeypatch.setattr(geoserver_client.settings, "GEOSERVER_USERNAME", "admin") + monkeypatch.setattr(geoserver_client.settings, "GEOSERVER_PASSWORD", "secret") + monkeypatch.setattr( + geoserver_client.settings, + "GEOSERVER_CLIENT_CACHE_SECONDS", + 2_592_000, + ) + http_client = httpx.Client( + base_url="http://geoserver.example/geoserver", + transport=httpx.MockTransport(handler), + ) + client = geoserver_client.GeoServerAdminClient(client=http_client) + + layers = client.create_project_workspace( + workspace="demo", + database_name="demo", + map_bbox=(1.0, 2.0, 3.0, 4.0), + ) + + assert layers == geoserver_client.PROJECT_LAYER_NAMES + assert cache_updates == ["2592000"] * len(layers) diff --git a/tests/unit/test_metadata_repository_projects.py b/tests/unit/test_metadata_repository_projects.py index 0633843..6f79718 100644 --- a/tests/unit/test_metadata_repository_projects.py +++ b/tests/unit/test_metadata_repository_projects.py @@ -5,6 +5,7 @@ from uuid import uuid4 import pytest from sqlalchemy.dialects import postgresql +from app.infra.db.metadb.repositories import metadata_repository from app.infra.db.metadb.repositories.metadata_repository import MetadataRepository @@ -49,3 +50,38 @@ async def test_list_projects_for_user_only_queries_active_projects(): assert "user_project_membership.user_id" in sql assert "projects.status = 'active'" in sql assert [item.code for item in projects] == ["active-project"] + + +@pytest.mark.anyio +async def test_create_provisioned_project_commits_project_routes_and_membership_once( + monkeypatch, +): + added = [] + session = SimpleNamespace( + add=added.append, + commit=AsyncMock(), + refresh=AsyncMock(), + ) + monkeypatch.setattr( + metadata_repository, + "_encrypt_database_secret", + lambda value: f"encrypted::{value}", + ) + user_id = uuid4() + + project = await MetadataRepository(session).create_provisioned_project( + name="Demo", + code="demo", + description=None, + gs_workspace="demo", + map_extent={"bbox": [1, 2, 3, 4], "zoom": 14}, + creator_user_id=user_id, + business_dsn="postgresql://business/demo", + timescale_dsn="postgresql://timescale/demo", + ) + + assert project.status == "active" + assert [record.db_role for record in added[1:3]] == ["biz_data", "iot_data"] + assert added[3].user_id == user_id + session.commit.assert_awaited_once() + session.refresh.assert_awaited_once_with(project) diff --git a/tests/unit/test_model_replace.py b/tests/unit/test_model_replace.py new file mode 100644 index 0000000..7dc7503 --- /dev/null +++ b/tests/unit/test_model_replace.py @@ -0,0 +1,31 @@ +from app.native.wndb.core.model_replace import _copy_out_statement + + +def test_copy_out_statement_relabels_project_geometry_to_target_srid() -> None: + statement = _copy_out_statement( + "gis", + "labels", + ["id", "label", "geom"], + source_geometry_srids={"geom": 900914}, + target_geometry_srids={"geom": 990001}, + ) + + assert statement.as_string(None) == ( + 'copy (select "id", "label", ' + 'st_setsrid("geom", 990001) as "geom" ' + 'from "gis"."labels") to stdout' + ) + + +def test_copy_out_statement_keeps_direct_copy_when_srid_matches() -> None: + statement = _copy_out_statement( + "gis", + "node_geometries", + ["node_id", "geom"], + source_geometry_srids={"geom": 900914}, + target_geometry_srids={"geom": 900914}, + ) + + assert statement.as_string(None) == ( + 'copy "gis"."node_geometries" ("node_id", "geom") to stdout' + ) diff --git a/tests/unit/test_project_inp.py b/tests/unit/test_project_inp.py new file mode 100644 index 0000000..4be3008 --- /dev/null +++ b/tests/unit/test_project_inp.py @@ -0,0 +1,45 @@ +from pathlib import Path + +import pytest + +from app.services import project_inp + + +def test_temporary_project_inp_exports_unique_file_and_removes_it( + monkeypatch, tmp_path +) -> None: + exported_paths: list[Path] = [] + + def fake_dump_inp(project_code: str, path: str, version: str) -> None: + assert project_code == "demo" + assert version == "2" + exported_path = Path(path) + exported_path.write_text("[TITLE]\ndemo", encoding="utf-8") + exported_paths.append(exported_path) + + monkeypatch.setattr(project_inp, "PROJECT_INP_DIRECTORY", tmp_path) + monkeypatch.setattr(project_inp, "dump_inp", fake_dump_inp) + + with project_inp.temporary_project_inp( + "demo", purpose="sensor-placement" + ) as first_path: + assert first_path.read_text(encoding="utf-8") == "[TITLE]\ndemo" + with project_inp.temporary_project_inp( + "demo", purpose="sensor-placement" + ) as second_path: + assert second_path.exists() + + assert first_path != second_path + assert exported_paths == [first_path, second_path] + assert all(not path.exists() for path in exported_paths) + + +def test_temporary_project_inp_removes_empty_export(monkeypatch, tmp_path) -> None: + monkeypatch.setattr(project_inp, "PROJECT_INP_DIRECTORY", tmp_path) + monkeypatch.setattr(project_inp, "dump_inp", lambda *_args: None) + + with pytest.raises(ValueError, match="INP 导出失败"): + with project_inp.temporary_project_inp("missing", purpose="burst"): + pass + + assert list(tmp_path.iterdir()) == [] diff --git a/tests/unit/test_project_provisioning.py b/tests/unit/test_project_provisioning.py new file mode 100644 index 0000000..b77b5d0 --- /dev/null +++ b/tests/unit/test_project_provisioning.py @@ -0,0 +1,158 @@ +from pathlib import Path + +import pytest + +from app.services import project_provisioning + + +class _FakeGeoServer: + calls: list[tuple] = [] + + def __enter__(self): + return self + + def __exit__(self, *_args): + return None + + def check_ready(self): + self.calls.append(("geoserver_ready",)) + + def workspace_exists(self, workspace): + self.calls.append(("workspace_exists", workspace)) + return False + + def create_project_workspace(self, **kwargs): + self.calls.append(("create_workspace", kwargs)) + return ("junctions", "pipes") + + def delete_workspace(self, workspace): + self.calls.append(("delete_workspace", workspace)) + + +def _arrange_preflight(monkeypatch, calls): + schema_template = project_provisioning.settings.WNDB_SCHEMA_TEMPLATE_DB_NAME + monkeypatch.setattr( + project_provisioning, + "have_project", + lambda name: name == schema_template, + ) + monkeypatch.setattr( + project_provisioning, + "require_timescale_schema_template", + lambda: calls.append(("require_timescale_template",)), + ) + monkeypatch.setattr( + project_provisioning, + "ensure_replication_worker_capacity", + lambda: calls.append(("replication_capacity",)), + ) + monkeypatch.setattr( + project_provisioning, + "timescale_database_exists", + lambda _name: False, + ) + monkeypatch.setattr(project_provisioning, "GeoServerAdminClient", _FakeGeoServer) + + +def test_provision_project_creates_resources_in_dependency_order(monkeypatch): + calls: list[tuple] = [] + _FakeGeoServer.calls = calls + _arrange_preflight(monkeypatch, calls) + monkeypatch.setattr( + project_provisioning, + "create_project", + lambda name: calls.append(("create_business", name)), + ) + monkeypatch.setattr( + project_provisioning, + "network_update", + lambda path, name: calls.append(("import_model", Path(path).name, name)), + ) + monkeypatch.setattr( + project_provisioning, + "get_project_map_bbox", + lambda name: (1.0, 2.0, 3.0, 4.0), + ) + monkeypatch.setattr( + project_provisioning, + "create_project_model_template", + lambda name: calls.append(("create_model_template", name)) or "demo_template", + ) + monkeypatch.setattr( + project_provisioning, + "create_timescale_database", + lambda name: calls.append(("create_timescale", name)), + ) + + result = project_provisioning.provision_project_infrastructure( + code="demo", + workspace="demo", + inp_path="model.inp", + ) + + ordered_actions = [call[0] for call in calls] + assert ordered_actions.index("create_business") < ordered_actions.index("import_model") + assert ordered_actions.index("import_model") < ordered_actions.index("create_model_template") + assert ordered_actions.index("create_model_template") < ordered_actions.index("create_timescale") + assert ordered_actions.index("create_timescale") < ordered_actions.index("create_workspace") + assert result.model_template == "demo_template" + assert result.map_bbox == (1.0, 2.0, 3.0, 4.0) + + +def test_provision_project_rolls_back_created_resources_in_reverse_order(monkeypatch): + calls: list[tuple] = [] + _FakeGeoServer.calls = calls + _arrange_preflight(monkeypatch, calls) + monkeypatch.setattr( + project_provisioning, + "create_project", + lambda name: calls.append(("create_business", name)), + ) + monkeypatch.setattr(project_provisioning, "network_update", lambda *_args: None) + monkeypatch.setattr( + project_provisioning, + "get_project_map_bbox", + lambda _name: (1.0, 2.0, 3.0, 4.0), + ) + monkeypatch.setattr( + project_provisioning, + "create_project_model_template", + lambda name: calls.append(("create_model_template", name)) or "demo_template", + ) + monkeypatch.setattr( + project_provisioning, + "create_timescale_database", + lambda _name: (_ for _ in ()).throw(RuntimeError("timescale unavailable")), + ) + monkeypatch.setattr( + project_provisioning, + "delete_project_model_template", + lambda name: calls.append(("delete_model_template", name)), + ) + monkeypatch.setattr( + project_provisioning, + "delete_project", + lambda name: calls.append(("delete_business", name)), + ) + + with pytest.raises(project_provisioning.ProjectProvisioningError) as exc_info: + project_provisioning.provision_project_infrastructure( + code="demo", + workspace="demo", + inp_path="model.inp", + ) + + assert exc_info.value.stage == "timescaledb" + assert [call[0] for call in calls[-2:]] == [ + "delete_model_template", + "delete_business", + ] + + +@pytest.mark.parametrize( + "code", + ["Demo", "1demo", "demo-project", "demo_template", "postgres"], +) +def test_validate_project_code_rejects_unsafe_or_reserved_names(code): + with pytest.raises(ValueError): + project_provisioning.validate_project_code(code) diff --git a/tests/unit/test_project_routing.py b/tests/unit/test_project_routing.py index 11edf0a..03e2ac9 100644 --- a/tests/unit/test_project_routing.py +++ b/tests/unit/test_project_routing.py @@ -7,6 +7,7 @@ from app.infra.db.project_routing import ( get_active_project_routing, get_project_database_name, get_project_pgconn_string, + get_schema_template_database_name, get_project_template_database_name, get_project_timescale_pgconn_string, ) @@ -37,12 +38,12 @@ def test_project_database_uses_exact_routing_dsn_for_project_code() -> None: ) -def test_business_template_keeps_server_and_timescale_ignores_legacy_db_name( +def test_business_template_follows_physical_project_and_timescale_uses_routing( monkeypatch, ) -> None: monkeypatch.setattr( - "app.infra.db.project_routing.settings.WNDB_TEMPLATE_DB_NAME", - "tjwater_v2_template", + "app.infra.db.project_routing.settings.WNDB_SCHEMA_TEMPLATE_DB_NAME", + "tjwater_v2_schema_template", ) with activate_project_routing(_routing()): template_name = get_project_template_database_name("project_a") @@ -54,7 +55,7 @@ def test_business_template_keeps_server_and_timescale_ignores_legacy_db_name( assert business == { "user": "biz_user", "password": "biz_password", - "dbname": "tjwater_v2_template", + "dbname": "biz_database_template", "host": "biz.example", "port": "5432", "sslmode": "require", @@ -71,24 +72,32 @@ def test_business_template_keeps_server_and_timescale_ignores_legacy_db_name( def test_project_code_resolves_to_physical_business_database(monkeypatch) -> None: monkeypatch.setattr( - "app.infra.db.project_routing.settings.WNDB_TEMPLATE_DB_NAME", - "tjwater_v2_template", + "app.infra.db.project_routing.settings.WNDB_SCHEMA_TEMPLATE_DB_NAME", + "tjwater_v2_schema_template", ) with activate_project_routing(_routing()): assert get_project_database_name("project_a") == "biz_database" assert get_project_database_name("temporary_run") == "temporary_run" - assert get_project_template_database_name("project_a") == "tjwater_v2_template" + assert ( + get_project_template_database_name("project_a") + == "biz_database_template" + ) -def test_template_falls_back_to_config_outside_project_routing(monkeypatch) -> None: +def test_project_template_uses_database_name_outside_project_routing() -> None: + assert ( + get_project_template_database_name("physical_project") + == "physical_project_template" + ) + + +def test_schema_template_comes_from_config(monkeypatch) -> None: monkeypatch.setattr( - "app.infra.db.project_routing.settings.WNDB_TEMPLATE_DB_NAME", - "tjwater_v2_template", + "app.infra.db.project_routing.settings.WNDB_SCHEMA_TEMPLATE_DB_NAME", + "tjwater_v2_schema_template", ) - assert get_project_template_database_name("ignored-project-code") == ( - "tjwater_v2_template" - ) + assert get_schema_template_database_name() == "tjwater_v2_schema_template" def test_project_routing_is_nested_and_request_local() -> None: diff --git a/tests/unit/test_wndb_importer.py b/tests/unit/test_wndb_importer.py index 606c803..d96c70e 100644 --- a/tests/unit/test_wndb_importer.py +++ b/tests/unit/test_wndb_importer.py @@ -39,7 +39,9 @@ def test_read_inp_refreshes_after_committed_replace_when_cleanup_fails( monkeypatch.setattr( importer, "copy_project", - lambda _source, _target: events.append("copy"), + lambda source, _target, **kwargs: events.append( + f"copy:{source}:{kwargs.get('allow_template_source')}" + ), ) monkeypatch.setattr(importer, "project_transaction", fake_transaction) monkeypatch.setattr( @@ -63,4 +65,10 @@ def test_read_inp_refreshes_after_committed_replace_when_cleanup_fails( ) assert importer.read_inp("project_a", "model.inp") is True - assert events == ["copy", "parse", "replace", "cleanup-failed", "refresh"] + assert events == [ + "copy:tjwater_v2_schema_template:True", + "parse", + "replace", + "cleanup-failed", + "refresh", + ] diff --git a/tests/unit/test_wndb_options.py b/tests/unit/test_wndb_options.py new file mode 100644 index 0000000..6feb639 --- /dev/null +++ b/tests/unit/test_wndb_options.py @@ -0,0 +1,38 @@ +from app.native.wndb.model.options_v3 import _inp_in_option_v3, inp_in_option_v3 + + +def _operation_by_type(section: list[str], option_type: str) -> dict[str, str]: + result = _inp_in_option_v3(section) + return next(operation for operation in result.operations if operation["type"] == option_type) + + +def test_v3_import_preserves_legacy_options_from_standard_inp() -> None: + section = [ + "Units LPS", + "Pattern PAT_BASE", + "Quality None mg/L", + "Trials 80", + "Unbalanced Continue 10", + "CHECKFREQ 2", + ] + + legacy = _operation_by_type(section, "option") + v3 = _operation_by_type(section, "option_v3") + + assert legacy["PATTERN"] == "PAT_BASE" + assert legacy["QUALITY"] == "None mg/L" + assert legacy["TRIALS"] == "80" + assert legacy["UNBALANCED"] == "CONTINUE 10" + assert legacy["CHECKFREQ"] == "2" + assert v3["DEMAND_PATTERN"] == "PAT_BASE" + assert v3["QUALITY_MODEL"] == "NONE" + + statement = inp_in_option_v3(section) + assert "engine_version = 'legacy' and key = 'PATTERN'" in statement + assert "engine_version = 'v3' and key = 'DEMAND_PATTERN'" in statement + + +def test_v3_import_ignores_empty_option_lines() -> None: + legacy = _operation_by_type(["", " ", "Pattern PAT_BASE"], "option") + + assert legacy["PATTERN"] == "PAT_BASE" diff --git a/tests/unit/test_wndb_projects.py b/tests/unit/test_wndb_projects.py index 468015b..354af80 100644 --- a/tests/unit/test_wndb_projects.py +++ b/tests/unit/test_wndb_projects.py @@ -43,6 +43,15 @@ class _FakeConnection: return self._cursor +class _SequenceCursor(_FakeCursor): + def __init__(self, rows: list[dict]) -> None: + super().__init__() + self._fetch_rows = iter(rows) + + def fetchone(self): + return next(self._fetch_rows) + + def _admin_connection(cursor: _FakeCursor): @contextmanager def connection(): @@ -51,14 +60,22 @@ def _admin_connection(cursor: _FakeCursor): return connection +def _project_connection(cursor: _FakeCursor): + @contextmanager + def connection(_name): + yield _FakeConnection(cursor) + + return connection + + @pytest.mark.parametrize( "name", [ "postgres", - "project", "system_hub", "SYSTEM_HUB", "tjwater_v2_template", + "tjwater_v2_schema_template", "another_template", ], ) @@ -109,12 +126,12 @@ def test_create_project_allows_the_protected_template_as_source(monkeypatch) -> projects.create_project("project_a") - assert closed == ["tjwater_v2_template", "project_a"] - assert any(call[1] == ("tjwater_v2_template",) for call in cursor.calls) + assert closed == ["tjwater_v2_schema_template", "project_a"] + assert any(call[1] == ("tjwater_v2_schema_template",) for call in cursor.calls) assert any( isinstance(call[0], str) and call[0].startswith("select pg_terminate_backend") - and call[1] == ("tjwater_v2_template",) + and call[1] == ("tjwater_v2_schema_template",) for call in cursor.calls ) assert any("create database" in str(call[0]).lower() for call in cursor.calls) @@ -127,8 +144,7 @@ def test_list_project_excludes_metadata_database(monkeypatch) -> None: assert projects.list_project() == ["project_a"] excluded = cursor.calls[0][1][0] assert "system_hub" in excluded - assert "project" in excluded - assert "tjwater_v2_template" in excluded + assert "tjwater_v2_schema_template" in excluded def test_delete_project_uses_routed_physical_database_name(monkeypatch) -> None: @@ -177,6 +193,63 @@ def test_temporary_database_capacity_rejects_creation_at_limit( assert any("pg_advisory_unlock" in str(statement) for statement, _ in cursor.calls) +def test_project_model_template_requires_ready_active_subscription(monkeypatch) -> None: + cursor = _SequenceCursor( + [ + { + "subscriptions": 1, + "enabled_subscriptions": 1, + "active_workers": 1, + }, + {"relations": 32, "pending_relations": 0}, + ] + ) + closed: list[str] = [] + monkeypatch.setattr(projects, "project_connection", _project_connection(cursor)) + monkeypatch.setattr(projects, "close_project_pool", closed.append) + + projects._ensure_project_model_template_ready("project_a_template") + + assert closed == ["project_a_template"] + + +@pytest.mark.parametrize( + ("status", "relations", "message"), + [ + ( + { + "subscriptions": 1, + "enabled_subscriptions": 1, + "active_workers": 0, + }, + {"relations": 32, "pending_relations": 0}, + "subscription is not active", + ), + ( + { + "subscriptions": 1, + "enabled_subscriptions": 1, + "active_workers": 1, + }, + {"relations": 32, "pending_relations": 1}, + "still synchronizing", + ), + ], +) +def test_project_model_template_rejects_unready_subscription( + monkeypatch, status, relations, message +) -> None: + cursor = _SequenceCursor([status, relations]) + closed: list[str] = [] + monkeypatch.setattr(projects, "project_connection", _project_connection(cursor)) + monkeypatch.setattr(projects, "close_project_pool", closed.append) + + with pytest.raises(RuntimeError, match=message): + projects._ensure_project_model_template_ready("project_a_template") + + assert closed == ["project_a_template"] + + def test_temporary_project_database_cleans_up_after_failure(monkeypatch) -> None: calls: list[tuple[str, ...]] = [] monkeypatch.setattr( @@ -187,7 +260,9 @@ def test_temporary_project_database_cleans_up_after_failure(monkeypatch) -> None monkeypatch.setattr( projects, "copy_project", - lambda source, target: calls.append(("copy", source, target)), + lambda source, target, **kwargs: calls.append( + ("copy", source, target, str(kwargs.get("allow_template_source"))) + ), ) monkeypatch.setattr(projects, "have_project", lambda name: True) monkeypatch.setattr( @@ -195,26 +270,13 @@ def test_temporary_project_database_cleans_up_after_failure(monkeypatch) -> None "delete_project", lambda name: calls.append(("delete", name)), ) - monkeypatch.setattr( - "app.native.wndb.core.model_replace.replace_project_model", - lambda target, source, *, copy_source_scada: calls.append( - ("clone", target, source, str(copy_source_scada)) - ), - ) - monkeypatch.setattr( - "app.native.wndb.core.database.refresh_materialized_views_after_commit", - lambda name: calls.append(("refresh", name)), - ) - with pytest.raises(RuntimeError, match="analysis failed"): with projects.temporary_project_database("project_a", "age") as name: assert name == "isolated_run" raise RuntimeError("analysis failed") assert calls == [ - ("copy", "tjwater_v2_template", "isolated_run"), - ("clone", "isolated_run", "project_a", "True"), - ("refresh", "isolated_run"), + ("copy", "project_a_template", "isolated_run", "True"), ("delete", "isolated_run"), ] @@ -229,7 +291,9 @@ def test_temporary_template_database_does_not_clone_a_project(monkeypatch) -> No monkeypatch.setattr( projects, "copy_project", - lambda source, target: calls.append(("copy", source, target)), + lambda source, target, **kwargs: calls.append( + ("copy", source, target, str(kwargs.get("allow_template_source"))) + ), ) monkeypatch.setattr(projects, "have_project", lambda name: True) monkeypatch.setattr( @@ -242,7 +306,12 @@ def test_temporary_template_database_does_not_clone_a_project(monkeypatch) -> No assert name == "empty_conversion" assert calls == [ - ("copy", "tjwater_v2_template", "empty_conversion"), + ( + "copy", + "tjwater_v2_schema_template", + "empty_conversion", + "True", + ), ("delete", "empty_conversion"), ]