From 31e2728db128ae92316fa4cc666589fe2af48ea9 Mon Sep 17 00:00:00 2001 From: Jiang Date: Thu, 30 Jul 2026 11:01:45 +0800 Subject: [PATCH] refactor(api): unify scheme query endpoints --- app/api/v1/endpoints/burst_detection.py | 65 +---------- app/api/v1/endpoints/burst_location.py | 65 +---------- app/api/v1/endpoints/leakage.py | 67 +---------- app/api/v1/endpoints/schemes.py | 36 +++++- app/services/burst_location.py | 2 +- app/services/scheme_management.py | 121 +++++++++++++++++++- app/services/tjnetwork.py | 32 +++++- cli/tjwater_cli/commands_analysis.py | 42 +++++-- cli/tjwater_cli/registry.py | 12 +- cli/tjwater_cli_endpoint_scope.md | 6 +- tests/api/test_leakage_endpoints.py | 33 ------ tests/api/test_schemes_endpoints.py | 102 +++++++++++++++++ tests/unit/test_burst_location_service.py | 2 +- tests/unit/test_scheme_list_filter.py | 133 ++++++++++++++++++++++ 14 files changed, 457 insertions(+), 261 deletions(-) create mode 100644 tests/api/test_schemes_endpoints.py create mode 100644 tests/unit/test_scheme_list_filter.py diff --git a/app/api/v1/endpoints/burst_detection.py b/app/api/v1/endpoints/burst_detection.py index e1cb237..9d188d8 100644 --- a/app/api/v1/endpoints/burst_detection.py +++ b/app/api/v1/endpoints/burst_detection.py @@ -1,13 +1,11 @@ from datetime import datetime from typing import Any -from fastapi import APIRouter, Depends, HTTPException, Query, Path, Body +from fastapi import APIRouter, Depends, HTTPException, Body from pydantic import BaseModel, Field from app.auth.keycloak_dependencies import get_current_keycloak_username from app.services.burst_detection import ( - get_burst_detection_scheme_detail, - list_burst_detection_schemes, run_burst_detection, ) @@ -78,64 +76,3 @@ async def detect_burst( return run_burst_detection(**data.model_dump(), username=username) except Exception as exc: raise HTTPException(status_code=400, detail=str(exc)) - - -@router.get( - "/schemes/", - summary="查询爆管检测方案列表", - description="获取指定网络的所有爆管检测方案" -) -async def query_burst_detection_schemes( - network: str = Query(..., description="管网名称(或数据库名称)"), - query_date: datetime | None = Query(None, description="查询日期(可选)"), -) -> list[dict[str, Any]]: - """ - 获取爆管检测方案列表。 - - 查询指定网络的所有已配置的爆管检测方案, - 可按日期进行筛选。 - - Args: - network: 管网名称(或数据库名称) - query_date: 查询日期(可选) - - Returns: - 爆管检测方案列表 - - Raises: - HTTPException: 当查询失败时 - """ - try: - return list_burst_detection_schemes(network=network, query_date=query_date) - except Exception as exc: - raise HTTPException(status_code=400, detail=str(exc)) - - -@router.get( - "/schemes/{scheme_name}", - summary="获取爆管检测方案详情", - description="获取指定爆管检测方案的详细信息" -) -async def query_burst_detection_scheme_detail( - network: str = Query(..., description="管网名称(或数据库名称)"), - scheme_name: str = Path(..., description="爆管检测方案名称"), -) -> dict[str, Any]: - """ - 获取爆管检测方案详情。 - - 查询指定爆管检测方案的完整配置和参数信息。 - - Args: - network: 管网名称(或数据库名称) - scheme_name: 爆管检测方案名称 - - Returns: - 包含方案详情的字典 - - Raises: - HTTPException: 当查询失败时 - """ - try: - return get_burst_detection_scheme_detail(network=network, scheme_name=scheme_name) - except Exception as exc: - raise HTTPException(status_code=400, detail=str(exc)) diff --git a/app/api/v1/endpoints/burst_location.py b/app/api/v1/endpoints/burst_location.py index 0bb36e2..fa5995a 100644 --- a/app/api/v1/endpoints/burst_location.py +++ b/app/api/v1/endpoints/burst_location.py @@ -3,13 +3,11 @@ from datetime import datetime from typing import Literal -from fastapi import APIRouter, Depends, HTTPException, Query, Path, Body +from fastapi import APIRouter, Depends, HTTPException, Body from pydantic import BaseModel, Field from app.auth.keycloak_dependencies import get_current_keycloak_username from app.services.burst_location import ( - get_burst_location_scheme_detail, - list_burst_location_schemes, run_burst_location_by_network, ) @@ -68,64 +66,3 @@ async def locate_burst( return run_burst_location_by_network(**data.model_dump(), username=username) except (TypeError, ValueError) as exc: raise HTTPException(status_code=400, detail=str(exc)) - - -@router.get( - "/schemes/", - summary="查询爆管定位方案列表", - description="获取指定网络的所有爆管定位方案" -) -async def query_burst_schemes( - network: str = Query(..., description="管网名称(或数据库名称)"), - query_date: datetime | None = Query(None, description="查询日期(可选)") -) -> list[dict[str, Any]]: - """ - 获取爆管定位方案列表。 - - 查询指定网络的所有已配置的爆管定位方案, - 可按日期进行筛选。 - - Args: - network: 管网名称(或数据库名称) - query_date: 查询日期(可选) - - Returns: - 爆管定位方案列表 - - Raises: - HTTPException: 当查询失败时 - """ - try: - return list_burst_location_schemes(network=network, query_date=query_date) - except Exception as exc: - raise HTTPException(status_code=400, detail=str(exc)) - - -@router.get( - "/schemes/{scheme_name}", - summary="获取爆管定位方案详情", - description="获取指定爆管定位方案的详细信息" -) -async def query_burst_scheme_detail( - network: str = Query(..., description="管网名称(或数据库名称)"), - scheme_name: str = Path(..., description="爆管定位方案名称") -) -> dict[str, Any]: - """ - 获取爆管定位方案详情。 - - 查询指定爆管定位方案的完整配置和参数信息。 - - Args: - network: 管网名称(或数据库名称) - scheme_name: 爆管定位方案名称 - - Returns: - 包含方案详情的字典 - - Raises: - HTTPException: 当查询失败时 - """ - try: - return get_burst_location_scheme_detail(network=network, scheme_name=scheme_name) - except Exception as exc: - raise HTTPException(status_code=400, detail=str(exc)) diff --git a/app/api/v1/endpoints/leakage.py b/app/api/v1/endpoints/leakage.py index 0261e26..c57f95b 100644 --- a/app/api/v1/endpoints/leakage.py +++ b/app/api/v1/endpoints/leakage.py @@ -2,13 +2,11 @@ import os from typing import Any from datetime import datetime -from fastapi import APIRouter, Depends, HTTPException, Query, Path, Body +from fastapi import APIRouter, Depends, HTTPException, Body from pydantic import BaseModel, Field from app.auth.keycloak_dependencies import get_current_keycloak_username from app.services.leakage_identifier import ( - get_leakage_identify_scheme_detail, - list_leakage_identify_schemes, run_leakage_identification, ) @@ -68,66 +66,3 @@ async def identify_leakage( return run_leakage_identification(**data.model_dump(), username=username) except Exception as exc: raise HTTPException(status_code=400, detail=str(exc)) - - -@router.get( - "/schemes/", - summary="查询漏损识别方案列表", - description="获取指定网络的所有漏损识别方案" -) -async def query_leakage_schemes( - network: str = Query(..., description="管网名称(或数据库名称)"), - query_date: datetime | None = Query(None, description="查询日期(可选)") -) -> list[dict[str, Any]]: - """ - 获取漏损识别方案列表。 - - 查询指定网络的所有已配置的漏损识别方案, - 可按日期进行筛选。 - - Args: - network: 管网名称(或数据库名称) - query_date: 查询日期(可选) - - Returns: - 漏损识别方案列表 - - Raises: - HTTPException: 当查询失败时 - """ - try: - return list_leakage_identify_schemes(network=network, query_date=query_date) - except Exception as exc: - raise HTTPException(status_code=400, detail=str(exc)) - - -@router.get( - "/schemes/{scheme_name}", - summary="获取漏损识别方案详情", - description="获取指定漏损识别方案的详细信息" -) -async def query_leakage_scheme_detail( - network: str = Query(..., description="管网名称(或数据库名称)"), - scheme_name: str = Path(..., description="漏损识别方案名称") -) -> dict[str, Any]: - """ - 获取漏损识别方案详情。 - - 查询指定漏损识别方案的完整配置和参数信息。 - - Args: - network: 管网名称(或数据库名称) - scheme_name: 漏损识别方案名称 - - Returns: - 包含方案详情的字典 - - Raises: - HTTPException: 当查询失败时 - """ - try: - return get_leakage_identify_scheme_detail( - network=network, scheme_name=scheme_name - ) - except Exception as exc: - raise HTTPException(status_code=400, detail=str(exc)) diff --git a/app/api/v1/endpoints/schemes.py b/app/api/v1/endpoints/schemes.py index a08746e..7195ad4 100644 --- a/app/api/v1/endpoints/schemes.py +++ b/app/api/v1/endpoints/schemes.py @@ -1,6 +1,9 @@ -from fastapi import APIRouter, Query -from typing import Any, List, Dict +from datetime import datetime +from fastapi import APIRouter, HTTPException, Path, Query +from typing import Any from app.services.tjnetwork import get_scheme_schema, get_scheme, get_all_schemes +from app.services.scheme_management import query_scheme_detail +from app.services.time_api import extract_date router = APIRouter() @@ -24,10 +27,35 @@ async def fastapi_get_scheme(network: str = Query(..., description="管网名称 @router.get("/schemes", summary="获取所有方案", description="获取指定网络的所有方案信息") @router.get("/getallschemes/", summary="获取所有方案(旧路径)", description="获取指定网络的所有方案信息", deprecated=True) -async def fastapi_get_all_schemes(network: str = Query(..., description="管网名称(或数据库名称)")) -> list[dict[Any, Any]]: +async def fastapi_get_all_schemes( + network: str = Query(..., description="管网名称(或数据库名称)"), + scheme_type: str | None = Query(None, description="方案类型;为空时返回全部类型"), + query_date: datetime | None = Query(None, description="查询日期(可选)"), +) -> list[dict[Any, Any]]: """ 获取所有方案列表 返回指定网络中所有可用的方案 """ - return get_all_schemes(network) + parsed_date = ( + extract_date(query_date, field_name="query_date") + if query_date is not None + else None + ) + return get_all_schemes(network, scheme_type=scheme_type, query_date=parsed_date) + + +@router.get("/schemes/{scheme_name}", summary="获取方案详情", description="按方案类型获取指定方案详情") +async def fastapi_get_scheme_detail( + scheme_name: str = Path(..., description="方案名称"), + network: str = Query(..., description="管网名称(或数据库名称)"), + scheme_type: str | None = Query(None, description="方案类型;为空时返回通用方案详情"), +) -> dict[Any, Any]: + result = query_scheme_detail( + name=network, + scheme_name=scheme_name, + scheme_type=scheme_type, + ) + if not result: + raise HTTPException(status_code=404, detail=f"Scheme {scheme_name} not found") + return result diff --git a/app/services/burst_location.py b/app/services/burst_location.py index 049637b..589d122 100644 --- a/app/services/burst_location.py +++ b/app/services/burst_location.py @@ -476,7 +476,7 @@ def _get_simulation_scheme_burst_ids( ) -> list[str]: if not scheme_name: return [] - rows = query_scheme_list(network) or [] + rows = query_scheme_list(network, scheme_type=scheme_type) or [] for row in rows: if len(row) < 7: continue diff --git a/app/services/scheme_management.py b/app/services/scheme_management.py index 0bb1f11..aaa5c04 100644 --- a/app/services/scheme_management.py +++ b/app/services/scheme_management.py @@ -154,10 +154,16 @@ def delete_scheme_info(name: str, scheme_name: str) -> None: # 2025/03/23 -def query_scheme_list(name: str) -> list: +def query_scheme_list( + name: str, + scheme_type: str | None = None, + query_date: date | None = None, +) -> list: """ 查询pg数据库中的scheme_list,按照 create_time 降序排列,离现在时间最近的记录排在最前面 :param name: 项目名称(数据库名称) + :param scheme_type: 方案类型;为空时返回全部类型 + :param query_date: 查询日期;为空时不按日期过滤 :return: 返回查询结果的所有行 """ try: @@ -166,8 +172,38 @@ def query_scheme_list(name: str) -> list: # 连接到 PostgreSQL 数据库(这里是数据库 "bb") with psycopg.connect(conn_string) as conn: with conn.cursor() as cur: - # 按 create_time 降序排列 - cur.execute("SELECT * FROM scheme_list ORDER BY create_time DESC") + if scheme_type and query_date is not None: + cur.execute( + """ + SELECT * + FROM scheme_list + WHERE scheme_type = %s AND DATE(create_time) = %s + ORDER BY create_time DESC + """, + (scheme_type, query_date), + ) + elif scheme_type: + cur.execute( + """ + SELECT * + FROM scheme_list + WHERE scheme_type = %s + ORDER BY create_time DESC + """, + (scheme_type,), + ) + elif query_date is not None: + cur.execute( + """ + SELECT * + FROM scheme_list + WHERE DATE(create_time) = %s + ORDER BY create_time DESC + """, + (query_date,), + ) + else: + cur.execute("SELECT * FROM scheme_list ORDER BY create_time DESC") rows = cur.fetchall() return rows @@ -175,6 +211,85 @@ def query_scheme_list(name: str) -> list: print(f"查询错误:{e}") +def _filter_scheme_detail_scope( + result: dict, + name: str, + scheme_type: str | None = None, +) -> dict: + if not result: + return {} + if scheme_type and result.get("scheme_type") != scheme_type: + return {} + network = result.get("network") + if network not in (None, name): + return {} + return result + + +def query_scheme_detail( + name: str, + scheme_name: str, + scheme_type: str | None = None, +) -> dict: + if scheme_type == "dma_leak_identification": + return _filter_scheme_detail_scope( + query_leakage_identify_scheme_detail(name, scheme_name), + name, + scheme_type, + ) + if scheme_type == "burst_detection": + return _filter_scheme_detail_scope( + query_burst_detection_scheme_detail(name, scheme_name), + name, + scheme_type, + ) + if scheme_type == "burst_location": + return _filter_scheme_detail_scope( + query_burst_location_scheme_detail(name, scheme_name), + name, + scheme_type, + ) + + conn_string = get_pgconn_string(db_name=name) + with psycopg.connect(conn_string) as conn: + with conn.cursor() as cur: + if scheme_type: + cur.execute( + """ + SELECT scheme_id, scheme_name, scheme_type, username, create_time, scheme_start_time, scheme_detail + FROM public.scheme_list + WHERE scheme_name = %s AND scheme_type = %s + LIMIT 1 + """, + (scheme_name, scheme_type), + ) + else: + cur.execute( + """ + SELECT scheme_id, scheme_name, scheme_type, username, create_time, scheme_start_time, scheme_detail + FROM public.scheme_list + WHERE scheme_name = %s + LIMIT 1 + """, + (scheme_name,), + ) + row = cur.fetchone() + if row is None: + return {} + detail = row[6] if isinstance(row[6], dict) else {} + return _filter_scheme_detail_scope({ + "scheme_id": row[0], + "scheme_name": row[1], + "scheme_type": row[2], + "username": row[3], + "create_time": row[4], + "scheme_start_time": row[5], + "scheme_detail": detail, + "network": detail.get("network"), + "result_payload": detail.get("result_payload", {}), + }, name, scheme_type) + + def store_leakage_identify_result( name: str, scheme_name: str, diff --git a/app/services/tjnetwork.py b/app/services/tjnetwork.py index a8447e0..25f4a0f 100644 --- a/app/services/tjnetwork.py +++ b/app/services/tjnetwork.py @@ -1312,8 +1312,34 @@ def get_scheme_schema(name: str) -> dict[str, dict[str, Any]]: def get_scheme(name: str, schema_name: str) -> dict[str, Any]: return api.get_scheme(name, schema_name) -def get_all_schemes(name: str) -> list[dict[str, Any]]: - return api.get_all_schemes(name) +def get_all_schemes( + name: str, + scheme_type: str | None = None, + query_date: Any | None = None, +) -> list[dict[str, Any]]: + if scheme_type is None and query_date is None: + return api.get_all_schemes(name) + + from app.services.scheme_management import query_scheme_list + + rows = query_scheme_list(name, scheme_type=scheme_type, query_date=query_date) or [] + columns = [ + "scheme_id", + "scheme_name", + "scheme_type", + "username", + "create_time", + "scheme_start_time", + "scheme_detail", + ] + result = [] + for row in rows: + item = dict(zip(columns, row, strict=False)) + detail = item.get("scheme_detail") + if isinstance(detail, dict) and detail.get("network") not in (None, name): + continue + result.append(item) + return result ############################################################ # pipe_risk_probability 41 @@ -1345,5 +1371,3 @@ def get_all_sensor_placements(name: str) -> list[dict[Any, Any]]: def get_all_burst_locate_results(name: str) -> list[dict[Any, Any]]: return api.get_all_burst_locate_results(name) - - diff --git a/cli/tjwater_cli/commands_analysis.py b/cli/tjwater_cli/commands_analysis.py index c4b673a..d43b3eb 100644 --- a/cli/tjwater_cli/commands_analysis.py +++ b/cli/tjwater_cli/commands_analysis.py @@ -306,8 +306,11 @@ def analysis_leakage_schemes_list(ctx: typer.Context) -> None: ctx, summary="读取漏损方案列表成功", method="GET", - path="/leakage/schemes/", - params={"network": require_network(runtime)}, + path="/schemes", + params={ + "network": require_network(runtime), + "scheme_type": "dma_leak_identification", + }, require_auth=True, require_network_ctx=True, ) @@ -323,8 +326,11 @@ def analysis_leakage_schemes_get( ctx, summary="读取漏损方案详情成功", method="GET", - path=f"/leakage/schemes/{scheme_name}", - params={"network": require_network(runtime)}, + path=f"/schemes/{scheme_name}", + params={ + "network": require_network(runtime), + "scheme_type": "dma_leak_identification", + }, require_auth=True, require_network_ctx=True, ) @@ -362,8 +368,11 @@ def analysis_burst_detection_schemes_list(ctx: typer.Context) -> None: ctx, summary="读取爆管检测方案列表成功", method="GET", - path="/burst-detection/schemes/", - params={"network": require_network(runtime)}, + path="/schemes", + params={ + "network": require_network(runtime), + "scheme_type": "burst_detection", + }, require_auth=True, require_network_ctx=True, ) @@ -379,8 +388,11 @@ def analysis_burst_detection_schemes_get( ctx, summary="读取爆管检测方案详情成功", method="GET", - path=f"/burst-detection/schemes/{scheme_name}", - params={"network": require_network(runtime)}, + path=f"/schemes/{scheme_name}", + params={ + "network": require_network(runtime), + "scheme_type": "burst_detection", + }, require_auth=True, require_network_ctx=True, ) @@ -438,8 +450,11 @@ def analysis_burst_location_schemes_list(ctx: typer.Context) -> None: ctx, summary="读取爆管定位方案列表成功", method="GET", - path="/burst-location/schemes/", - params={"network": require_network(runtime)}, + path="/schemes", + params={ + "network": require_network(runtime), + "scheme_type": "burst_location", + }, require_auth=True, require_network_ctx=True, ) @@ -455,8 +470,11 @@ def analysis_burst_location_schemes_get( ctx, summary="读取爆管定位方案详情成功", method="GET", - path=f"/burst-location/schemes/{scheme_name}", - params={"network": require_network(runtime)}, + path=f"/schemes/{scheme_name}", + params={ + "network": require_network(runtime), + "scheme_type": "burst_location", + }, require_auth=True, require_network_ctx=True, ) diff --git a/cli/tjwater_cli/registry.py b/cli/tjwater_cli/registry.py index 540f84e..14a9e37 100644 --- a/cli/tjwater_cli/registry.py +++ b/cli/tjwater_cli/registry.py @@ -247,13 +247,13 @@ COMMAND_DOCS: dict[tuple[str, ...], CommandDoc] = { ("analysis", "leakage", "schemes", "list"): CommandDoc( path=("analysis", "leakage", "schemes", "list"), summary="列出漏损方案", - description="调用 /leakage/schemes/。", + description="调用 /schemes,并传入 scheme_type=dma_leak_identification。", examples=("tjwater-cli analysis leakage schemes list",), ), ("analysis", "leakage", "schemes", "get"): CommandDoc( path=("analysis", "leakage", "schemes", "get"), summary="读取漏损方案详情", - description="调用 /leakage/schemes/{scheme_name}。", + description="调用 /schemes/{scheme_name},并传入 scheme_type=dma_leak_identification。", examples=("tjwater-cli analysis leakage schemes get my_scheme",), ), ("analysis", "burst-detection", "detect"): CommandDoc( @@ -270,13 +270,13 @@ COMMAND_DOCS: dict[tuple[str, ...], CommandDoc] = { ("analysis", "burst-detection", "schemes", "list"): CommandDoc( path=("analysis", "burst-detection", "schemes", "list"), summary="列出爆管检测方案", - description="调用 /burst-detection/schemes/。", + description="调用 /schemes,并传入 scheme_type=burst_detection。", examples=("tjwater-cli analysis burst-detection schemes list",), ), ("analysis", "burst-detection", "schemes", "get"): CommandDoc( path=("analysis", "burst-detection", "schemes", "get"), summary="读取爆管检测方案详情", - description="调用 /burst-detection/schemes/{scheme_name}。", + description="调用 /schemes/{scheme_name},并传入 scheme_type=burst_detection。", examples=("tjwater-cli analysis burst-detection schemes get my_scheme",), ), ("analysis", "burst-location", "locate"): CommandDoc( @@ -303,13 +303,13 @@ COMMAND_DOCS: dict[tuple[str, ...], CommandDoc] = { ("analysis", "burst-location", "schemes", "list"): CommandDoc( path=("analysis", "burst-location", "schemes", "list"), summary="列出爆管定位方案", - description="调用 /burst-location/schemes/。", + description="调用 /schemes,并传入 scheme_type=burst_location。", examples=("tjwater-cli analysis burst-location schemes list",), ), ("analysis", "burst-location", "schemes", "get"): CommandDoc( path=("analysis", "burst-location", "schemes", "get"), summary="读取爆管定位方案详情", - description="调用 /burst-location/schemes/{scheme_name}。", + description="调用 /schemes/{scheme_name},并传入 scheme_type=burst_location。", examples=("tjwater-cli analysis burst-location schemes get my_scheme",), ), ("analysis", "risk", "pipe-now"): CommandDoc( diff --git a/cli/tjwater_cli_endpoint_scope.md b/cli/tjwater_cli_endpoint_scope.md index 5af650f..86da44a 100644 --- a/cli/tjwater_cli_endpoint_scope.md +++ b/cli/tjwater_cli_endpoint_scope.md @@ -202,11 +202,11 @@ app/api/v1/endpoints/risk.py | `tjwater-cli analysis contaminant --start-time TIME --duration SEC --source-node NODE --concentration VALUE --scheme SCHEME [--pattern PATTERN]` | `GET /contaminant-simulation` | 污染物模拟 | | `tjwater-cli analysis sensor-placement kmeans --count N` | `GET /pressuresensorplacementkmeans/` | 基于 kmeans 的传感器放置分析;不包含创建方案 | | `tjwater-cli analysis leakage identify --scheme SCHEME --start-time TIME --end-time TIME` | `POST /leakage/identify/` | 漏损识别 | -| `tjwater-cli analysis leakage schemes list\|get` | `GET /leakage/schemes/`、`GET /leakage/schemes/{scheme_name}` | 漏损方案查询 | +| `tjwater-cli analysis leakage schemes list\|get` | `GET /schemes?scheme_type=dma_leak_identification`、`GET /schemes/{scheme_name}?scheme_type=dma_leak_identification` | 漏损方案查询 | | `tjwater-cli analysis burst-detection detect --scheme SCHEME --start-time TIME --end-time TIME` | `POST /burst-detection/detect/` | 爆管检测 | -| `tjwater-cli analysis burst-detection schemes list\|get` | `GET /burst-detection/schemes/`、`GET /burst-detection/schemes/{scheme_name}` | 爆管检测方案查询 | +| `tjwater-cli analysis burst-detection schemes list\|get` | `GET /schemes?scheme_type=burst_detection`、`GET /schemes/{scheme_name}?scheme_type=burst_detection` | 爆管检测方案查询 | | `tjwater-cli analysis burst-location locate --scheme SCHEME --start-time TIME --end-time TIME` | `POST /burst-location/locate/` | 爆管定位 | -| `tjwater-cli analysis burst-location schemes list\|get` | `GET /burst-location/schemes/`、`GET /burst-location/schemes/{scheme_name}` | 爆管定位方案查询 | +| `tjwater-cli analysis burst-location schemes list\|get` | `GET /schemes?scheme_type=burst_location`、`GET /schemes/{scheme_name}?scheme_type=burst_location` | 爆管定位方案查询 | | `tjwater-cli analysis risk pipe-now --pipe PIPE` | `GET /getpiperiskprobabilitynow/` | 单条管道当前风险 | | `tjwater-cli analysis risk pipe-history --pipe PIPE` | `GET /getpiperiskprobability/` | 单条管道历史风险 | | `tjwater-cli analysis risk network` | `GET /getnetworkpiperiskprobabilitynow/`、`GET /getpiperiskprobabilitygeometries/` | 当前 project 全网风险 | diff --git a/tests/api/test_leakage_endpoints.py b/tests/api/test_leakage_endpoints.py index db8ec69..8be295b 100644 --- a/tests/api/test_leakage_endpoints.py +++ b/tests/api/test_leakage_endpoints.py @@ -1,7 +1,5 @@ from fastapi import FastAPI from fastapi.testclient import TestClient -from types import SimpleNamespace - from app.api.v1.endpoints import leakage as leakage_endpoint @@ -35,34 +33,3 @@ def test_identify_leakage_success(monkeypatch): ) assert response.status_code == 200 assert response.json()["area_count"] == 0 - - -def test_query_leakage_schemes_success(monkeypatch): - monkeypatch.setattr( - leakage_endpoint, - "list_leakage_identify_schemes", - lambda network, query_date=None: [ - {"scheme_name": "dma_001", "scheme_type": "dma_leak_identification"} - ], - ) - client = _build_client() - response = client.get("/api/v1/leakage/schemes/", params={"network": "demo"}) - assert response.status_code == 200 - assert response.json()[0]["scheme_name"] == "dma_001" - - -def test_query_leakage_scheme_detail_success(monkeypatch): - monkeypatch.setattr( - leakage_endpoint, - "get_leakage_identify_scheme_detail", - lambda network, scheme_name: { - "scheme_name": scheme_name, - "rows": [{"Area": "1", "LeakageFlow_m3_per_s": 0.1}], - }, - ) - client = _build_client() - response = client.get( - "/api/v1/leakage/schemes/dma_001", params={"network": "demo"} - ) - assert response.status_code == 200 - assert response.json()["scheme_name"] == "dma_001" diff --git a/tests/api/test_schemes_endpoints.py b/tests/api/test_schemes_endpoints.py new file mode 100644 index 0000000..319a5ad --- /dev/null +++ b/tests/api/test_schemes_endpoints.py @@ -0,0 +1,102 @@ +from datetime import date + +from fastapi import FastAPI +from fastapi.testclient import TestClient + +from app.api.v1.endpoints import schemes as schemes_endpoint + + +def _build_client() -> TestClient: + app = FastAPI() + app.include_router(schemes_endpoint.router, prefix="/api/v1") + return TestClient(app) + + +def test_get_schemes_forwards_optional_scheme_type(monkeypatch): + captured = {} + + def fake_get_all_schemes(network, scheme_type=None, query_date=None): + captured["network"] = network + captured["scheme_type"] = scheme_type + captured["query_date"] = query_date + return [ + { + "scheme_id": 1, + "scheme_name": "burst_case", + "scheme_type": scheme_type, + } + ] + + monkeypatch.setattr(schemes_endpoint, "get_all_schemes", fake_get_all_schemes) + + response = _build_client().get( + "/api/v1/schemes", + params={"network": "demo", "scheme_type": "burst_analysis"}, + ) + + assert response.status_code == 200 + assert captured == { + "network": "demo", + "scheme_type": "burst_analysis", + "query_date": None, + } + assert response.json()[0]["scheme_type"] == "burst_analysis" + + +def test_get_schemes_forwards_query_date(monkeypatch): + captured = {} + + def fake_get_all_schemes(network, scheme_type=None, query_date=None): + captured["network"] = network + captured["scheme_type"] = scheme_type + captured["query_date"] = query_date + return [] + + monkeypatch.setattr(schemes_endpoint, "get_all_schemes", fake_get_all_schemes) + + response = _build_client().get( + "/api/v1/schemes", + params={ + "network": "demo", + "scheme_type": "dma_leak_identification", + "query_date": "2026-01-02T00:00:00+08:00", + }, + ) + + assert response.status_code == 200 + assert captured == { + "network": "demo", + "scheme_type": "dma_leak_identification", + "query_date": date(2026, 1, 2), + } + + +def test_get_scheme_detail_forwards_scheme_type(monkeypatch): + captured = {} + + def fake_query_scheme_detail(name, scheme_name, scheme_type=None): + captured["name"] = name + captured["scheme_name"] = scheme_name + captured["scheme_type"] = scheme_type + return { + "scheme_name": scheme_name, + "scheme_type": scheme_type, + "rows": [{"Area": "1", "LeakageFlow_m3_per_s": 0.1}], + } + + monkeypatch.setattr( + schemes_endpoint, "query_scheme_detail", fake_query_scheme_detail + ) + + response = _build_client().get( + "/api/v1/schemes/dma_001", + params={"network": "demo", "scheme_type": "dma_leak_identification"}, + ) + + assert response.status_code == 200 + assert captured == { + "name": "demo", + "scheme_name": "dma_001", + "scheme_type": "dma_leak_identification", + } + assert response.json()["scheme_name"] == "dma_001" diff --git a/tests/unit/test_burst_location_service.py b/tests/unit/test_burst_location_service.py index a82bf4b..c80f813 100644 --- a/tests/unit/test_burst_location_service.py +++ b/tests/unit/test_burst_location_service.py @@ -197,7 +197,7 @@ def test_run_burst_location_uses_single_timerange_with_burst_source_split(monkey monkeypatch.setattr( module, "query_scheme_list", - lambda name: [ + lambda name, scheme_type=None: [ ( 1, "BurstSchemeA", diff --git a/tests/unit/test_scheme_list_filter.py b/tests/unit/test_scheme_list_filter.py new file mode 100644 index 0000000..567c1fb --- /dev/null +++ b/tests/unit/test_scheme_list_filter.py @@ -0,0 +1,133 @@ +from app.services import scheme_management, tjnetwork + + +class _FakeCursor: + def __init__(self): + self.calls = [] + + def __enter__(self): + return self + + def __exit__(self, *_exc_info): + return False + + def execute(self, statement, params=None): + self.calls.append((str(statement), params)) + + def fetchall(self): + return [] + + +class _FakeConnection: + def __init__(self, cursor): + self._cursor = cursor + + def __enter__(self): + return self + + def __exit__(self, *_exc_info): + return False + + def cursor(self): + return self._cursor + + +def test_query_scheme_list_pushes_scheme_type_into_sql(monkeypatch): + cursor = _FakeCursor() + monkeypatch.setattr( + scheme_management, "get_pgconn_string", lambda db_name=None: "postgres://test" + ) + monkeypatch.setattr( + scheme_management.psycopg, "connect", lambda _conn_string: _FakeConnection(cursor) + ) + + assert scheme_management.query_scheme_list("demo", scheme_type="burst_analysis") == [] + + statement, params = cursor.calls[0] + assert "WHERE scheme_type = %s" in statement + assert params == ("burst_analysis",) + + +def test_get_all_schemes_filters_central_scheme_list_by_type(monkeypatch): + captured = {} + + def fake_query_scheme_list(name, scheme_type=None, query_date=None): + captured["name"] = name + captured["scheme_type"] = scheme_type + captured["query_date"] = query_date + return [ + ( + 7, + "burst_case", + "burst_analysis", + "alice", + "2026-01-01T00:00:00+08:00", + "2026-01-01T01:00:00+08:00", + {"burst_ID": ["P1"]}, + ) + ] + + monkeypatch.setattr( + scheme_management, "query_scheme_list", fake_query_scheme_list + ) + + result = tjnetwork.get_all_schemes("demo", scheme_type="burst_analysis") + + assert captured == { + "name": "demo", + "scheme_type": "burst_analysis", + "query_date": None, + } + assert result == [ + { + "scheme_id": 7, + "scheme_name": "burst_case", + "scheme_type": "burst_analysis", + "username": "alice", + "create_time": "2026-01-01T00:00:00+08:00", + "scheme_start_time": "2026-01-01T01:00:00+08:00", + "scheme_detail": {"burst_ID": ["P1"]}, + } + ] + + +def test_query_scheme_detail_rejects_wrong_specialized_type(monkeypatch): + monkeypatch.setattr( + scheme_management, + "query_burst_detection_scheme_detail", + lambda name, scheme_name: { + "scheme_name": scheme_name, + "scheme_type": "burst_analysis", + "network": name, + }, + ) + + assert ( + scheme_management.query_scheme_detail( + "demo", + "same_name", + scheme_type="burst_detection", + ) + == {} + ) + + +def test_query_scheme_detail_rejects_wrong_network(monkeypatch): + monkeypatch.setattr( + scheme_management, + "query_burst_location_scheme_detail", + lambda name, scheme_name: { + "scheme_name": scheme_name, + "scheme_type": "burst_location", + "network": "other_network", + }, + ) + + assert ( + scheme_management.query_scheme_detail( + "demo", + "same_name", + scheme_type="burst_location", + ) + == {} + )