feat(projects): automate project infrastructure provisioning
Generic Container CI/CD / test-build-publish (push) Successful in 1m13s
Server CI/CD v2 / build-test-publish-and-deploy (push) Successful in 1m13s

This commit is contained in:
2026-09-11 10:57:51 +08:00
parent 10a7a66a41
commit 90b02057bc
38 changed files with 2345 additions and 189 deletions
+6
View File
@@ -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,
+174
View File
@@ -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="导入桌面端水力模型",
+12 -1
View File
@@ -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"
+9
View File
@@ -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)
@@ -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,
@@ -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"))
+10 -4
View File
@@ -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:
+165
View File
@@ -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}",),
)
+1
View File
@@ -0,0 +1 @@
"""GeoServer administration adapters."""
+224
View File
@@ -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}"
)
+105 -1
View File
@@ -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,
+191
View File
@@ -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)
+113 -23
View File
@@ -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:
+6 -2
View File
@@ -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)
+31 -8
View File
@@ -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
+15 -29
View File
@@ -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
+20 -38
View File
@@ -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
+35
View File
@@ -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)
+216
View File
@@ -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,
)
+22 -15
View File
@@ -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,