feat(projects): automate project infrastructure provisioning
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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
@@ -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"
|
||||
|
||||
@@ -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"))
|
||||
@@ -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:
|
||||
|
||||
@@ -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}",),
|
||||
)
|
||||
@@ -0,0 +1 @@
|
||||
"""GeoServer administration adapters."""
|
||||
@@ -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}"
|
||||
)
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user