from typing import Any from datetime import datetime from typing import Literal from uuid import UUID from fastapi import APIRouter, Depends, HTTPException, Body from pydantic import BaseModel, Field from starlette.concurrency import run_in_threadpool from app.auth.keycloak_dependencies import get_current_keycloak_username from app.services.burst_location import ( run_burst_location_by_network, ) router = APIRouter() class BurstLocationRequest(BaseModel): """爆管定位请求模型""" network: str = Field(..., description="管网名称(或数据库名称)") data_source: Literal["monitoring", "simulation"] = Field("monitoring", description="数据来源:monitoring(监测)或simulation(模拟)") pressure_scada_ids: list[str] | None = Field(None, description="压力SCADA传感器ID列表") burst_pressure: dict[str, float] | list[dict[str, Any]] | None = Field(None, description="爆管时的压力数据") normal_pressure: dict[str, float] | list[dict[str, Any]] | None = Field(None, description="正常时的压力数据") burst_leakage: float = Field(..., description="爆管时的漏水量") flow_scada_ids: list[str] | None = Field(None, description="流量SCADA传感器ID列表") burst_flow: dict[str, float] | list[dict[str, Any]] | None = Field(None, description="爆管时的流量数据") normal_flow: dict[str, float] | list[dict[str, Any]] | None = Field(None, description="正常时的流量数据") min_dpressure: float = Field(2.0, description="最小压力差(bar)") basic_pressure: float = Field(10.0, description="基准压力(bar)") scada_burst_start: datetime | None = Field(None, description="爆管/模拟方案开始时间") scada_burst_end: datetime | None = Field(None, description="爆管/模拟方案结束时间") scada_normal_start: datetime | None = Field(None, description="监测数据正常工况开始时间") scada_normal_end: datetime | None = Field(None, description="监测数据正常工况结束时间") use_scada_flow: bool = Field(False, description="是否使用SCADA流量数据") scheme_name: str | None = Field(None, description="爆管定位运行名称") simulation_run_id: UUID | None = Field(None, description="分析模拟运行 ID") @router.post( "/burst-locations", summary="执行爆管定位", description="基于压力和流量数据定位管网中的爆管位置" ) async def locate_burst( data: BurstLocationRequest = Body(..., description="爆管定位请求数据"), username: str = Depends(get_current_keycloak_username), ) -> dict[str, Any]: """ 执行爆管定位分析。 使用压力和流量SCADA数据,通过对比爆管和正常状态下的数据差异, 定位管网中的爆管位置。 Args: data: 包含管网名称(或数据库名称)、压力、流量数据及相关参数的请求体 username: 当前认证用户名 Returns: 包含定位结果的字典 Raises: HTTPException: 当数据类型或值不正确时 """ try: return await run_in_threadpool( run_burst_location_by_network, **data.model_dump(), username=username, ) except (TypeError, ValueError) as exc: raise HTTPException(status_code=400, detail=str(exc))